Skip to main content

mz_mysql_util/
tunnel.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10use 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/// Configures an optional tunnel for use when connecting to a MySQL
33/// database.
34#[derive(Debug, PartialEq, Clone)]
35pub enum TunnelConfig {
36    /// Establish a direct TCP connection to the database host.
37    /// If `resolved_ips` is not None, the provided IPs will be used
38    /// rather than resolving the hostname.
39    Direct {
40        resolved_ips: Option<BTreeSet<IpAddr>>,
41    },
42    /// Establish a TCP connection to the database via an SSH tunnel.
43    /// This means first establishing an SSH connection to a bastion host,
44    /// and then opening a separate connection from that host to the database.
45    /// This is commonly referred by vendors as a "direct SSH tunnel", in
46    /// opposition to "reverse SSH tunnel", which is currently unsupported.
47    Ssh { config: SshTunnelConfig },
48    /// Establish a TCP connection to the database via an AWS PrivateLink
49    /// service.
50    AwsPrivatelink {
51        /// The ID of the AWS PrivateLink service.
52        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);
59/// The `wait_timeout` to set on connections used during snapshotting, chosen
60/// to comfortably outlast even very long-running snapshots.
61pub const DEFAULT_SNAPSHOT_WAIT_TIMEOUT: Duration = Duration::from_secs(48 * 60 * 60);
62/// The minimum value MySQL accepts for `wait_timeout`.
63pub const MIN_SNAPSHOT_WAIT_TIMEOUT: Duration = Duration::from_secs(1);
64/// The maximum value MySQL accepts for `wait_timeout` on Windows, which is
65/// lower than the Unix maximum and so the portable upper bound.
66pub 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    // Snapshot-related configs
72    pub snapshot_max_execution_time: Option<Duration>,
73    pub snapshot_lock_wait_timeout: Option<Duration>,
74    pub snapshot_wait_timeout: Option<Duration>,
75
76    // Socket-related configs
77    pub tcp_keepalive: Option<Duration>,
78
79    // Connection timeout.  This timeout covers creating an authenticated connection
80    // (e.g. includes network connection, TLS handshake, authentication, etc.).
81    // If the connection has not been established in that time, it is considered an error.
82    pub connect_timeout: Option<Duration>,
83    // There are other timeout options on `mysql_async::OptsBuilder`
84    // (e.g. `conn_ttl` and `wait_timeout`) that could be exposed
85    // but they only apply to connection pools, which we are not currently using.
86}
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        // Verify values are within valid ranges
109        // Note we error log but do not fail as this is called in a non-fallible
110        // LD-sync in the adapter.
111
112        // https://dev.mysql.com/doc/refman/8.0/en/server-system-variables.html#sysvar_lock_wait_timeout
113        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        // https://dev.mysql.com/doc/refman/8.0/en/server-system-variables.html#sysvar_wait_timeout
124        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        // https://dev.mysql.com/doc/refman/8.0/en/server-system-variables.html#sysvar_max_execution_time
139        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    /// Apply relevant timeout configurations to a `mysql_async::OptsBuilder`.
183    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/// A MySQL connection with an optional SSH tunnel handle.
192///
193/// This wrapper is intended to be used in place of `mysql_async::Conn` to
194/// keep the SSH tunnel alive for the lifecycle of the connection by holding
195/// a reference to the tunnel handle.
196#[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/// Configuration for MySQL connections.
229///
230/// This wraps [`mysql_async::Opts`] to allow the configuration of a
231/// tunnel via a [`TunnelConfig`].
232#[derive(Clone, Debug)]
233pub struct Config {
234    inner: Opts,
235    tunnel: TunnelConfig,
236    // Whether to poll I/O for this connection in a tokio task
237    // TODO(roshan): Make this apply to queries on the returned connection, not just the initial
238    // connection.
239    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            // Cleartext plugin must be enabled for IAM authentication, for security,
306            // the network traffic is SSL/TLS encrypted.  The cleartext plugin is built
307            // into the MySQL client library.
308            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                // Override the connection host and port for the actual TCP connection to point to
341                // the local tunnel instead.
342                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                        // If the TLS configuration will validate the hostname, we need to set
349                        // the TLS hostname back to the actual upstream host and not the hostname
350                        // of the local SSH tunnel
351                        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                // Override the connection host for the actual TCP connection to point to
368                // the privatelink hostname instead.
369                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                        // If the TLS configuration will validate the hostname, we need to set
374                        // the TLS hostname back to the actual upstream host and not the
375                        // privatelink hostname.
376                        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        // NOTE: mysql_async panics on some server-controlled handshake input, for example an
397        // auth switch to `parsec`, which is a bare `panic!` without mysql_common's
398        // `client_parsec` feature. Outside a catch scope the panic hook aborts the process, so a
399        // server we connect to could take down environmentd or clusterd. The catch scope is
400        // task-local, so it has to wrap `Conn::new` inside the spawned task.
401        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                    // match instead of impl From<> for MySqlError so we can capture the timeout value
430                    TimeoutError::DeadlineElapsed => MySqlError::ConnectionTimeout(connect_timeout),
431                    TimeoutError::Inner(e) => e,
432                })
433        } else {
434            connection_future.await
435        }
436    }
437}