diff --git a/sqlx-mysql/src/connection/tls.rs b/sqlx-mysql/src/connection/tls.rs index dddb4c6aa5..57ac79e0cc 100644 --- a/sqlx-mysql/src/connection/tls.rs +++ b/sqlx-mysql/src/connection/tls.rs @@ -86,7 +86,7 @@ pub(super) async fn maybe_upgrade( tls::handshake( stream.socket.into_inner(), - &options.host, + options.tls_server_name.as_deref().unwrap_or(&options.host), connector, MapStream { server_version: stream.server_version, diff --git a/sqlx-mysql/src/options/mod.rs b/sqlx-mysql/src/options/mod.rs index 95413b6b11..506d626bc4 100644 --- a/sqlx-mysql/src/options/mod.rs +++ b/sqlx-mysql/src/options/mod.rs @@ -71,6 +71,7 @@ pub struct MySqlConnectOptions { pub(crate) username: String, pub(crate) password: Option, pub(crate) database: Option, + pub(crate) tls_server_name: Option, pub(crate) ssl_options: SslOptions, pub(crate) statement_cache_capacity: usize, pub(crate) charset: String, @@ -111,6 +112,7 @@ impl MySqlConnectOptions { database: None, charset: String::from("utf8mb4"), collation: None, + tls_server_name: None, ssl_options: SslOptions { ssl_mode: MySqlSslMode::Preferred, ssl_ca: None, @@ -138,6 +140,23 @@ impl MySqlConnectOptions { self } + /// Overrides the TLS server name used for SNI and hostname verification. + /// + /// By default, the host from `MySqlConnectOptions` is used. + /// + /// # Example + /// + /// ```rust + /// # use sqlx_mysql::MySqlConnectOptions; + /// let _options = MySqlConnectOptions::new() + /// .host("haproxy.example.com") + /// .tls_server_name("mysql.example.com"); + /// ``` + pub fn tls_server_name(mut self, server_name: &str) -> Self { + self.tls_server_name = Some(server_name.to_owned()); + self + } + /// Sets the port to connect to at the server host. /// /// The default port for MySQL is `3306`. @@ -560,3 +579,14 @@ impl MySqlConnectOptions { self.collation.as_deref() } } + +#[cfg(test)] +mod tests { + use super::MySqlConnectOptions; + + #[test] + fn tls_server_name_is_stored() { + let opts = MySqlConnectOptions::new().tls_server_name("sni.example.com"); + assert_eq!(opts.tls_server_name.as_deref(), Some("sni.example.com")); + } +}