1use aws_types::SdkConfig;
11use mysql_async::{Conn, Opts, OptsBuilder};
12use std::collections::BTreeSet;
13use std::future::Future;
14use std::net::IpAddr;
15use std::ops::{Deref, DerefMut};
16use std::panic::AssertUnwindSafe;
17use std::pin::Pin;
18use std::time::Duration;
19
20use mz_ore::future::{InTask, OreFutureExt, TimeoutError};
21use mz_ore::option::OptionExt;
22use mz_ore::task::spawn;
23use mz_repr::CatalogItemId;
24use mz_ssh_util::tunnel::{SshTimeoutConfig, SshTunnelConfig};
25use mz_ssh_util::tunnel_manager::{ManagedSshTunnelHandle, SshTunnelManager};
26use serde::{Deserialize, Serialize};
27use tracing::{error, info, warn};
28
29use crate::MySqlError;
30use crate::aws_rds::rds_auth_token;
31
32#[derive(Debug, PartialEq, Clone)]
35pub enum TunnelConfig {
36 Direct {
40 resolved_ips: Option<BTreeSet<IpAddr>>,
41 },
42 Ssh { config: SshTunnelConfig },
48 AwsPrivatelink {
51 connection_id: CatalogItemId,
53 },
54}
55
56pub const DEFAULT_TCP_KEEPALIVE: Duration = Duration::from_secs(60);
57pub const DEFAULT_SNAPSHOT_MAX_EXECUTION_TIME: Duration = Duration::ZERO;
58pub const DEFAULT_SNAPSHOT_LOCK_WAIT_TIMEOUT: Duration = Duration::from_secs(3600);
59pub const DEFAULT_SNAPSHOT_WAIT_TIMEOUT: Duration = Duration::from_secs(48 * 60 * 60);
62pub const MIN_SNAPSHOT_WAIT_TIMEOUT: Duration = Duration::from_secs(1);
64pub const MAX_SNAPSHOT_WAIT_TIMEOUT: Duration = Duration::from_secs(2147483);
67pub const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(60);
68
69#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
70pub struct TimeoutConfig {
71 pub snapshot_max_execution_time: Option<Duration>,
73 pub snapshot_lock_wait_timeout: Option<Duration>,
74 pub snapshot_wait_timeout: Option<Duration>,
75
76 pub tcp_keepalive: Option<Duration>,
78
79 pub connect_timeout: Option<Duration>,
83 }
87
88impl Default for TimeoutConfig {
89 fn default() -> Self {
90 Self {
91 snapshot_max_execution_time: Some(DEFAULT_SNAPSHOT_MAX_EXECUTION_TIME),
92 snapshot_lock_wait_timeout: Some(DEFAULT_SNAPSHOT_LOCK_WAIT_TIMEOUT),
93 snapshot_wait_timeout: Some(DEFAULT_SNAPSHOT_WAIT_TIMEOUT),
94 tcp_keepalive: Some(DEFAULT_TCP_KEEPALIVE),
95 connect_timeout: Some(DEFAULT_CONNECT_TIMEOUT),
96 }
97 }
98}
99
100impl TimeoutConfig {
101 pub fn build(
102 snapshot_max_execution_time: Duration,
103 snapshot_lock_wait_timeout: Duration,
104 snapshot_wait_timeout: Duration,
105 tcp_keepalive: Duration,
106 connect_timeout: Duration,
107 ) -> Self {
108 let snapshot_lock_wait_timeout = if snapshot_lock_wait_timeout.as_secs() > 31536000 {
114 error!(
115 "snapshot_lock_wait_timeout is too large: {}. Maximum is 31536000.",
116 snapshot_lock_wait_timeout.as_secs()
117 );
118 Some(DEFAULT_SNAPSHOT_LOCK_WAIT_TIMEOUT)
119 } else {
120 Some(snapshot_lock_wait_timeout)
121 };
122
123 let snapshot_wait_timeout = if snapshot_wait_timeout > MAX_SNAPSHOT_WAIT_TIMEOUT
125 || snapshot_wait_timeout < MIN_SNAPSHOT_WAIT_TIMEOUT
126 {
127 error!(
128 "snapshot_wait_timeout is out of range: {}. Must be within [{}, {}].",
129 snapshot_wait_timeout.as_secs(),
130 MIN_SNAPSHOT_WAIT_TIMEOUT.as_secs(),
131 MAX_SNAPSHOT_WAIT_TIMEOUT.as_secs(),
132 );
133 Some(DEFAULT_SNAPSHOT_WAIT_TIMEOUT)
134 } else {
135 Some(snapshot_wait_timeout)
136 };
137
138 let snapshot_max_execution_time = if snapshot_max_execution_time.as_millis() > 4294967295 {
140 error!(
141 "snapshot_max_execution_time is too large: {}. Maximum is 4294967295.",
142 snapshot_max_execution_time.as_secs()
143 );
144 Some(DEFAULT_SNAPSHOT_MAX_EXECUTION_TIME)
145 } else {
146 Some(snapshot_max_execution_time)
147 };
148
149 let tcp_keepalive = match u32::try_from(tcp_keepalive.as_millis()) {
150 Err(_) => {
151 error!(
152 "tcp_keepalive is too large: {}. Maximum is {}.",
153 tcp_keepalive.as_millis(),
154 u32::MAX,
155 );
156 Some(DEFAULT_TCP_KEEPALIVE)
157 }
158 Ok(_) => Some(tcp_keepalive),
159 };
160
161 let connect_timeout = match u32::try_from(connect_timeout.as_millis()) {
162 Err(_) => {
163 error!(
164 "connect_timeout is too large: {}. Maximum is {}.",
165 connect_timeout.as_millis(),
166 u32::MAX,
167 );
168 Some(DEFAULT_CONNECT_TIMEOUT)
169 }
170 Ok(_) => Some(connect_timeout),
171 };
172
173 Self {
174 snapshot_max_execution_time,
175 snapshot_lock_wait_timeout,
176 snapshot_wait_timeout,
177 tcp_keepalive,
178 connect_timeout,
179 }
180 }
181
182 pub fn apply_to_opts(&self, mut opts_builder: OptsBuilder) -> Result<OptsBuilder, MySqlError> {
184 if let Some(tcp_keepalive) = self.tcp_keepalive {
185 opts_builder = opts_builder.tcp_keepalive(Some(tcp_keepalive));
186 }
187 Ok(opts_builder)
188 }
189}
190
191#[derive(Debug)]
197pub struct MySqlConn {
198 conn: Conn,
199 _ssh_tunnel_handle: Option<ManagedSshTunnelHandle>,
200}
201
202impl Deref for MySqlConn {
203 type Target = Conn;
204
205 fn deref(&self) -> &Self::Target {
206 &self.conn
207 }
208}
209
210impl DerefMut for MySqlConn {
211 fn deref_mut(&mut self) -> &mut Self::Target {
212 &mut self.conn
213 }
214}
215
216impl MySqlConn {
217 pub async fn disconnect(mut self) -> Result<(), MySqlError> {
218 self.conn.disconnect().await?;
219 self._ssh_tunnel_handle.take();
220 Ok(())
221 }
222
223 pub fn take(self) -> (Conn, Option<ManagedSshTunnelHandle>) {
224 (self.conn, self._ssh_tunnel_handle)
225 }
226}
227
228#[derive(Clone, Debug)]
233pub struct Config {
234 inner: Opts,
235 tunnel: TunnelConfig,
236 in_task: InTask,
240 ssh_timeout_config: SshTimeoutConfig,
241 mysql_timeout_config: TimeoutConfig,
242 aws_config: Option<SdkConfig>,
243}
244
245impl Config {
246 pub fn new(
247 builder: OptsBuilder,
248 tunnel: TunnelConfig,
249 ssh_timeout_config: SshTimeoutConfig,
250 in_task: InTask,
251 mysql_timeout_config: TimeoutConfig,
252 aws_config: Option<SdkConfig>,
253 ) -> Result<Self, MySqlError> {
254 let opts = mysql_timeout_config.apply_to_opts(builder)?;
255 Ok(Self {
256 inner: opts.into(),
257 tunnel,
258 in_task,
259 ssh_timeout_config,
260 mysql_timeout_config,
261 aws_config,
262 })
263 }
264
265 pub async fn connect(
266 &self,
267 task_name: &str,
268 ssh_tunnel_manager: &SshTunnelManager,
269 ) -> Result<MySqlConn, MySqlError> {
270 let address = format!(
271 "mysql:://{}@{}:{}/{}",
272 self.inner.user().display_or("<unknown-user>"),
273 self.inner.ip_or_hostname(),
274 self.inner.tcp_port(),
275 self.inner.db_name().display_or("<unknown-dbname>"),
276 );
277 info!(%task_name, %address, "connecting");
278 match self.connect_internal(ssh_tunnel_manager).await {
279 Ok(t) => {
280 info!(%task_name, %address, "connected");
281 Ok(t)
282 }
283 Err(e) => {
284 warn!(%task_name, %address, "connection failed: {e:#}");
285 Err(e)
286 }
287 }
288 }
289
290 fn address(&self) -> (&str, u16) {
291 (self.inner.ip_or_hostname(), self.inner.tcp_port())
292 }
293
294 async fn connect_internal(
295 &self,
296 ssh_tunnel_manager: &SshTunnelManager,
297 ) -> Result<MySqlConn, MySqlError> {
298 let mut opts_builder = OptsBuilder::from_opts(self.inner.clone());
299
300 if let Some(aws_config) = &self.aws_config {
301 let (host, port) = self.address();
302 let username = self.inner.user().expect("MySQL: username required");
303
304 let token = rds_auth_token(host, port, username, aws_config).await?;
305 opts_builder = opts_builder
309 .pass(Some(token.to_string()))
310 .enable_cleartext_plugin(true);
311 }
312
313 match &self.tunnel {
314 TunnelConfig::Direct { resolved_ips } => {
315 opts_builder = opts_builder.resolved_ips(
316 resolved_ips
317 .clone()
318 .map(|ips| ips.into_iter().collect::<Vec<_>>()),
319 );
320
321 Ok(MySqlConn {
322 conn: self.connect_with_timeout(opts_builder).await?,
323 _ssh_tunnel_handle: None,
324 })
325 }
326 TunnelConfig::Ssh { config } => {
327 let (host, port) = self.address();
328 let tunnel = ssh_tunnel_manager
329 .connect(
330 config.clone(),
331 host,
332 port,
333 self.ssh_timeout_config,
334 self.in_task,
335 )
336 .await
337 .map_err(MySqlError::Ssh)?;
338
339 let tunnel_addr = tunnel.local_addr();
340 opts_builder = opts_builder
343 .ip_or_hostname(tunnel_addr.ip().to_string())
344 .tcp_port(tunnel_addr.port());
345
346 if let Some(ssl_opts) = self.inner.ssl_opts() {
347 if !ssl_opts.skip_domain_validation() {
348 opts_builder = opts_builder.ssl_opts(Some(
352 ssl_opts.clone().with_danger_tls_hostname_override(Some(
353 self.inner.ip_or_hostname().to_string(),
354 )),
355 ));
356 }
357 }
358
359 Ok(MySqlConn {
360 conn: self.connect_with_timeout(opts_builder).await?,
361 _ssh_tunnel_handle: Some(tunnel),
362 })
363 }
364 TunnelConfig::AwsPrivatelink { connection_id } => {
365 let privatelink_host = mz_cloud_resources::vpc_endpoint_name(*connection_id);
366
367 let mut opts_builder = opts_builder.ip_or_hostname(privatelink_host);
370
371 if let Some(ssl_opts) = self.inner.ssl_opts() {
372 if !ssl_opts.skip_domain_validation() {
373 opts_builder = opts_builder.ssl_opts(Some(
377 ssl_opts.clone().with_danger_tls_hostname_override(Some(
378 self.inner.ip_or_hostname().to_string(),
379 )),
380 ));
381 }
382 }
383
384 Ok(MySqlConn {
385 conn: self.connect_with_timeout(opts_builder).await?,
386 _ssh_tunnel_handle: None,
387 })
388 }
389 }
390 }
391
392 async fn connect_with_timeout(
393 &self,
394 opts_builder: OptsBuilder,
395 ) -> Result<mysql_async::Conn, MySqlError> {
396 let connect = async move {
402 match AssertUnwindSafe(Conn::new(opts_builder))
403 .ore_catch_unwind()
404 .await
405 {
406 Ok(result) => result.map_err(MySqlError::from),
407 Err(payload) => {
408 let message = payload
409 .downcast_ref::<&str>()
410 .map(|s| s.to_string())
411 .or_else(|| payload.downcast_ref::<String>().cloned())
412 .unwrap_or_else(|| "unknown panic".to_string());
413 error!("mysql connection attempt panicked: {message}");
414 Err(MySqlError::ConnectionPanicked(message))
415 }
416 }
417 };
418 let connection_future: Pin<Box<dyn Future<Output = Result<Conn, MySqlError>> + Send>> =
419 if let InTask::Yes = self.in_task {
420 Box::pin(spawn(|| "mysql_connect".to_string(), connect).abort_on_drop())
421 } else {
422 Box::pin(connect)
423 };
424
425 if let Some(connect_timeout) = self.mysql_timeout_config.connect_timeout {
426 mz_ore::future::timeout(connect_timeout, connection_future)
427 .await
428 .map_err(|err| match err {
429 TimeoutError::DeadlineElapsed => MySqlError::ConnectionTimeout(connect_timeout),
431 TimeoutError::Inner(e) => e,
432 })
433 } else {
434 connection_future.await
435 }
436 }
437}