Skip to main content

mz_storage_types/
connections.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
10//! Connection types.
11
12use std::borrow::Cow;
13use std::collections::{BTreeMap, BTreeSet};
14use std::fmt;
15use std::net::SocketAddr;
16use std::sync::Arc;
17use std::time::SystemTime;
18
19use anyhow::{Context, anyhow};
20use async_trait::async_trait;
21use aws_credential_types::provider::{ProvideCredentials, SharedCredentialsProvider};
22use aws_sigv4::http_request::{SignableBody, SignableRequest, SigningSettings, sign};
23use aws_sigv4::sign::v4;
24// Aliased to avoid colliding with `mz_ccsr::tls::Identity`.
25use aws_smithy_runtime_api::client::identity::Identity as AwsIdentity;
26use base64::Engine;
27use http::{HeaderMap, HeaderName, HeaderValue};
28use iceberg::Catalog;
29use iceberg::CatalogBuilder;
30use iceberg::TableIdent;
31use iceberg::io::{
32    GCS_CREDENTIALS_JSON, GCS_DISABLE_CONFIG_LOAD, GCS_DISABLE_VM_METADATA, GCS_USER_PROJECT,
33    S3_ACCESS_KEY_ID, S3_DISABLE_EC2_METADATA, S3_REGION, S3_SECRET_ACCESS_KEY,
34};
35use iceberg_catalog_rest::{
36    OAuth2TokenProvider, REST_CATALOG_PROP_URI, REST_CATALOG_PROP_WAREHOUSE, RequestAuthenticator,
37    RestCatalogBuilder, TokenProvider,
38};
39use iceberg_storage_opendal::{
40    AwsCredential, CustomAwsCredentialLoader, CustomAzdlsCredentialLoader,
41    CustomGcsCredentialLoader, OpenDalStorageFactory, ProvideCredential,
42};
43use itertools::Itertools;
44use mz_ccsr::tls::{Certificate, Identity};
45use mz_cloud_resources::{AwsExternalIdPrefix, CloudResourceReader, vpc_endpoint_host};
46use mz_dyncfg::ConfigSet;
47use mz_kafka_util::client::{
48    BrokerAddr, BrokerRewrite, HostMappingRules, MzClientContext, MzKafkaError, TunnelConfig,
49    TunnelingClientContext,
50};
51use mz_mysql_util::{MySqlConn, MySqlError};
52use mz_ore::assert_none;
53use mz_ore::error::ErrorExt;
54use mz_ore::future::{InTask, OreFutureExt};
55use mz_ore::netio::resolve_address;
56use mz_ore::num::NonNeg;
57use mz_ore::str::StrExt;
58use mz_repr::{CatalogItemId, GlobalId};
59use mz_secrets::SecretsReader;
60use mz_sql_parser::ast::ConnectionRulePattern;
61use mz_ssh_util::keys::SshKeyPair;
62use mz_ssh_util::tunnel::SshTunnelConfig;
63use mz_ssh_util::tunnel_manager::{ManagedSshTunnelHandle, SshTunnelManager};
64use mz_tracing::CloneableEnvFilter;
65use rdkafka::ClientContext;
66use rdkafka::config::FromClientConfigAndContext;
67use rdkafka::consumer::{BaseConsumer, Consumer};
68use regex::Regex;
69use reqsign_core::time::Timestamp;
70use reqwest_0_12::Request;
71use serde::{Deserialize, Deserializer, Serialize};
72use tokio::net;
73use tokio::runtime::Handle;
74use tokio_postgres::config::SslMode;
75use tracing::{debug, info, warn};
76use url::Url;
77
78use crate::AlterCompatible;
79use crate::configuration::StorageConfiguration;
80use crate::connections::aws::{
81    AwsAuth, AwsConnection, AwsConnectionReference, AwsConnectionValidationError,
82};
83use crate::connections::gcp::{GcpConnectionReference, GcpTokenProvider};
84use crate::connections::string_or_secret::StringOrSecret;
85use crate::controller::AlterError;
86use crate::dyncfgs::{
87    ENFORCE_EXTERNAL_ADDRESSES, KAFKA_CLIENT_ID_ENRICHMENT_RULES,
88    KAFKA_DEFAULT_AWS_PRIVATELINK_ENDPOINT_IDENTIFICATION_ALGORITHM, KAFKA_RECONNECT_BACKOFF,
89    KAFKA_RECONNECT_BACKOFF_MAX, KAFKA_RETRY_BACKOFF, KAFKA_RETRY_BACKOFF_MAX,
90};
91use crate::errors::{ContextCreationError, CsrConnectError};
92
93pub mod aws;
94pub mod gcp;
95mod iceberg_credentials;
96pub mod inline;
97pub mod string_or_secret;
98
99/// The OAuth2 form field naming the scopes a token is requested for.
100///
101/// Materialize drives the OAuth2 exchange itself rather than through the `credential`,
102/// `oauth2-server-uri`, and `scope` catalog properties, so that one token object serves both
103/// catalog requests and storage-credential refreshes.
104const OAUTH2_PARAM_SCOPE: &str = "scope";
105
106const REST_CATALOG_PROP_OAUTH2_SERVER_URI: &str = "oauth2-server-uri";
107/// The prefix marking a catalog property that `iceberg-rust` turns into a header on every REST
108/// request, the same convention the Iceberg Java client uses.
109const REST_CATALOG_HEADER_PROP_PREFIX: &str = "header.";
110/// Requests catalog-vended storage credentials, carried as a header by way of
111/// [`REST_CATALOG_HEADER_PROP_PREFIX`].
112const REST_CATALOG_PROP_ACCESS_DELEGATION: &str = "header.X-Iceberg-Access-Delegation";
113
114/// A credential loader that wraps an aws-sdk-rust credentials provider for use with
115/// iceberg/OpenDAL. This allows us to provide refreshable credentials from the AWS SDK
116/// credential chain (including the full assume role chain) to OpenDAL's S3 implementation.
117///
118/// We use this instead of OpenDAL's built-in assume role support because Materialize
119/// has a runtime-defined credential chain (ambient → jump role → user role with external ID)
120/// that can't be expressed via OpenDAL's static configuration properties.
121#[derive(Debug)]
122struct AwsSdkCredentialLoader {
123    /// The underlying AWS SDK credentials provider. For assume role auth, this provider
124    /// already handles the full chain: ambient creds -> jump role -> user role.
125    provider: SharedCredentialsProvider,
126}
127
128impl AwsSdkCredentialLoader {
129    fn new(provider: SharedCredentialsProvider) -> Self {
130        Self { provider }
131    }
132}
133
134impl ProvideCredential for AwsSdkCredentialLoader {
135    type Credential = AwsCredential;
136
137    async fn provide_credential(
138        &self,
139        _ctx: &reqsign_core::Context,
140    ) -> reqsign_core::Result<Option<Self::Credential>> {
141        let creds = self.provider.provide_credentials().await.map_err(|e| {
142            warn!(
143                error = %e.display_with_causes(),
144                "failed to load AWS credentials for Iceberg FileIO from SDK provider"
145            );
146            reqsign_core::Error::credential_invalid(
147                "failed to load AWS credentials from SDK provider for Iceberg FileIO \
148                 (credential source may be temporarily unavailable)",
149            )
150            .with_source(e)
151        })?;
152
153        // Propagate the SDK's expiry whenever it reports one. reqsign treats a `None` expiry as
154        // "valid forever", so dropping it would leave OpenDAL signing with stale assume-role
155        // credentials rather than asking us for fresh ones.
156        let expires_in = creds.expiry().map(aws_expiry_to_timestamp).transpose()?;
157
158        Ok(Some(AwsCredential {
159            access_key_id: creds.access_key_id().to_string(),
160            secret_access_key: creds.secret_access_key().to_string(),
161            session_token: creds.session_token().map(|s| s.to_string()),
162            expires_in,
163        }))
164    }
165}
166
167/// Converts an AWS SDK credential expiry into reqsign's [`Timestamp`].
168///
169/// Both failure modes require a nonsensical expiry (before the Unix epoch, or beyond year
170/// 292278994), so they are reported as errors rather than silently dropped, which would make the
171/// credential look non-expiring.
172fn aws_expiry_to_timestamp(expiry: SystemTime) -> reqsign_core::Result<Timestamp> {
173    let millis = expiry
174        .duration_since(SystemTime::UNIX_EPOCH)
175        .map_err(|e| {
176            reqsign_core::Error::unexpected("AWS credential expiry precedes the Unix epoch")
177                .with_source(e)
178        })?
179        .as_millis();
180    let millis = i64::try_from(millis).map_err(|e| {
181        reqsign_core::Error::unexpected("AWS credential expiry overflows a millisecond timestamp")
182            .with_source(e)
183    })?;
184    Timestamp::from_millisecond(millis)
185}
186
187/// Signs each outgoing REST-catalog request with AWS SigV4.
188///
189/// Holds a [`SharedCredentialsProvider`] (not static `Credentials`) so each
190/// request signs with refreshable creds from Materialize's chain
191/// (ambient -> jump role -> user role w/ external ID).
192struct Sigv4Authenticator {
193    provider: SharedCredentialsProvider,
194    region: String,
195    /// The AWS signing name. `"s3tables"` for AWS S3 Tables REST catalog.
196    signing_name: String,
197}
198
199impl std::fmt::Debug for Sigv4Authenticator {
200    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
201        f.debug_struct("Sigv4Authenticator")
202            .field("region", &self.region)
203            .field("signing_name", &self.signing_name)
204            .finish_non_exhaustive()
205    }
206}
207
208fn sigv4_err(e: impl Into<anyhow::Error>) -> iceberg::Error {
209    iceberg::Error::new(iceberg::ErrorKind::DataInvalid, "AWS SigV4").with_source(e)
210}
211
212#[async_trait]
213impl RequestAuthenticator for Sigv4Authenticator {
214    async fn authenticate_request(&self, req: &mut Request) -> iceberg::Result<()> {
215        let creds = self
216            .provider
217            .provide_credentials()
218            .await
219            .map_err(sigv4_err)?;
220        let identity: AwsIdentity = creds.into();
221        let params = v4::SigningParams::builder()
222            .identity(&identity)
223            .region(&self.region)
224            .name(&self.signing_name)
225            .time(SystemTime::now())
226            .settings(SigningSettings::default())
227            .build()
228            .map_err(sigv4_err)?
229            .into();
230        let body: &[u8] = req
231            .body()
232            .map(|b| match b.as_bytes() {
233                Some(b) => Ok(b),
234                None => Err(iceberg::Error::new(
235                    iceberg::ErrorKind::FeatureUnsupported,
236                    "SigV4 Authenticator cannot sign a streaming request body.",
237                )),
238            })
239            .transpose()?
240            .unwrap_or_default();
241        let headers = req
242            .headers()
243            .iter()
244            .map(|(k, v)| {
245                Ok((
246                    k.as_str(),
247                    v.to_str().map_err(|_| {
248                        iceberg::Error::new(
249                            iceberg::ErrorKind::DataInvalid,
250                            format!("header '{}' value is not all visible ASCII", k),
251                        )
252                    })?,
253                ))
254            })
255            .collect::<iceberg::Result<Vec<(&str, &str)>>>()?;
256        let signable = SignableRequest::new(
257            req.method().as_str(),
258            req.url().as_str(),
259            headers.into_iter(),
260            SignableBody::Bytes(body),
261        )
262        .map_err(sigv4_err)?;
263        let (instructions, _sig) = sign(signable, &params).map_err(sigv4_err)?.into_parts();
264        let (new_headers, new_query) = instructions.into_parts();
265        for header in new_headers {
266            let mut value = HeaderValue::from_str(header.value()).map_err(sigv4_err)?;
267            value.set_sensitive(header.sensitive());
268            req.headers_mut()
269                .insert(HeaderName::from_static(header.name()), value);
270        }
271        if !new_query.is_empty() {
272            let url = req.url_mut();
273            let mut pairs = url.query_pairs_mut();
274            for (name, value) in new_query {
275                pairs.append_pair(name, &value);
276            }
277        }
278        Ok(())
279    }
280
281    // SigV4 is stateless: nothing to cache, invalidate, or refresh.
282    async fn invalidate_cache(&self) -> iceberg::Result<()> {
283        Ok(())
284    }
285    async fn regenerate_cache(&self) -> iceberg::Result<()> {
286        Ok(())
287    }
288}
289
290/// An extension trait for [`SecretsReader`]
291#[async_trait::async_trait]
292trait SecretsReaderExt {
293    /// `SecretsReader::read`, but optionally run in a task.
294    async fn read_in_task_if(
295        &self,
296        in_task: InTask,
297        id: CatalogItemId,
298    ) -> Result<Vec<u8>, anyhow::Error>;
299
300    /// `SecretsReader::read_string`, but optionally run in a task.
301    async fn read_string_in_task_if(
302        &self,
303        in_task: InTask,
304        id: CatalogItemId,
305    ) -> Result<String, anyhow::Error>;
306}
307
308#[async_trait::async_trait]
309impl SecretsReaderExt for Arc<dyn SecretsReader> {
310    async fn read_in_task_if(
311        &self,
312        in_task: InTask,
313        id: CatalogItemId,
314    ) -> Result<Vec<u8>, anyhow::Error> {
315        let sr = Arc::clone(self);
316        async move { sr.read(id).await }
317            .run_in_task_if(in_task, || "secrets_reader_read".to_string())
318            .await
319    }
320    async fn read_string_in_task_if(
321        &self,
322        in_task: InTask,
323        id: CatalogItemId,
324    ) -> Result<String, anyhow::Error> {
325        let sr = Arc::clone(self);
326        async move { sr.read_string(id).await }
327            .run_in_task_if(in_task, || "secrets_reader_read".to_string())
328            .await
329    }
330}
331
332/// Extra context to pass through when instantiating a connection for a source
333/// or sink.
334///
335/// Should be kept cheaply cloneable.
336#[derive(Debug, Clone)]
337pub struct ConnectionContext {
338    /// An opaque identifier for the environment in which this process is
339    /// running.
340    ///
341    /// The storage layer is intentionally unaware of the structure within this
342    /// identifier. Higher layers of the stack can make use of that structure,
343    /// but the storage layer should be oblivious to it.
344    pub environment_id: String,
345    /// The level for librdkafka's logs.
346    pub librdkafka_log_level: tracing::Level,
347    /// A prefix for an external ID to use for all AWS AssumeRole operations.
348    pub aws_external_id_prefix: Option<AwsExternalIdPrefix>,
349    /// The ARN for a Materialize-controlled role to assume before assuming
350    /// a customer's requested role for an AWS connection.
351    pub aws_connection_role_arn: Option<String>,
352    /// A secrets reader.
353    pub secrets_reader: Arc<dyn SecretsReader>,
354    /// A cloud resource reader, if supported in this configuration.
355    pub cloud_resource_reader: Option<Arc<dyn CloudResourceReader>>,
356    /// A manager for SSH tunnels.
357    pub ssh_tunnel_manager: SshTunnelManager,
358}
359
360impl ConnectionContext {
361    /// Constructs a new connection context from command line arguments.
362    ///
363    /// **WARNING:** it is critical for security that the `aws_external_id` be
364    /// provided by the operator of the Materialize service (i.e., via a CLI
365    /// argument or environment variable) and not the end user of Materialize
366    /// (e.g., via a configuration option in a SQL statement). See
367    /// [`AwsExternalIdPrefix`] for details.
368    pub fn from_cli_args(
369        environment_id: String,
370        startup_log_level: &CloneableEnvFilter,
371        aws_external_id_prefix: Option<AwsExternalIdPrefix>,
372        aws_connection_role_arn: Option<String>,
373        secrets_reader: Arc<dyn SecretsReader>,
374        cloud_resource_reader: Option<Arc<dyn CloudResourceReader>>,
375    ) -> ConnectionContext {
376        ConnectionContext {
377            environment_id,
378            librdkafka_log_level: mz_ore::tracing::crate_level(
379                &startup_log_level.clone().into(),
380                "librdkafka",
381            ),
382            aws_external_id_prefix,
383            aws_connection_role_arn,
384            secrets_reader,
385            cloud_resource_reader,
386            ssh_tunnel_manager: SshTunnelManager::default(),
387        }
388    }
389
390    /// Constructs a new connection context for usage in tests.
391    pub fn for_tests(secrets_reader: Arc<dyn SecretsReader>) -> ConnectionContext {
392        ConnectionContext {
393            environment_id: "test-environment-id".into(),
394            librdkafka_log_level: tracing::Level::INFO,
395            aws_external_id_prefix: Some(
396                AwsExternalIdPrefix::new_from_cli_argument_or_environment_variable(
397                    "test-aws-external-id-prefix",
398                )
399                .expect("infallible"),
400            ),
401            aws_connection_role_arn: Some(
402                "arn:aws:iam::123456789000:role/MaterializeConnection".into(),
403            ),
404            secrets_reader,
405            cloud_resource_reader: None,
406            ssh_tunnel_manager: SshTunnelManager::default(),
407        }
408    }
409}
410
411#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
412pub enum Connection<C: ConnectionAccess = InlinedConnection> {
413    Kafka(KafkaConnection<C>),
414    Csr(CsrConnection<C>),
415    GlueSchemaRegistry(GlueSchemaRegistryConnection<C>),
416    Postgres(PostgresConnection<C>),
417    Ssh(SshConnection),
418    Aws(AwsConnection),
419    AwsPrivatelink(AwsPrivatelinkConnection),
420    Gcp(gcp::GcpConnection),
421    MySql(MySqlConnection<C>),
422    SqlServer(SqlServerConnectionDetails<C>),
423    IcebergCatalog(IcebergCatalogConnection<C>),
424}
425
426impl<R: ConnectionResolver> IntoInlineConnection<Connection, R>
427    for Connection<ReferencedConnection>
428{
429    fn into_inline_connection(self, r: R) -> Connection {
430        match self {
431            Connection::Kafka(kafka) => Connection::Kafka(kafka.into_inline_connection(r)),
432            Connection::Csr(csr) => Connection::Csr(csr.into_inline_connection(r)),
433            Connection::GlueSchemaRegistry(glue) => {
434                Connection::GlueSchemaRegistry(glue.into_inline_connection(r))
435            }
436            Connection::Postgres(pg) => Connection::Postgres(pg.into_inline_connection(r)),
437            Connection::Ssh(ssh) => Connection::Ssh(ssh),
438            Connection::Aws(aws) => Connection::Aws(aws),
439            Connection::AwsPrivatelink(awspl) => Connection::AwsPrivatelink(awspl),
440            Connection::Gcp(gcp) => Connection::Gcp(gcp),
441            Connection::MySql(mysql) => Connection::MySql(mysql.into_inline_connection(r)),
442            Connection::SqlServer(sql_server) => {
443                Connection::SqlServer(sql_server.into_inline_connection(r))
444            }
445            Connection::IcebergCatalog(iceberg) => {
446                Connection::IcebergCatalog(iceberg.into_inline_connection(r))
447            }
448        }
449    }
450}
451
452impl<C: ConnectionAccess> Connection<C> {
453    /// Whether this connection should be validated by default on creation.
454    pub fn validate_by_default(&self) -> bool {
455        match self {
456            Connection::Kafka(conn) => conn.validate_by_default(),
457            Connection::Csr(conn) => conn.validate_by_default(),
458            Connection::GlueSchemaRegistry(conn) => conn.validate_by_default(),
459            Connection::Postgres(conn) => conn.validate_by_default(),
460            Connection::Ssh(conn) => conn.validate_by_default(),
461            Connection::Aws(conn) => conn.validate_by_default(),
462            Connection::AwsPrivatelink(conn) => conn.validate_by_default(),
463            Connection::Gcp(conn) => conn.validate_by_default(),
464            Connection::MySql(conn) => conn.validate_by_default(),
465            Connection::SqlServer(conn) => conn.validate_by_default(),
466            Connection::IcebergCatalog(conn) => conn.validate_by_default(),
467        }
468    }
469}
470
471impl Connection<InlinedConnection> {
472    /// Validates this connection by attempting to connect to the upstream system.
473    pub async fn validate(
474        &self,
475        id: CatalogItemId,
476        storage_configuration: &StorageConfiguration,
477    ) -> Result<(), ConnectionValidationError> {
478        match self {
479            Connection::Kafka(conn) => conn.validate(id, storage_configuration).await?,
480            Connection::Csr(conn) => conn.validate(id, storage_configuration).await?,
481            Connection::GlueSchemaRegistry(conn) => {
482                conn.validate(id, storage_configuration).await?
483            }
484            Connection::Postgres(conn) => {
485                conn.validate(id, storage_configuration).await?;
486            }
487            Connection::Ssh(conn) => conn.validate(id, storage_configuration).await?,
488            Connection::Aws(conn) => conn.validate(id, storage_configuration).await?,
489            Connection::AwsPrivatelink(conn) => conn.validate(id, storage_configuration).await?,
490            Connection::Gcp(conn) => conn.validate(id, storage_configuration).await?,
491            Connection::MySql(conn) => {
492                conn.validate(id, storage_configuration).await?;
493            }
494            Connection::SqlServer(conn) => {
495                conn.validate(id, storage_configuration).await?;
496            }
497            Connection::IcebergCatalog(conn) => conn.validate(id, storage_configuration).await?,
498        }
499        Ok(())
500    }
501
502    pub fn unwrap_kafka(self) -> <InlinedConnection as ConnectionAccess>::Kafka {
503        match self {
504            Self::Kafka(conn) => conn,
505            o => unreachable!("{o:?} is not a Kafka connection"),
506        }
507    }
508
509    pub fn unwrap_pg(self) -> <InlinedConnection as ConnectionAccess>::Pg {
510        match self {
511            Self::Postgres(conn) => conn,
512            o => unreachable!("{o:?} is not a Postgres connection"),
513        }
514    }
515
516    pub fn unwrap_mysql(self) -> <InlinedConnection as ConnectionAccess>::MySql {
517        match self {
518            Self::MySql(conn) => conn,
519            o => unreachable!("{o:?} is not a MySQL connection"),
520        }
521    }
522
523    pub fn unwrap_sql_server(self) -> <InlinedConnection as ConnectionAccess>::SqlServer {
524        match self {
525            Self::SqlServer(conn) => conn,
526            o => unreachable!("{o:?} is not a SQL Server connection"),
527        }
528    }
529
530    pub fn unwrap_aws(self) -> <InlinedConnection as ConnectionAccess>::Aws {
531        match self {
532            Self::Aws(conn) => conn,
533            o => unreachable!("{o:?} is not an AWS connection"),
534        }
535    }
536
537    pub fn unwrap_gcp(self) -> <InlinedConnection as ConnectionAccess>::Gcp {
538        match self {
539            Self::Gcp(conn) => conn,
540            o => unreachable!("{o:?} is not a GCP connection"),
541        }
542    }
543
544    pub fn unwrap_ssh(self) -> <InlinedConnection as ConnectionAccess>::Ssh {
545        match self {
546            Self::Ssh(conn) => conn,
547            o => unreachable!("{o:?} is not an SSH connection"),
548        }
549    }
550
551    pub fn unwrap_csr(self) -> <InlinedConnection as ConnectionAccess>::Csr {
552        match self {
553            Self::Csr(conn) => conn,
554            o => unreachable!("{o:?} is not a Kafka connection"),
555        }
556    }
557
558    pub fn unwrap_glue_schema_registry(
559        self,
560    ) -> <InlinedConnection as ConnectionAccess>::GlueSchemaRegistry {
561        match self {
562            Self::GlueSchemaRegistry(conn) => conn,
563            o => unreachable!("{o:?} is not an AWS Glue Schema Registry connection"),
564        }
565    }
566
567    pub fn unwrap_iceberg_catalog(self) -> <InlinedConnection as ConnectionAccess>::IcebergCatalog {
568        match self {
569            Self::IcebergCatalog(conn) => conn,
570            o => unreachable!("{o:?} is not an Iceberg catalog connection"),
571        }
572    }
573}
574
575/// An error returned by [`Connection::validate`].
576#[derive(thiserror::Error, Debug)]
577pub enum ConnectionValidationError {
578    #[error(transparent)]
579    Postgres(#[from] PostgresConnectionValidationError),
580    #[error(transparent)]
581    MySql(#[from] MySqlConnectionValidationError),
582    #[error(transparent)]
583    SqlServer(#[from] SqlServerConnectionValidationError),
584    #[error(transparent)]
585    Aws(#[from] AwsConnectionValidationError),
586    #[error(transparent)]
587    Gcp(#[from] gcp::GcpConnectionValidationError),
588    #[error(transparent)]
589    AwsPrivatelinkServiceName(#[from] InvalidAwsPrivatelinkServiceName),
590    #[error("{}", .0.display_with_causes())]
591    Other(#[from] anyhow::Error),
592}
593
594impl ConnectionValidationError {
595    /// Reports additional details about the error, if any are available.
596    pub fn detail(&self) -> Option<String> {
597        match self {
598            ConnectionValidationError::Postgres(e) => e.detail(),
599            ConnectionValidationError::MySql(e) => e.detail(),
600            ConnectionValidationError::SqlServer(e) => e.detail(),
601            ConnectionValidationError::Aws(e) => e.detail(),
602            ConnectionValidationError::Gcp(e) => e.detail(),
603            ConnectionValidationError::AwsPrivatelinkServiceName(_) => None,
604            ConnectionValidationError::Other(_) => None,
605        }
606    }
607
608    /// Reports a hint for the user about how the error could be fixed.
609    pub fn hint(&self) -> Option<String> {
610        match self {
611            ConnectionValidationError::Postgres(e) => e.hint(),
612            ConnectionValidationError::MySql(e) => e.hint(),
613            ConnectionValidationError::SqlServer(e) => e.hint(),
614            ConnectionValidationError::Aws(e) => e.hint(),
615            ConnectionValidationError::Gcp(e) => e.hint(),
616            ConnectionValidationError::AwsPrivatelinkServiceName(e) => Some(e.hint()),
617            ConnectionValidationError::Other(_) => None,
618        }
619    }
620}
621
622impl<C: ConnectionAccess> AlterCompatible for Connection<C> {
623    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
624        match (self, other) {
625            (Self::Aws(s), Self::Aws(o)) => s.alter_compatible(id, o),
626            (Self::AwsPrivatelink(s), Self::AwsPrivatelink(o)) => s.alter_compatible(id, o),
627            (Self::Gcp(s), Self::Gcp(o)) => s.alter_compatible(id, o),
628            (Self::Ssh(s), Self::Ssh(o)) => s.alter_compatible(id, o),
629            (Self::Csr(s), Self::Csr(o)) => s.alter_compatible(id, o),
630            (Self::Kafka(s), Self::Kafka(o)) => s.alter_compatible(id, o),
631            (Self::Postgres(s), Self::Postgres(o)) => s.alter_compatible(id, o),
632            (Self::MySql(s), Self::MySql(o)) => s.alter_compatible(id, o),
633            _ => {
634                tracing::warn!(
635                    "Connection incompatible:\nself:\n{:#?}\n\nother\n{:#?}",
636                    self,
637                    other
638                );
639                Err(AlterError { id })
640            }
641        }
642    }
643}
644
645/// Auth mechanism for Iceberg REST catalogs.
646#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
647pub enum IcebergCatalogAuth<C: ConnectionAccess = InlinedConnection> {
648    /// Use Iceberg catalog REST API's standard OAuth flow.
649    OAuth {
650        /// client_id:client_secret
651        credential: StringOrSecret,
652        /// OAuth2 scope
653        scope: Option<String>,
654        /// Where to exchange `credential` for a bearer token.
655        ///
656        /// `None` uses the endpoint the Iceberg REST specification defines relative to the
657        /// catalog URL, `<url>/v1/oauth/tokens`. Catalogs that host their token endpoint
658        /// elsewhere, or behind an auth gateway that will not serve an unauthenticated
659        /// exchange, need this override.
660        server_url: Option<String>,
661    },
662    Gcp(GcpConnectionReference<C>),
663}
664
665#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
666pub struct RestIcebergCatalog<C: ConnectionAccess = InlinedConnection> {
667    pub auth: IcebergCatalogAuth<C>,
668    /// The warehouse for REST catalogs
669    pub warehouse: Option<String>,
670    /// Which form of storage-access delegation to request from the catalog, if any.
671    ///
672    /// `None` means "do not request storage-access delegation".
673    /// If we do not have permission to request delegated access but request it anyway,
674    /// a catalog can reject our whole request,
675    /// even if we have our own storage credentials to fall back on.
676    pub access_delegation: Option<IcebergAccessDelegation>,
677    /// Which object store the catalog's tables live in.
678    ///
679    /// Defaults to S3. A REST catalog does not tell us this, so a table backed
680    /// by GCS or ADLS is unreadable until the connection says so.
681    pub storage_provider: IcebergStorageProvider,
682}
683
684/// The value Materialize sends in the Iceberg REST `X-Iceberg-Access-Delegation`
685/// header, naming how the catalog should grant access to table storage.
686#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
687pub enum IcebergAccessDelegation {
688    /// Ask the catalog to mint temporary, table-scoped storage credentials.
689    VendedCredentials,
690}
691
692impl IcebergAccessDelegation {
693    /// The header value, as spelled in the Iceberg REST specification.
694    pub fn as_header_value(&self) -> &'static str {
695        match self {
696            IcebergAccessDelegation::VendedCredentials => "vended-credentials",
697        }
698    }
699}
700
701#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
702pub struct S3TablesRestIcebergCatalog<C: ConnectionAccess = InlinedConnection> {
703    /// The AWS connection details, for s3tables
704    pub aws_connection: AwsConnectionReference<C>,
705    /// The warehouse for s3tables
706    pub warehouse: String,
707}
708
709impl<R: ConnectionResolver> IntoInlineConnection<IcebergCatalogAuth, R>
710    for IcebergCatalogAuth<ReferencedConnection>
711{
712    fn into_inline_connection(self, r: R) -> IcebergCatalogAuth {
713        match self {
714            IcebergCatalogAuth::Gcp(x) => IcebergCatalogAuth::Gcp(x.into_inline_connection(&r)),
715            IcebergCatalogAuth::OAuth {
716                credential,
717                scope,
718                server_url,
719            } => IcebergCatalogAuth::OAuth {
720                credential,
721                scope,
722                server_url,
723            },
724        }
725    }
726}
727
728impl<R: ConnectionResolver> IntoInlineConnection<RestIcebergCatalog, R>
729    for RestIcebergCatalog<ReferencedConnection>
730{
731    fn into_inline_connection(self, r: R) -> RestIcebergCatalog {
732        RestIcebergCatalog {
733            auth: self.auth.into_inline_connection(&r),
734            warehouse: self.warehouse,
735            access_delegation: self.access_delegation,
736            storage_provider: self.storage_provider,
737        }
738    }
739}
740
741impl<R: ConnectionResolver> IntoInlineConnection<S3TablesRestIcebergCatalog, R>
742    for S3TablesRestIcebergCatalog<ReferencedConnection>
743{
744    fn into_inline_connection(self, r: R) -> S3TablesRestIcebergCatalog {
745        S3TablesRestIcebergCatalog {
746            aws_connection: self.aws_connection.into_inline_connection(&r),
747            warehouse: self.warehouse,
748        }
749    }
750}
751
752#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
753pub enum IcebergCatalogType {
754    Rest,
755    S3TablesRest,
756}
757
758/// Which object store holds the data files of a REST catalog's tables.
759///
760/// The catalog protocol says nothing about this: a REST catalog hands back
761/// storage locations and credentials, and the client has to already know how to
762/// talk to that store. So it is configured per connection rather than
763/// discovered.
764#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
765pub enum IcebergStorageProvider {
766    S3,
767    Gcs,
768    Adls,
769}
770
771impl IcebergStorageProvider {
772    /// The name as spelled in SQL.
773    pub fn as_str(&self) -> &'static str {
774        match self {
775            IcebergStorageProvider::S3 => "s3",
776            IcebergStorageProvider::Gcs => "gcs",
777            IcebergStorageProvider::Adls => "adls",
778        }
779    }
780}
781
782impl Default for IcebergStorageProvider {
783    fn default() -> Self {
784        IcebergStorageProvider::S3
785    }
786}
787
788#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
789pub enum IcebergCatalogImpl<C: ConnectionAccess = InlinedConnection> {
790    Rest(RestIcebergCatalog<C>),
791    S3TablesRest(S3TablesRestIcebergCatalog<C>),
792}
793
794impl<R: ConnectionResolver> IntoInlineConnection<IcebergCatalogImpl, R>
795    for IcebergCatalogImpl<ReferencedConnection>
796{
797    fn into_inline_connection(self, r: R) -> IcebergCatalogImpl {
798        match self {
799            IcebergCatalogImpl::Rest(rest) => {
800                IcebergCatalogImpl::Rest(rest.into_inline_connection(r))
801            }
802            IcebergCatalogImpl::S3TablesRest(s3tables) => {
803                IcebergCatalogImpl::S3TablesRest(s3tables.into_inline_connection(r))
804            }
805        }
806    }
807}
808
809#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
810pub struct IcebergCatalogConnection<C: ConnectionAccess = InlinedConnection> {
811    /// The catalog impl impl of that catalog
812    pub catalog: IcebergCatalogImpl<C>,
813    /// Where the catalog is located
814    pub uri: reqwest_0_12::Url,
815}
816
817impl AlterCompatible for IcebergCatalogConnection {
818    fn alter_compatible(&self, id: GlobalId, _other: &Self) -> Result<(), AlterError> {
819        Err(AlterError { id })
820    }
821}
822
823impl<R: ConnectionResolver> IntoInlineConnection<IcebergCatalogConnection, R>
824    for IcebergCatalogConnection<ReferencedConnection>
825{
826    fn into_inline_connection(self, r: R) -> IcebergCatalogConnection {
827        IcebergCatalogConnection {
828            catalog: self.catalog.into_inline_connection(&r),
829            uri: self.uri,
830        }
831    }
832}
833
834impl<C: ConnectionAccess> IcebergCatalogConnection<C> {
835    fn validate_by_default(&self) -> bool {
836        true
837    }
838}
839
840impl IcebergCatalogConnection<InlinedConnection> {
841    /// Connects to the catalog.
842    ///
843    /// `table` names the table this handle will be used against. It is needed only to keep
844    /// catalog-vended storage credentials refreshed, which the REST specification scopes to a
845    /// single table. Passing `None` leaves the connection on whatever credentials the catalog
846    /// supplies at `loadTable` time, which expire.
847    pub async fn connect(
848        &self,
849        storage_configuration: &StorageConfiguration,
850        in_task: InTask,
851        table: Option<&TableIdent>,
852    ) -> Result<Arc<dyn Catalog>, anyhow::Error> {
853        match self.catalog {
854            IcebergCatalogImpl::S3TablesRest(ref s3tables) => {
855                // S3 Tables signs every request with SigV4 off a refreshable AWS provider, so it
856                // has no vended credential to keep alive.
857                self.connect_s3tables(s3tables, storage_configuration, in_task)
858                    .await
859            }
860            IcebergCatalogImpl::Rest(ref rest) => {
861                self.connect_rest(rest, storage_configuration, in_task, table)
862                    .await
863            }
864        }
865    }
866
867    pub fn catalog_type(&self) -> IcebergCatalogType {
868        match self.catalog {
869            IcebergCatalogImpl::S3TablesRest(_) => IcebergCatalogType::S3TablesRest,
870            IcebergCatalogImpl::Rest(_) => IcebergCatalogType::Rest,
871        }
872    }
873
874    pub fn s3tables_catalog(&self) -> Option<&S3TablesRestIcebergCatalog> {
875        match &self.catalog {
876            IcebergCatalogImpl::S3TablesRest(s3tables) => Some(s3tables),
877            IcebergCatalogImpl::Rest(_) => None,
878        }
879    }
880
881    pub fn rest_catalog(&self) -> Option<&RestIcebergCatalog> {
882        match &self.catalog {
883            IcebergCatalogImpl::Rest(rest) => Some(rest),
884            IcebergCatalogImpl::S3TablesRest(_) => None,
885        }
886    }
887
888    async fn connect_s3tables(
889        &self,
890        s3tables: &S3TablesRestIcebergCatalog,
891        storage_configuration: &StorageConfiguration,
892        in_task: InTask,
893    ) -> Result<Arc<dyn Catalog>, anyhow::Error> {
894        let secret_reader = &storage_configuration.connection_context.secrets_reader;
895        let aws_ref = &s3tables.aws_connection;
896
897        let aws_region = aws_ref
898            .connection
899            .region
900            .clone()
901            .unwrap_or_else(|| "us-east-1".to_string());
902
903        let mut props = vec![
904            (S3_REGION.to_string(), aws_region.clone()),
905            (S3_DISABLE_EC2_METADATA.to_string(), "true".to_string()),
906            (
907                REST_CATALOG_PROP_WAREHOUSE.to_string(),
908                s3tables.warehouse.clone(),
909            ),
910            (REST_CATALOG_PROP_URI.to_string(), self.uri.to_string()),
911        ];
912
913        let aws_auth = aws_ref.connection.auth.clone();
914
915        if let AwsAuth::Credentials(creds) = &aws_auth {
916            props.push((
917                S3_ACCESS_KEY_ID.to_string(),
918                creds
919                    .access_key_id
920                    .get_string(in_task, secret_reader)
921                    .await?,
922            ));
923            props.push((
924                S3_SECRET_ACCESS_KEY.to_string(),
925                secret_reader.read_string(creds.secret_access_key).await?,
926            ));
927        }
928
929        // Sign REST catalog requests with the Materialize AWS credential chain
930        // via a custom `RequestAuthenticator`. For AssumeRole auth, also feed
931        // the chain to OpenDAL's S3 loader so data-file IO uses the same creds.
932        //
933        // For AssumeRole auth, the provider serves cached credentials that a
934        // background task keeps fresh, so no request through this catalog ever
935        // waits on STS. The task lives as long as the catalog holds the
936        // provider.
937        let credentials_provider = match &aws_auth {
938            // NOTE: This branch never contacts the connection's ENDPOINT.
939            // REST requests go to the catalog URI and the STS calls use the
940            // SDK defaults. The endpoint is still validated so a forbidden
941            // one is rejected rather than silently ignored.
942            AwsAuth::AssumeRole(assume_role) => {
943                aws_ref.connection.validate_endpoint(
944                    ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
945                )?;
946                assume_role
947                    .prefetch_credentials(
948                        &storage_configuration.connection_context,
949                        aws_ref.connection_id,
950                        storage_configuration.config_set(),
951                        format!("aws-connection-{}", aws_ref.connection_id),
952                    )
953                    .await
954                    .with_context(|| {
955                        format!(
956                            "failed to initialize AssumeRole credentials for S3 Tables Iceberg \
957                             catalog (catalog uri: {}, warehouse: {})",
958                            self.uri, s3tables.warehouse
959                        )
960                    })?
961            }
962            AwsAuth::Credentials(_) => {
963                let aws_config = aws_ref
964                    .connection
965                    .load_sdk_config(
966                        &storage_configuration.connection_context,
967                        aws_ref.connection_id,
968                        in_task,
969                        ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
970                    )
971                    .await
972                    .with_context(|| {
973                        format!(
974                            "failed to load AWS SDK config for S3 Tables Iceberg catalog \
975                             (connection id: {}, auth method: {}, catalog uri: {}, warehouse: {})",
976                            aws_ref.connection_id,
977                            aws_ref.connection.auth_method(),
978                            self.uri,
979                            s3tables.warehouse
980                        )
981                    })?;
982                aws_config
983                    .credentials_provider()
984                    .ok_or_else(|| anyhow!("aws_config missing credentials provider"))?
985            }
986        };
987
988        let authenticator = Arc::new(Sigv4Authenticator {
989            provider: credentials_provider.clone(),
990            region: aws_region.clone(),
991            signing_name: "s3tables".to_string(),
992        });
993
994        // N.B. We're using the AWS credentials from the catalog connection for the storage layer
995        //   even though the sink comes with its own (unused) AWS credentials for storage.
996        let customized_credential_load = if matches!(aws_auth, AwsAuth::AssumeRole(_)) {
997            Some(CustomAwsCredentialLoader::new(AwsSdkCredentialLoader::new(
998                credentials_provider,
999            )))
1000        } else {
1001            None
1002        };
1003
1004        let storage_factory = Arc::new(OpenDalStorageFactory::S3 {
1005            customized_credential_load,
1006        });
1007
1008        let catalog = RestCatalogBuilder::default()
1009            .with_storage_factory(storage_factory)
1010            .with_authenticator(authenticator)
1011            .load("IcebergCatalog", props.into_iter().collect())
1012            .await
1013            .with_context(|| {
1014                format!(
1015                    "failed to create S3 Tables Iceberg catalog \
1016                     (connection id: {}, catalog uri: {}, warehouse: {})",
1017                    aws_ref.connection_id, self.uri, s3tables.warehouse
1018                )
1019            })?;
1020
1021        Ok(Arc::new(catalog))
1022    }
1023
1024    /// Collects the headers `iceberg-rust` puts on every catalog request out of the `header.*`
1025    /// props it takes them from.
1026    ///
1027    /// The credential endpoints Materialize calls directly bypass the catalog client, so without
1028    /// this they would reach the same server missing headers it may require, `x-goog-user-project`
1029    /// on a GCP-hosted catalog among them.
1030    fn catalog_headers(props: &BTreeMap<String, String>) -> Result<HeaderMap, anyhow::Error> {
1031        props
1032            .iter()
1033            .filter_map(|(k, v)| {
1034                k.strip_prefix(REST_CATALOG_HEADER_PROP_PREFIX)
1035                    .map(|name| (name, v))
1036            })
1037            .map(|(name, value)| {
1038                let name = HeaderName::try_from(name)
1039                    .with_context(|| format!("invalid Iceberg catalog header name: {name}"))?;
1040                let value = HeaderValue::try_from(value)
1041                    .with_context(|| format!("invalid Iceberg catalog header value for {name}"))?;
1042                Ok((name, value))
1043            })
1044            .collect()
1045    }
1046
1047    /// Resolves the endpoint that vends storage credentials for `table`, or `None` if this
1048    /// connection has no use for one.
1049    ///
1050    /// A loader built on the returned endpoint takes sole responsibility for storage credentials,
1051    /// so `None` is the signal to leave the catalog's static `storage-credentials` props in force.
1052    /// It means either that the connection did not ask for delegation, or that the caller named no
1053    /// table, and the specification scopes vended credentials to a single table.
1054    async fn vended_credential_endpoint(
1055        &self,
1056        rest: &RestIcebergCatalog,
1057        client: &reqwest_0_12::Client,
1058        token: &Arc<dyn TokenProvider>,
1059        headers: &HeaderMap,
1060        table: Option<&TableIdent>,
1061    ) -> Result<Option<Url>, anyhow::Error> {
1062        match (&rest.access_delegation, table) {
1063            (Some(IcebergAccessDelegation::VendedCredentials), Some(table)) => Ok(Some(
1064                iceberg_credentials::table_credentials_endpoint(
1065                    &self.uri,
1066                    client,
1067                    token,
1068                    headers,
1069                    rest.warehouse.as_deref(),
1070                    table,
1071                )
1072                .await?,
1073            )),
1074            _ => Ok(None),
1075        }
1076    }
1077
1078    /// Builds a GCS storage factory, refreshing vended credentials from `endpoint` if there is one.
1079    fn gcs_storage_factory(
1080        endpoint: Option<Url>,
1081        client: &reqwest_0_12::Client,
1082        token: &Arc<dyn TokenProvider>,
1083        headers: &HeaderMap,
1084    ) -> OpenDalStorageFactory {
1085        OpenDalStorageFactory::Gcs {
1086            customized_credential_load: endpoint.map(|endpoint| {
1087                CustomGcsCredentialLoader::new(iceberg_credentials::VendedCredentialLoader::new(
1088                    client.clone(),
1089                    endpoint,
1090                    Arc::clone(token),
1091                    headers.clone(),
1092                ))
1093            }),
1094        }
1095    }
1096
1097    async fn connect_rest(
1098        &self,
1099        rest: &RestIcebergCatalog,
1100        storage_configuration: &StorageConfiguration,
1101        in_task: InTask,
1102        table: Option<&TableIdent>,
1103    ) -> Result<Arc<dyn Catalog>, anyhow::Error> {
1104        let mut props = BTreeMap::from([(
1105            REST_CATALOG_PROP_URI.to_string(),
1106            self.uri.to_string().clone(),
1107        )]);
1108
1109        if let Some(warehouse) = &rest.warehouse {
1110            props.insert(REST_CATALOG_PROP_WAREHOUSE.to_string(), warehouse.clone());
1111        }
1112
1113        // One client for catalog requests, OAuth token requests, and credential refreshes, so all
1114        // three share a connection pool. `iceberg-rust` would otherwise default to its own.
1115        let client = reqwest_0_12::Client::new();
1116
1117        // Catalog auth is configured through a combination of `props` and `.with_authenticator(...)`,
1118        // which happen at different stages of the [`RestCatalogBuilder`] -> [`RestCatalog`]
1119        // construction pipeline.
1120        let (storage_factory, custom_authenticator) = match &rest.auth {
1121            IcebergCatalogAuth::OAuth {
1122                credential,
1123                scope,
1124                server_url,
1125            } => {
1126                let credential = credential
1127                    .get_string(
1128                        in_task,
1129                        &storage_configuration.connection_context.secrets_reader,
1130                    )
1131                    .await
1132                    .map_err(|e| anyhow!("failed to read Iceberg catalog credential: {e}"))?;
1133
1134                if let Some(server_url) = server_url {
1135                    // The OAuth2 exchange POSTs the catalog credential to this URL, so a URL
1136                    // aimed inside our own network turns the connection into a request forger
1137                    // against, say, a cloud metadata endpoint. Resolve it and reject private
1138                    // addresses, the same check every other host we dial directly gets.
1139                    //
1140                    // NOTE: the resolved addresses are only checked, not pinned. The catalog
1141                    // client offers no hook to dial a pre-resolved address, so a name that
1142                    // resolves differently between this check and the request slips through.
1143                    // Kafka and Confluent Schema Registry connections do pin theirs.
1144                    let url = Url::parse(server_url).with_context(|| {
1145                        format!("invalid OAUTH2 SERVER URL for Iceberg catalog: {server_url}")
1146                    })?;
1147                    let host = url.host_str().ok_or_else(|| {
1148                        anyhow!("OAUTH2 SERVER URL for Iceberg catalog has no host: {server_url}")
1149                    })?;
1150                    resolve_address(
1151                        host,
1152                        ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
1153                    )
1154                    .await
1155                    .with_context(|| {
1156                        format!("OAUTH2 SERVER URL for Iceberg catalog is not resolvable to an external address: {server_url}")
1157                    })?;
1158
1159                    props.insert(
1160                        REST_CATALOG_PROP_OAUTH2_SERVER_URI.to_string(),
1161                        server_url.clone(),
1162                    );
1163                }
1164
1165                // Materialize builds an OAuth2 provider shared across both catalog requests
1166                // and the vended credentials refresh below.
1167                let token_endpoint = match server_url {
1168                    Some(server_url) => server_url.clone(),
1169                    // Matches `iceberg-rust`'s default when no `oauth2-server-uri` is configured.
1170                    None => format!(
1171                        "{}/v1/oauth/tokens",
1172                        self.uri.as_str().trim_end_matches('/')
1173                    ),
1174                };
1175                let (client_id, client_secret) = match credential.split_once(':') {
1176                    Some((client_id, client_secret)) => {
1177                        (Some(client_id.to_string()), client_secret.to_string())
1178                    }
1179                    None => (None, credential),
1180                };
1181                let oauth_params = BTreeMap::from([(
1182                    OAUTH2_PARAM_SCOPE.to_string(),
1183                    // The default `iceberg-rust` applies when the connection names no scope.
1184                    scope.clone().unwrap_or_else(|| "catalog".to_string()),
1185                )]);
1186                let token: Arc<dyn TokenProvider> = Arc::new(OAuth2TokenProvider::new(
1187                    client.clone(),
1188                    client_id,
1189                    client_secret,
1190                    token_endpoint,
1191                    // The token request needs none of the catalog's headers, and
1192                    // `OAuth2TokenProvider` sets the form content type itself.
1193                    HeaderMap::new(),
1194                    oauth_params.into_iter().collect(),
1195                ));
1196
1197                let headers = Self::catalog_headers(&props)?;
1198                let endpoint = self
1199                    .vended_credential_endpoint(rest, &client, &token, &headers, table)
1200                    .await?;
1201
1202                (
1203                    // The catalog tells us where the data lives but not what
1204                    // kind of store it is, so the connection has to say.
1205                    match rest.storage_provider {
1206                        // When used with MinIO, Polaris returns a config with:
1207                        //   s3.access-key-id, s3.secret-access-key, s3.endpoint, ...
1208                        // `iceberg-rust` forwards these props to `opendal`. When the catalog
1209                        // vends instead, it returns per-table `storage-credentials` that
1210                        // `iceberg-rust` wires into the same FileIO.
1211                        // N.B. This is not confirmed to work with other catalog & storage implementations.
1212                        IcebergStorageProvider::S3 => OpenDalStorageFactory::S3 {
1213                            customized_credential_load: endpoint.map(|endpoint| {
1214                                CustomAwsCredentialLoader::new(
1215                                    iceberg_credentials::VendedCredentialLoader::new(
1216                                        client.clone(),
1217                                        endpoint,
1218                                        Arc::clone(&token),
1219                                        headers.clone(),
1220                                    ),
1221                                )
1222                            }),
1223                        },
1224                        IcebergStorageProvider::Gcs => {
1225                            Self::gcs_storage_factory(endpoint, &client, &token, &headers)
1226                        }
1227                        IcebergStorageProvider::Adls => OpenDalStorageFactory::Azdls {
1228                            customized_credential_load: endpoint.map(|endpoint| {
1229                                CustomAzdlsCredentialLoader::new(
1230                                    iceberg_credentials::VendedCredentialLoader::new(
1231                                        client.clone(),
1232                                        endpoint,
1233                                        Arc::clone(&token),
1234                                        headers.clone(),
1235                                    ),
1236                                )
1237                            }),
1238                        },
1239                    },
1240                    // NOTE: We construct our own OAuth authenticator for the Catalog client instead of using the one built in.
1241                    // This means we ignore auth overrides from `/v1/config` (e.g. `oauth2-server-uri`).
1242                    // This is okay because users can set these configs from Mz SQL.
1243                    Some(iceberg_catalog_rest::BearerTokenAuthenticator::new(token)),
1244                )
1245            }
1246            IcebergCatalogAuth::Gcp(gcp_connection_reference) => {
1247                let (creds_json, service_account) = gcp_connection_reference
1248                    .connection
1249                    .read_credentials(storage_configuration)
1250                    .await
1251                    .map_err(|e| anyhow!("failed to parse GCP service account JSON: {e}"))?;
1252
1253                props.insert(
1254                    GCS_CREDENTIALS_JSON.to_owned(),
1255                    base64::engine::general_purpose::STANDARD.encode(creds_json),
1256                );
1257                // We supplied a service account key. Don't look elsewhere for GCP credentials.
1258                props.insert(GCS_DISABLE_VM_METADATA.to_owned(), "true".to_owned());
1259                props.insert(GCS_DISABLE_CONFIG_LOAD.to_owned(), "true".to_owned());
1260                if let Some(project_id) = service_account.project_id() {
1261                    props.insert(GCS_USER_PROJECT.to_owned(), project_id.to_owned());
1262                    props.insert(
1263                        "header.x-goog-user-project".to_owned(),
1264                        project_id.to_owned(),
1265                    );
1266                }
1267
1268                // The service account authenticates catalog requests whether or not the catalog
1269                // vends storage credentials, and doubles as the token source for refreshing them.
1270                let token: Arc<dyn TokenProvider> = Arc::new(GcpTokenProvider { service_account });
1271                let headers = Self::catalog_headers(&props)?;
1272                let endpoint = self
1273                    .vended_credential_endpoint(rest, &client, &token, &headers, table)
1274                    .await?;
1275
1276                (
1277                    // A GCP-hosted catalog vends GCS credentials, so the storage provider is not
1278                    // in question here the way it is for a generic REST catalog.
1279                    //
1280                    // NOTE: with delegation the service account key above stops governing storage
1281                    // and only authenticates the catalog. That is not a change in precedence:
1282                    // OpenDAL already preferred the vended token over a credential file, and a
1283                    // loader only keeps that token from expiring.
1284                    Self::gcs_storage_factory(endpoint, &client, &token, &headers),
1285                    Some(iceberg_catalog_rest::BearerTokenAuthenticator::new(token)),
1286                )
1287            }
1288        };
1289
1290        // Inserted after the storage factory is built, so that the loaders above carry only the
1291        // headers the catalog client would send on an ordinary request. Each adds this one itself,
1292        // since the credentials endpoint is the one request that always asks for delegation.
1293        if let Some(delegation) = &rest.access_delegation {
1294            props.insert(
1295                REST_CATALOG_PROP_ACCESS_DELEGATION.to_string(),
1296                delegation.as_header_value().to_string(),
1297            );
1298        }
1299
1300        let mut catalog = RestCatalogBuilder::default()
1301            .with_storage_factory(Arc::new(storage_factory))
1302            .with_client(client);
1303        if let Some(auth) = custom_authenticator {
1304            catalog = catalog.with_authenticator(Arc::new(auth));
1305        }
1306        let catalog = catalog
1307            .load("IcebergCatalog", props.into_iter().collect())
1308            .await
1309            .map_err(|e| anyhow!("failed to create Iceberg catalog: {e}"))?;
1310        Ok(Arc::new(catalog))
1311    }
1312
1313    async fn validate(
1314        &self,
1315        _id: CatalogItemId,
1316        storage_configuration: &StorageConfiguration,
1317    ) -> Result<(), ConnectionValidationError> {
1318        // Validation only lists namespaces, so it needs no table-scoped credentials.
1319        let catalog = self
1320            .connect(storage_configuration, InTask::No, None)
1321            .await
1322            .map_err(|e| {
1323                ConnectionValidationError::Other(anyhow!("failed to connect to catalog: {e}"))
1324            })?;
1325
1326        // If we can list namespaces, the connection is valid.
1327        catalog.list_namespaces(None).await.map_err(|e| {
1328            ConnectionValidationError::Other(anyhow!("failed to list namespaces: {e}"))
1329        })?;
1330
1331        Ok(())
1332    }
1333}
1334
1335#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
1336pub struct AwsPrivatelinkConnection {
1337    pub service_name: String,
1338    pub availability_zones: Vec<String>,
1339}
1340
1341impl AlterCompatible for AwsPrivatelinkConnection {
1342    fn alter_compatible(&self, _id: GlobalId, _other: &Self) -> Result<(), AlterError> {
1343        // Every element of the AwsPrivatelinkConnection connection is configurable.
1344        Ok(())
1345    }
1346}
1347
1348/// A `SERVICE NAME` that cannot name an AWS VPC endpoint service.
1349#[derive(Clone, Debug, Eq, PartialEq)]
1350pub struct InvalidAwsPrivatelinkServiceName {
1351    pub name: String,
1352}
1353
1354impl fmt::Display for InvalidAwsPrivatelinkServiceName {
1355    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1356        write!(
1357            f,
1358            "invalid AWS PrivateLink service name {}",
1359            self.name.quoted()
1360        )
1361    }
1362}
1363
1364impl std::error::Error for InvalidAwsPrivatelinkServiceName {}
1365
1366impl InvalidAwsPrivatelinkServiceName {
1367    /// Explains how to find the right value.
1368    pub fn hint(&self) -> String {
1369        "SERVICE NAME must name an AWS VPC endpoint service, for example \
1370         `com.amazonaws.vpce.us-east-1.vpce-svc-0e123abc123198abc`. Endpoint service names are \
1371         listed in the AWS console under VPC > Endpoint services."
1372            .into()
1373    }
1374}
1375
1376impl AwsPrivatelinkConnection {
1377    /// Checks that `service_name` could name an AWS VPC endpoint service.
1378    ///
1379    /// Every endpoint service name starts with `com.amazonaws.`, whether the
1380    /// service is customer-owned (`com.amazonaws.vpce.<region>.vpce-svc-<id>`)
1381    /// or AWS-managed (`com.amazonaws.<region>.<service>`). Only the prefix is
1382    /// checked, so a name AWS would accept is never rejected here.
1383    pub fn check_service_name(service_name: &str) -> Result<(), InvalidAwsPrivatelinkServiceName> {
1384        if service_name.starts_with("com.amazonaws.") {
1385            return Ok(());
1386        }
1387        Err(InvalidAwsPrivatelinkServiceName {
1388            name: service_name.to_string(),
1389        })
1390    }
1391}
1392
1393#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
1394pub struct KafkaTlsConfig {
1395    pub identity: Option<TlsIdentity>,
1396    pub root_cert: Option<StringOrSecret>,
1397}
1398
1399#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
1400pub struct KafkaSaslConfig<C: ConnectionAccess = InlinedConnection> {
1401    pub mechanism: String,
1402    pub username: StringOrSecret,
1403    pub password: Option<CatalogItemId>,
1404    pub aws: Option<AwsConnectionReference<C>>,
1405}
1406
1407impl<R: ConnectionResolver> IntoInlineConnection<KafkaSaslConfig, R>
1408    for KafkaSaslConfig<ReferencedConnection>
1409{
1410    fn into_inline_connection(self, r: R) -> KafkaSaslConfig {
1411        KafkaSaslConfig {
1412            mechanism: self.mechanism,
1413            username: self.username,
1414            password: self.password,
1415            aws: self.aws.map(|aws| aws.into_inline_connection(&r)),
1416        }
1417    }
1418}
1419
1420/// Specifies a Kafka broker in a [`KafkaConnection`].
1421#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
1422pub struct KafkaBroker<C: ConnectionAccess = InlinedConnection> {
1423    /// The address of the Kafka broker.
1424    pub address: String,
1425    /// An optional tunnel to use when connecting to the broker.
1426    pub tunnel: Tunnel<C>,
1427}
1428
1429impl<R: ConnectionResolver> IntoInlineConnection<KafkaBroker, R>
1430    for KafkaBroker<ReferencedConnection>
1431{
1432    fn into_inline_connection(self, r: R) -> KafkaBroker {
1433        let KafkaBroker { address, tunnel } = self;
1434        KafkaBroker {
1435            address,
1436            tunnel: tunnel.into_inline_connection(r),
1437        }
1438    }
1439}
1440
1441#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize, Default)]
1442pub struct KafkaTopicOptions {
1443    /// The replication factor for the topic.
1444    /// If `None`, the broker default will be used.
1445    pub replication_factor: Option<NonNeg<i32>>,
1446    /// The number of partitions to create.
1447    /// If `None`, the broker default will be used.
1448    pub partition_count: Option<NonNeg<i32>>,
1449    /// The initial configuration parameters for the topic.
1450    pub topic_config: BTreeMap<String, String>,
1451}
1452
1453#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
1454pub struct KafkaConnection<C: ConnectionAccess = InlinedConnection> {
1455    pub brokers: Vec<KafkaBroker<C>>,
1456    /// A tunnel through which to route traffic,
1457    /// that can be overridden for individual brokers
1458    /// in `brokers`.
1459    pub default_tunnel: Tunnel<C>,
1460    pub progress_topic: Option<String>,
1461    pub progress_topic_options: KafkaTopicOptions,
1462    pub options: BTreeMap<String, StringOrSecret>,
1463    pub tls: Option<KafkaTlsConfig>,
1464    pub sasl: Option<KafkaSaslConfig<C>>,
1465}
1466
1467impl<R: ConnectionResolver> IntoInlineConnection<KafkaConnection, R>
1468    for KafkaConnection<ReferencedConnection>
1469{
1470    fn into_inline_connection(self, r: R) -> KafkaConnection {
1471        let KafkaConnection {
1472            brokers,
1473            progress_topic,
1474            progress_topic_options,
1475            default_tunnel,
1476            options,
1477            tls,
1478            sasl,
1479        } = self;
1480
1481        let brokers = brokers
1482            .into_iter()
1483            .map(|broker| broker.into_inline_connection(&r))
1484            .collect();
1485
1486        KafkaConnection {
1487            brokers,
1488            progress_topic,
1489            progress_topic_options,
1490            default_tunnel: default_tunnel.into_inline_connection(&r),
1491            options,
1492            tls,
1493            sasl: sasl.map(|sasl| sasl.into_inline_connection(&r)),
1494        }
1495    }
1496}
1497
1498impl<C: ConnectionAccess> KafkaConnection<C> {
1499    /// Returns the name of the progress topic to use for the connection.
1500    ///
1501    /// The caller is responsible for providing the connection ID as it is not
1502    /// known to `KafkaConnection`.
1503    ///
1504    /// NOTE: the `mz_catalog.mz_kafka_connections` builtin materialized view
1505    /// reconstructs the default (`_materialize-progress-<env>-<conn>`) in SQL
1506    /// (see `MZ_KAFKA_CONNECTIONS` in `src/catalog/src/builtin/mz_catalog.rs`).
1507    /// Keep the two in sync.
1508    pub fn progress_topic(
1509        &self,
1510        connection_context: &ConnectionContext,
1511        connection_id: CatalogItemId,
1512    ) -> Cow<'_, str> {
1513        if let Some(progress_topic) = &self.progress_topic {
1514            Cow::Borrowed(progress_topic)
1515        } else {
1516            Cow::Owned(format!(
1517                "_materialize-progress-{}-{}",
1518                connection_context.environment_id, connection_id,
1519            ))
1520        }
1521    }
1522
1523    fn validate_by_default(&self) -> bool {
1524        true
1525    }
1526}
1527
1528impl KafkaConnection {
1529    /// Generates a string that can be used as the base for a configuration ID
1530    /// (e.g., `client.id`, `group.id`, `transactional.id`) for a Kafka source
1531    /// or sink.
1532    ///
1533    /// NOTE: the `mz_catalog.mz_kafka_sources` builtin materialized view
1534    /// reconstructs this exact `materialize-<env>-<conn>-<obj>` format in SQL
1535    /// (see `MZ_KAFKA_SOURCES` in `src/catalog/src/builtin/mz_catalog.rs`).
1536    /// The two must stay in sync. `test/testdrive/kafka-commit.td` guards this
1537    /// by feeding the view's reconstructed value into `kafka-verify-commit`,
1538    /// so a divergence here fails that test.
1539    pub fn id_base(
1540        connection_context: &ConnectionContext,
1541        connection_id: CatalogItemId,
1542        object_id: GlobalId,
1543    ) -> String {
1544        format!(
1545            "materialize-{}-{}-{}",
1546            connection_context.environment_id, connection_id, object_id,
1547        )
1548    }
1549
1550    /// Enriches the provided `client_id` according to any enrichment rules in
1551    /// the `kafka_client_id_enrichment_rules` configuration parameter.
1552    pub fn enrich_client_id(&self, configs: &ConfigSet, client_id: &mut String) {
1553        #[derive(Debug, Deserialize)]
1554        struct EnrichmentRule {
1555            #[serde(deserialize_with = "deserialize_regex")]
1556            pattern: Regex,
1557            payload: String,
1558        }
1559
1560        fn deserialize_regex<'de, D>(deserializer: D) -> Result<Regex, D::Error>
1561        where
1562            D: Deserializer<'de>,
1563        {
1564            let buf = String::deserialize(deserializer)?;
1565            Regex::new(&buf).map_err(serde::de::Error::custom)
1566        }
1567
1568        let rules = KAFKA_CLIENT_ID_ENRICHMENT_RULES.get(configs);
1569        let rules = match serde_json::from_value::<Vec<EnrichmentRule>>(rules) {
1570            Ok(rules) => rules,
1571            Err(e) => {
1572                warn!(%e, "failed to decode kafka_client_id_enrichment_rules");
1573                return;
1574            }
1575        };
1576
1577        // Check every rule against every broker. Rules are matched in the order
1578        // that they are specified. It is usually a configuration error if
1579        // multiple rules match the same list of Kafka brokers, but we
1580        // nonetheless want to provide well defined semantics.
1581        debug!(?self.brokers, "evaluating client ID enrichment rules");
1582        for rule in rules {
1583            let is_match = self
1584                .brokers
1585                .iter()
1586                .any(|b| rule.pattern.is_match(&b.address));
1587            debug!(?rule, is_match, "evaluated client ID enrichment rule");
1588            if is_match {
1589                client_id.push('-');
1590                client_id.push_str(&rule.payload);
1591            }
1592        }
1593    }
1594
1595    /// Creates a Kafka client for the connection.
1596    pub async fn create_with_context<C, T>(
1597        &self,
1598        storage_configuration: &StorageConfiguration,
1599        context: C,
1600        extra_options: &BTreeMap<&str, String>,
1601        in_task: InTask,
1602    ) -> Result<T, ContextCreationError>
1603    where
1604        C: ClientContext,
1605        T: FromClientConfigAndContext<TunnelingClientContext<C>>,
1606    {
1607        let mut options = self.options.clone();
1608
1609        // Ensure that Kafka topics are *not* automatically created when
1610        // consuming, producing, or fetching metadata for a topic. This ensures
1611        // that we don't accidentally create topics with the wrong number of
1612        // partitions.
1613        options.insert("allow.auto.create.topics".into(), "false".into());
1614
1615        let brokers = match &self.default_tunnel {
1616            Tunnel::AwsPrivatelink(t) => {
1617                assert!(&self.brokers.is_empty());
1618
1619                let algo = KAFKA_DEFAULT_AWS_PRIVATELINK_ENDPOINT_IDENTIFICATION_ALGORITHM
1620                    .get(storage_configuration.config_set());
1621                options.insert("ssl.endpoint.identification.algorithm".into(), algo.into());
1622
1623                // When using a default privatelink tunnel broker/brokers cannot be specified
1624                // instead the tunnel connection_id and port are used for the initial connection.
1625                format!(
1626                    "{}:{}",
1627                    vpc_endpoint_host(
1628                        t.connection_id,
1629                        None, // Default tunnel does not support availability zones.
1630                    ),
1631                    t.port.unwrap_or(9092)
1632                )
1633            }
1634            Tunnel::AwsPrivatelinks(_pl) => {
1635                let algo = KAFKA_DEFAULT_AWS_PRIVATELINK_ENDPOINT_IDENTIFICATION_ALGORITHM
1636                    .get(storage_configuration.config_set());
1637                options.insert("ssl.endpoint.identification.algorithm".into(), algo.into());
1638
1639                if self.brokers.is_empty() {
1640                    return Err(ContextCreationError::Other(anyhow::anyhow!(
1641                        "at least one static broker is required when using BROKER or BROKERS"
1642                    )));
1643                }
1644                self.brokers.iter().map(|b| &b.address).join(",")
1645            }
1646            _ => self.brokers.iter().map(|b| &b.address).join(","),
1647        };
1648        options.insert("bootstrap.servers".into(), brokers.clone().into());
1649        let security_protocol = match (self.tls.is_some(), self.sasl.is_some()) {
1650            (false, false) => "PLAINTEXT",
1651            (true, false) => "SSL",
1652            (false, true) => "SASL_PLAINTEXT",
1653            (true, true) => "SASL_SSL",
1654        };
1655        info!(
1656            "kafka: create_with_context bootstrap.servers={brokers}, security_protocol={security_protocol}"
1657        );
1658        options.insert("security.protocol".into(), security_protocol.into());
1659        if let Some(tls) = &self.tls {
1660            if let Some(root_cert) = &tls.root_cert {
1661                options.insert("ssl.ca.pem".into(), root_cert.clone());
1662            }
1663            if let Some(identity) = &tls.identity {
1664                options.insert("ssl.key.pem".into(), StringOrSecret::Secret(identity.key));
1665                options.insert("ssl.certificate.pem".into(), identity.cert.clone());
1666            }
1667        }
1668        if let Some(sasl) = &self.sasl {
1669            options.insert("sasl.mechanisms".into(), (&sasl.mechanism).into());
1670            options.insert("sasl.username".into(), sasl.username.clone());
1671            if let Some(password) = sasl.password {
1672                options.insert("sasl.password".into(), StringOrSecret::Secret(password));
1673            }
1674        }
1675
1676        options.insert(
1677            "retry.backoff.ms".into(),
1678            KAFKA_RETRY_BACKOFF
1679                .get(storage_configuration.config_set())
1680                .as_millis()
1681                .into(),
1682        );
1683        options.insert(
1684            "retry.backoff.max.ms".into(),
1685            KAFKA_RETRY_BACKOFF_MAX
1686                .get(storage_configuration.config_set())
1687                .as_millis()
1688                .into(),
1689        );
1690        options.insert(
1691            "reconnect.backoff.ms".into(),
1692            KAFKA_RECONNECT_BACKOFF
1693                .get(storage_configuration.config_set())
1694                .as_millis()
1695                .into(),
1696        );
1697        options.insert(
1698            "reconnect.backoff.max.ms".into(),
1699            KAFKA_RECONNECT_BACKOFF_MAX
1700                .get(storage_configuration.config_set())
1701                .as_millis()
1702                .into(),
1703        );
1704
1705        let mut config = mz_kafka_util::client::create_new_client_config(
1706            storage_configuration
1707                .connection_context
1708                .librdkafka_log_level,
1709            storage_configuration.parameters.kafka_timeout_config,
1710        );
1711        for (k, v) in options {
1712            config.set(
1713                k,
1714                v.get_string(
1715                    in_task,
1716                    &storage_configuration.connection_context.secrets_reader,
1717                )
1718                .await
1719                .context("reading kafka secret")?,
1720            );
1721        }
1722        for (k, v) in extra_options {
1723            config.set(*k, v);
1724        }
1725
1726        let aws_config = match self.sasl.as_ref().and_then(|sasl| sasl.aws.as_ref()) {
1727            None => None,
1728            Some(aws) => Some(
1729                aws.connection
1730                    .load_sdk_config(
1731                        &storage_configuration.connection_context,
1732                        aws.connection_id,
1733                        in_task,
1734                        ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
1735                    )
1736                    .await?,
1737            ),
1738        };
1739
1740        // TODO(roshan): Implement enforcement of external address validation once
1741        // rdkafka client has been updated to support providing multiple resolved
1742        // addresses for brokers
1743        let mut context = TunnelingClientContext::new(
1744            context,
1745            Handle::current(),
1746            storage_configuration
1747                .connection_context
1748                .ssh_tunnel_manager
1749                .clone(),
1750            storage_configuration.parameters.ssh_timeout_config,
1751            aws_config,
1752            in_task,
1753        );
1754
1755        match &self.default_tunnel {
1756            Tunnel::Direct => {
1757                // By default, don't offer a default override for broker address lookup.
1758            }
1759            Tunnel::AwsPrivatelink(pl) => {
1760                context.set_default_tunnel(TunnelConfig::StaticHost(
1761                    // Possible bug: We have been ignoring the configured port.
1762                    KafkaConnection::from_default_aws_privatelink(pl).host,
1763                ));
1764            }
1765            Tunnel::AwsPrivatelinks(pl) => {
1766                context.set_default_tunnel(TunnelConfig::Rules(
1767                    KafkaConnection::from_aws_privatelinks(pl),
1768                ));
1769            }
1770            Tunnel::Ssh(ssh_tunnel) => {
1771                let secret = storage_configuration
1772                    .connection_context
1773                    .secrets_reader
1774                    .read_in_task_if(in_task, ssh_tunnel.connection_id)
1775                    .await?;
1776                let key_pair = SshKeyPair::from_bytes(&secret)?;
1777
1778                // Ensure any ssh-bastion address we connect to is resolved to an external address.
1779                let resolved = resolve_address(
1780                    &ssh_tunnel.connection.host,
1781                    ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
1782                )
1783                .await?;
1784                context.set_default_tunnel(TunnelConfig::Ssh(SshTunnelConfig {
1785                    host: resolved
1786                        .iter()
1787                        .map(|a| a.to_string())
1788                        .collect::<BTreeSet<_>>(),
1789                    port: ssh_tunnel.connection.port,
1790                    user: ssh_tunnel.connection.user.clone(),
1791                    key_pair,
1792                }));
1793            }
1794        }
1795        info!(
1796            "kafka: tunnel config set to {}",
1797            match &self.default_tunnel {
1798                Tunnel::Direct => "Direct".to_string(),
1799                Tunnel::AwsPrivatelink(_) => "AwsPrivatelink (static host)".to_string(),
1800                Tunnel::AwsPrivatelinks(pl) =>
1801                    format!("AwsPrivatelinks ({} rules)", pl.rules.len()),
1802                Tunnel::Ssh(_) => "Ssh".to_string(),
1803            }
1804        );
1805
1806        // Here, we preemptively rewrite broker addresses.
1807        // In concept, this overlaps with 'TunnelingClientContext::resolve_broker_addr'.
1808        for broker in &self.brokers {
1809            let mut addr_parts = broker.address.splitn(2, ':');
1810            let addr = BrokerAddr {
1811                host: addr_parts
1812                    .next()
1813                    .context("BROKER is not address:port")?
1814                    .into(),
1815                port: addr_parts
1816                    .next()
1817                    .unwrap_or("9092")
1818                    .parse()
1819                    .context("parsing BROKER port")?,
1820            };
1821            match &broker.tunnel {
1822                Tunnel::Direct => {
1823                    // By default, don't override broker address lookup.
1824                    //
1825                    // N.B.
1826                    //
1827                    // We _could_ pre-setup the default ssh tunnel for all known brokers here, but
1828                    // we avoid doing because:
1829                    // - Its not necessary.
1830                    // - Not doing so makes it easier to test the `FailedDefaultSshTunnel` path
1831                    // in the `TunnelingClientContext`.
1832                }
1833                Tunnel::AwsPrivatelink(aws_privatelink) => {
1834                    context.add_broker_rewrite(
1835                        addr,
1836                        KafkaConnection::from_aws_privatelink(aws_privatelink),
1837                    );
1838                }
1839                Tunnel::AwsPrivatelinks(_) => unreachable!(
1840                    "Individually predefined brokers do not use rule-based PrivateLinks routing."
1841                ),
1842                Tunnel::Ssh(ssh_tunnel) => {
1843                    // Ensure any SSH bastion address we connect to is resolved to an external address.
1844                    let ssh_host_resolved = resolve_address(
1845                        &ssh_tunnel.connection.host,
1846                        ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
1847                    )
1848                    .await?;
1849                    context
1850                        .add_ssh_tunnel(
1851                            addr,
1852                            SshTunnelConfig {
1853                                host: ssh_host_resolved
1854                                    .iter()
1855                                    .map(|a| a.to_string())
1856                                    .collect::<BTreeSet<_>>(),
1857                                port: ssh_tunnel.connection.port,
1858                                user: ssh_tunnel.connection.user.clone(),
1859                                key_pair: SshKeyPair::from_bytes(
1860                                    &storage_configuration
1861                                        .connection_context
1862                                        .secrets_reader
1863                                        .read_in_task_if(in_task, ssh_tunnel.connection_id)
1864                                        .await?,
1865                                )?,
1866                            },
1867                        )
1868                        .await
1869                        .map_err(ContextCreationError::Ssh)?;
1870                }
1871            }
1872        }
1873
1874        Ok(mz_kafka_util::client::create_with_context(
1875            &config, context,
1876        )?)
1877    }
1878
1879    async fn validate(
1880        &self,
1881        _id: CatalogItemId,
1882        storage_configuration: &StorageConfiguration,
1883    ) -> Result<(), anyhow::Error> {
1884        let (context, error_rx) = MzClientContext::with_errors();
1885        let consumer: BaseConsumer<_> = self
1886            .create_with_context(
1887                storage_configuration,
1888                context,
1889                &BTreeMap::new(),
1890                // We are in a normal tokio context during validation, already.
1891                InTask::No,
1892            )
1893            .await?;
1894        let consumer = Arc::new(consumer);
1895
1896        let timeout = storage_configuration
1897            .parameters
1898            .kafka_timeout_config
1899            .fetch_metadata_timeout;
1900
1901        // librdkafka doesn't expose an API for determining whether a connection to
1902        // the Kafka cluster has been successfully established. So we make a
1903        // metadata request, though we don't care about the results, so that we can
1904        // report any errors making that request. If the request succeeds, we know
1905        // we were able to contact at least one broker, and that's a good proxy for
1906        // being able to contact all the brokers in the cluster.
1907        //
1908        // The downside of this approach is it produces a generic error message like
1909        // "metadata fetch error" with no additional details. The real networking
1910        // error is buried in the librdkafka logs, which are not visible to users.
1911        info!("kafka: starting connection validation via fetch_metadata (timeout={timeout:?})");
1912        let result = mz_ore::task::spawn_blocking(|| "kafka_get_metadata", {
1913            let consumer = Arc::clone(&consumer);
1914            move || consumer.fetch_metadata(None, timeout)
1915        })
1916        .await;
1917        info!(
1918            "kafka: connection validation result: {}",
1919            if result.is_ok() { "success" } else { "failed" },
1920        );
1921        match result {
1922            Ok(_) => Ok(()),
1923            // The error returned by `fetch_metadata` does not provide any details which makes for
1924            // a crappy user facing error message. For this reason we attempt to grab a better
1925            // error message from the client context, which should contain any error logs emitted
1926            // by librdkafka, and fallback to the generic error if there is nothing there.
1927            Err(err) => {
1928                // Multiple errors might have been logged during this validation but some are more
1929                // relevant than others. Specifically, we prefer non-internal errors over internal
1930                // errors since those give much more useful information to the users.
1931                let main_err = error_rx.try_iter().reduce(|cur, new| match cur {
1932                    MzKafkaError::Internal(_) => new,
1933                    _ => cur,
1934                });
1935
1936                // Don't drop the consumer until after we've drained the errors
1937                // channel. Dropping the consumer can introduce spurious errors.
1938                // See database-issues#7432.
1939                drop(consumer);
1940
1941                match main_err {
1942                    Some(err) => Err(err.into()),
1943                    None => Err(err.into()),
1944                }
1945            }
1946        }
1947    }
1948
1949    /// The "default" PrivateLink connection is used for bootstrapping Kafka.
1950    fn from_default_aws_privatelink(pl: &AwsPrivatelink) -> BrokerRewrite {
1951        BrokerRewrite {
1952            host: vpc_endpoint_host(
1953                pl.connection_id,
1954                None, // Default tunnel does not support availability zones.
1955            ),
1956            port: pl.port,
1957        }
1958    }
1959
1960    /// The "not default" PrivateLink connections are used for routing to specific Kafka brokers.
1961    fn from_aws_privatelink(pl: &AwsPrivatelink) -> BrokerRewrite {
1962        BrokerRewrite {
1963            host: vpc_endpoint_host(pl.connection_id, pl.availability_zone.as_deref()),
1964            port: pl.port,
1965        }
1966    }
1967
1968    fn from_aws_privatelink_rule(
1969        AwsPrivatelinkRule { pattern, to }: &AwsPrivatelinkRule,
1970    ) -> (mz_kafka_util::client::ConnectionRulePattern, BrokerRewrite) {
1971        (
1972            mz_kafka_util::client::ConnectionRulePattern {
1973                prefix_wildcard: pattern.prefix_wildcard,
1974                literal_match: pattern.literal_match.clone(),
1975                suffix_wildcard: pattern.suffix_wildcard,
1976            },
1977            KafkaConnection::from_aws_privatelink(to),
1978        )
1979    }
1980
1981    fn from_aws_privatelinks(pl: &AwsPrivatelinks) -> HostMappingRules {
1982        HostMappingRules {
1983            rules: pl
1984                .rules
1985                .iter()
1986                .map(KafkaConnection::from_aws_privatelink_rule)
1987                .collect_vec(),
1988        }
1989    }
1990}
1991
1992impl<C: ConnectionAccess> AlterCompatible for KafkaConnection<C> {
1993    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
1994        let KafkaConnection {
1995            brokers: _,
1996            default_tunnel: _,
1997            progress_topic,
1998            progress_topic_options,
1999            options: _,
2000            tls: _,
2001            sasl: _,
2002        } = self;
2003
2004        let compatibility_checks = [
2005            (progress_topic == &other.progress_topic, "progress_topic"),
2006            (
2007                progress_topic_options == &other.progress_topic_options,
2008                "progress_topic_options",
2009            ),
2010        ];
2011
2012        for (compatible, field) in compatibility_checks {
2013            if !compatible {
2014                tracing::warn!(
2015                    "KafkaConnection incompatible at {field}:\nself:\n{:#?}\n\nother\n{:#?}",
2016                    self,
2017                    other
2018                );
2019
2020                return Err(AlterError { id });
2021            }
2022        }
2023
2024        Ok(())
2025    }
2026}
2027
2028/// A connection to a Confluent Schema Registry.
2029#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
2030pub struct CsrConnection<C: ConnectionAccess = InlinedConnection> {
2031    /// The URL of the schema registry.
2032    pub url: Url,
2033    /// Trusted root TLS certificate in PEM format.
2034    pub tls_root_cert: Option<StringOrSecret>,
2035    /// An optional TLS client certificate for authentication with the schema
2036    /// registry.
2037    pub tls_identity: Option<TlsIdentity>,
2038    /// Optional HTTP authentication credentials for the schema registry.
2039    pub http_auth: Option<CsrConnectionHttpAuth>,
2040    /// A tunnel through which to route traffic.
2041    pub tunnel: Tunnel<C>,
2042}
2043
2044impl<R: ConnectionResolver> IntoInlineConnection<CsrConnection, R>
2045    for CsrConnection<ReferencedConnection>
2046{
2047    fn into_inline_connection(self, r: R) -> CsrConnection {
2048        let CsrConnection {
2049            url,
2050            tls_root_cert,
2051            tls_identity,
2052            http_auth,
2053            tunnel,
2054        } = self;
2055        CsrConnection {
2056            url,
2057            tls_root_cert,
2058            tls_identity,
2059            http_auth,
2060            tunnel: tunnel.into_inline_connection(r),
2061        }
2062    }
2063}
2064
2065impl<C: ConnectionAccess> CsrConnection<C> {
2066    fn validate_by_default(&self) -> bool {
2067        true
2068    }
2069}
2070
2071impl CsrConnection {
2072    /// Constructs a schema registry client from the connection.
2073    pub async fn connect(
2074        &self,
2075        storage_configuration: &StorageConfiguration,
2076        in_task: InTask,
2077    ) -> Result<mz_ccsr::Client, CsrConnectError> {
2078        let mut client_config = mz_ccsr::ClientConfig::new(self.url.clone());
2079        if let Some(root_cert) = &self.tls_root_cert {
2080            let root_cert = root_cert
2081                .get_string(
2082                    in_task,
2083                    &storage_configuration.connection_context.secrets_reader,
2084                )
2085                .await?;
2086            let root_cert = Certificate::from_pem(root_cert.as_bytes())?;
2087            client_config = client_config.add_root_certificate(root_cert);
2088        }
2089
2090        if let Some(tls_identity) = &self.tls_identity {
2091            let key = &storage_configuration
2092                .connection_context
2093                .secrets_reader
2094                .read_string_in_task_if(in_task, tls_identity.key)
2095                .await?;
2096            let cert = tls_identity
2097                .cert
2098                .get_string(
2099                    in_task,
2100                    &storage_configuration.connection_context.secrets_reader,
2101                )
2102                .await?;
2103            let ident = Identity::from_pem(key.as_bytes(), cert.as_bytes())?;
2104            client_config = client_config.identity(ident);
2105        }
2106
2107        if let Some(http_auth) = &self.http_auth {
2108            let username = http_auth
2109                .username
2110                .get_string(
2111                    in_task,
2112                    &storage_configuration.connection_context.secrets_reader,
2113                )
2114                .await?;
2115            let password = match http_auth.password {
2116                None => None,
2117                Some(password) => Some(
2118                    storage_configuration
2119                        .connection_context
2120                        .secrets_reader
2121                        .read_string_in_task_if(in_task, password)
2122                        .await?,
2123                ),
2124            };
2125            client_config = client_config.auth(username, password);
2126        }
2127
2128        // TODO: use types to enforce that the URL has a string hostname.
2129        let host = self
2130            .url
2131            .host_str()
2132            .ok_or_else(|| anyhow!("url missing host"))?;
2133        match &self.tunnel {
2134            Tunnel::Direct => {
2135                // Ensure any host we connect to is resolved to an external address.
2136                let resolved = resolve_address(
2137                    host,
2138                    ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
2139                )
2140                .await?;
2141                client_config = client_config.resolve_to_addrs(
2142                    host,
2143                    &resolved
2144                        .iter()
2145                        .map(|addr| SocketAddr::new(*addr, 0))
2146                        .collect::<Vec<_>>(),
2147                )
2148            }
2149            Tunnel::Ssh(ssh_tunnel) => {
2150                let ssh_tunnel = ssh_tunnel
2151                    .connect(
2152                        storage_configuration,
2153                        host,
2154                        // Honor the URL scheme's default port (443 for https,
2155                        // 80 for http) if no explicit port was provided.
2156                        self.url.port_or_known_default().unwrap_or(80),
2157                        in_task,
2158                    )
2159                    .await
2160                    .map_err(CsrConnectError::Ssh)?;
2161
2162                // Carefully inject the SSH tunnel into the client
2163                // configuration. This is delicate because we need TLS
2164                // verification to continue to use the remote hostname rather
2165                // than the tunnel hostname.
2166
2167                client_config = client_config
2168                    // `resolve_to_addrs` allows us to rewrite the hostname
2169                    // at the DNS level, which means the TCP connection is
2170                    // correctly routed through the tunnel, but TLS verification
2171                    // is still performed against the remote hostname.
2172                    // Unfortunately the port here is ignored if the URL also
2173                    // specifies a port...
2174                    .resolve_to_addrs(host, &[SocketAddr::new(ssh_tunnel.local_addr().ip(), 0)])
2175                    // ...so we also dynamically rewrite the URL to use the
2176                    // current port for the SSH tunnel.
2177                    //
2178                    // WARNING: this is brittle, because we only dynamically
2179                    // update the client configuration with the tunnel *port*,
2180                    // and not the hostname This works fine in practice, because
2181                    // only the SSH tunnel port will change if the tunnel fails
2182                    // and has to be restarted (the hostname is always
2183                    // 127.0.0.1)--but this is an an implementation detail of
2184                    // the SSH tunnel code that we're relying on.
2185                    .dynamic_url({
2186                        let remote_url = self.url.clone();
2187                        move || {
2188                            let mut url = remote_url.clone();
2189                            url.set_port(Some(ssh_tunnel.local_addr().port()))
2190                                .expect("cannot fail");
2191                            url
2192                        }
2193                    });
2194            }
2195            Tunnel::AwsPrivatelink(connection) => {
2196                assert_none!(connection.port);
2197
2198                let privatelink_host = mz_cloud_resources::vpc_endpoint_host(
2199                    connection.connection_id,
2200                    connection.availability_zone.as_deref(),
2201                );
2202                let addrs: Vec<_> = net::lookup_host((privatelink_host, 0))
2203                    .await
2204                    .context("resolving PrivateLink host")?
2205                    .collect();
2206                client_config = client_config.resolve_to_addrs(host, &addrs)
2207            }
2208            Tunnel::AwsPrivatelinks(_) => {
2209                unreachable!("MATCHING broker rules are only available for Kafka connections.");
2210            }
2211        }
2212
2213        Ok(client_config.build()?)
2214    }
2215
2216    async fn validate(
2217        &self,
2218        _id: CatalogItemId,
2219        storage_configuration: &StorageConfiguration,
2220    ) -> Result<(), anyhow::Error> {
2221        let client = self
2222            .connect(
2223                storage_configuration,
2224                // We are in a normal tokio context during validation, already.
2225                InTask::No,
2226            )
2227            .await?;
2228        client.list_subjects().await?;
2229        Ok(())
2230    }
2231}
2232
2233impl<C: ConnectionAccess> AlterCompatible for CsrConnection<C> {
2234    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
2235        let CsrConnection {
2236            tunnel,
2237            // All non-tunnel fields may change
2238            url: _,
2239            tls_root_cert: _,
2240            tls_identity: _,
2241            http_auth: _,
2242        } = self;
2243
2244        let compatibility_checks = [(tunnel.alter_compatible(id, &other.tunnel).is_ok(), "tunnel")];
2245
2246        for (compatible, field) in compatibility_checks {
2247            if !compatible {
2248                tracing::warn!(
2249                    "CsrConnection incompatible at {field}:\nself:\n{:#?}\n\nother\n{:#?}",
2250                    self,
2251                    other
2252                );
2253
2254                return Err(AlterError { id });
2255            }
2256        }
2257        Ok(())
2258    }
2259}
2260
2261/// A connection to an AWS Glue Schema Registry.
2262///
2263/// AWS credentials, region, and endpoint are inherited from the referenced
2264/// [`AwsConnection`]; this struct only carries the per-registry settings.
2265#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
2266pub struct GlueSchemaRegistryConnection<C: ConnectionAccess = InlinedConnection> {
2267    /// The referenced AWS connection that supplies credentials, region, and
2268    /// (optional) endpoint.
2269    pub aws_connection: AwsConnectionReference<C>,
2270    /// The Glue Schema Registry name within the AWS account/region.
2271    pub registry_name: String,
2272}
2273
2274impl<R: ConnectionResolver> IntoInlineConnection<GlueSchemaRegistryConnection, R>
2275    for GlueSchemaRegistryConnection<ReferencedConnection>
2276{
2277    fn into_inline_connection(self, r: R) -> GlueSchemaRegistryConnection {
2278        let GlueSchemaRegistryConnection {
2279            aws_connection,
2280            registry_name,
2281        } = self;
2282        GlueSchemaRegistryConnection {
2283            aws_connection: aws_connection.into_inline_connection(&r),
2284            registry_name,
2285        }
2286    }
2287}
2288
2289impl<C: ConnectionAccess> GlueSchemaRegistryConnection<C> {
2290    fn validate_by_default(&self) -> bool {
2291        // Matches CSR: default-validate so a bad registry name fails at
2292        // `CREATE CONNECTION` rather than surfacing later on first use.
2293        // Users can still opt out with `WITH (VALIDATE = false)`.
2294        true
2295    }
2296}
2297
2298impl GlueSchemaRegistryConnection {
2299    async fn validate(
2300        &self,
2301        _id: CatalogItemId,
2302        storage_configuration: &StorageConfiguration,
2303    ) -> Result<(), anyhow::Error> {
2304        let enforce_external_addresses =
2305            crate::dyncfgs::ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set());
2306        let sdk_config = self
2307            .aws_connection
2308            .connection
2309            .load_sdk_config(
2310                &storage_configuration.connection_context,
2311                self.aws_connection.connection_id,
2312                // We are in a normal tokio context during validation.
2313                InTask::No,
2314                enforce_external_addresses,
2315            )
2316            .await?;
2317        let client = mz_aws_glue_schema_registry::ClientConfig::new(sdk_config).build();
2318        match client.get_registry(&self.registry_name).await {
2319            Ok(_) => Ok(()),
2320            Err(mz_aws_glue_schema_registry::GetRegistryError::NotFound) => Err(anyhow!(
2321                "AWS Glue Schema Registry {:?} does not exist in the configured account/region",
2322                self.registry_name
2323            )),
2324            Err(err) => Err(anyhow::Error::new(err).context(format!(
2325                "failed to validate AWS Glue Schema Registry connection (registry={:?})",
2326                self.registry_name
2327            ))),
2328        }
2329    }
2330}
2331
2332impl<C: ConnectionAccess> AlterCompatible for GlueSchemaRegistryConnection<C> {
2333    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
2334        let GlueSchemaRegistryConnection {
2335            registry_name,
2336            // The referenced AWS connection itself may be swapped; matches
2337            // the permissive policy of MySqlConnection / SqlServerConnection.
2338            aws_connection: _,
2339        } = self;
2340
2341        let compatibility_checks = [(registry_name == &other.registry_name, "registry_name")];
2342
2343        for (compatible, field) in compatibility_checks {
2344            if !compatible {
2345                tracing::warn!(
2346                    "GlueSchemaRegistryConnection incompatible at {field}:\nself:\n{:#?}\n\nother\n{:#?}",
2347                    self,
2348                    other
2349                );
2350
2351                return Err(AlterError { id });
2352            }
2353        }
2354        Ok(())
2355    }
2356}
2357
2358/// A TLS key pair used for client identity.
2359#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
2360pub struct TlsIdentity {
2361    /// The client's TLS public certificate in PEM format.
2362    pub cert: StringOrSecret,
2363    /// The ID of the secret containing the client's TLS private key in PEM
2364    /// format.
2365    pub key: CatalogItemId,
2366}
2367
2368/// HTTP authentication credentials in a [`CsrConnection`].
2369#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
2370pub struct CsrConnectionHttpAuth {
2371    /// The username.
2372    pub username: StringOrSecret,
2373    /// The ID of the secret containing the password, if any.
2374    pub password: Option<CatalogItemId>,
2375}
2376
2377/// A connection to a PostgreSQL server.
2378#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
2379pub struct PostgresConnection<C: ConnectionAccess = InlinedConnection> {
2380    /// The hostname of the server.
2381    pub host: String,
2382    /// The port of the server.
2383    pub port: u16,
2384    /// The name of the database to connect to.
2385    pub database: String,
2386    /// The username to authenticate as.
2387    pub user: StringOrSecret,
2388    /// An optional password for authentication.
2389    pub password: Option<CatalogItemId>,
2390    /// A tunnel through which to route traffic.
2391    pub tunnel: Tunnel<C>,
2392    /// Whether to use TLS for encryption, authentication, or both.
2393    pub tls_mode: SslMode,
2394    /// An optional root TLS certificate in PEM format, to verify the server's
2395    /// identity.
2396    pub tls_root_cert: Option<StringOrSecret>,
2397    /// An optional TLS client certificate for authentication.
2398    pub tls_identity: Option<TlsIdentity>,
2399}
2400
2401impl<R: ConnectionResolver> IntoInlineConnection<PostgresConnection, R>
2402    for PostgresConnection<ReferencedConnection>
2403{
2404    fn into_inline_connection(self, r: R) -> PostgresConnection {
2405        let PostgresConnection {
2406            host,
2407            port,
2408            database,
2409            user,
2410            password,
2411            tunnel,
2412            tls_mode,
2413            tls_root_cert,
2414            tls_identity,
2415        } = self;
2416
2417        PostgresConnection {
2418            host,
2419            port,
2420            database,
2421            user,
2422            password,
2423            tunnel: tunnel.into_inline_connection(r),
2424            tls_mode,
2425            tls_root_cert,
2426            tls_identity,
2427        }
2428    }
2429}
2430
2431impl<C: ConnectionAccess> PostgresConnection<C> {
2432    fn validate_by_default(&self) -> bool {
2433        true
2434    }
2435}
2436
2437impl PostgresConnection<InlinedConnection> {
2438    pub async fn config(
2439        &self,
2440        secrets_reader: &Arc<dyn mz_secrets::SecretsReader>,
2441        storage_configuration: &StorageConfiguration,
2442        in_task: InTask,
2443    ) -> Result<mz_postgres_util::Config, anyhow::Error> {
2444        let params = &storage_configuration.parameters;
2445
2446        let mut config = tokio_postgres::Config::new();
2447        config
2448            .host(&self.host)
2449            .port(self.port)
2450            .dbname(&self.database)
2451            .user(&self.user.get_string(in_task, secrets_reader).await?)
2452            .ssl_mode(self.tls_mode)
2453            .application_name("materialize");
2454        if let Some(password) = self.password {
2455            let password = secrets_reader
2456                .read_string_in_task_if(in_task, password)
2457                .await?;
2458            config.password(password);
2459        }
2460        if let Some(tls_root_cert) = &self.tls_root_cert {
2461            let tls_root_cert = tls_root_cert.get_string(in_task, secrets_reader).await?;
2462            config.ssl_root_cert(tls_root_cert.as_bytes());
2463        }
2464        if let Some(tls_identity) = &self.tls_identity {
2465            let cert = tls_identity
2466                .cert
2467                .get_string(in_task, secrets_reader)
2468                .await?;
2469            let key = secrets_reader
2470                .read_string_in_task_if(in_task, tls_identity.key)
2471                .await?;
2472            config.ssl_cert(cert.as_bytes()).ssl_key(key.as_bytes());
2473        }
2474
2475        if let Some(connect_timeout) = params.pg_source_connect_timeout {
2476            config.connect_timeout(connect_timeout);
2477        }
2478        if let Some(keepalives_retries) = params.pg_source_tcp_keepalives_retries {
2479            config.keepalives_retries(keepalives_retries);
2480        }
2481        if let Some(keepalives_idle) = params.pg_source_tcp_keepalives_idle {
2482            config.keepalives_idle(keepalives_idle);
2483        }
2484        if let Some(keepalives_interval) = params.pg_source_tcp_keepalives_interval {
2485            config.keepalives_interval(keepalives_interval);
2486        }
2487        if let Some(tcp_user_timeout) = params.pg_source_tcp_user_timeout {
2488            config.tcp_user_timeout(tcp_user_timeout);
2489        }
2490
2491        let mut options = vec![];
2492        if let Some(wal_sender_timeout) = params.pg_source_wal_sender_timeout {
2493            options.push(format!(
2494                "--wal_sender_timeout={}",
2495                wal_sender_timeout.as_millis()
2496            ));
2497        };
2498        if params.pg_source_tcp_configure_server {
2499            if let Some(keepalives_retries) = params.pg_source_tcp_keepalives_retries {
2500                options.push(format!("--tcp_keepalives_count={}", keepalives_retries));
2501            }
2502            if let Some(keepalives_idle) = params.pg_source_tcp_keepalives_idle {
2503                options.push(format!(
2504                    "--tcp_keepalives_idle={}",
2505                    keepalives_idle.as_secs()
2506                ));
2507            }
2508            if let Some(keepalives_interval) = params.pg_source_tcp_keepalives_interval {
2509                options.push(format!(
2510                    "--tcp_keepalives_interval={}",
2511                    keepalives_interval.as_secs()
2512                ));
2513            }
2514            if let Some(tcp_user_timeout) = params.pg_source_tcp_user_timeout {
2515                options.push(format!(
2516                    "--tcp_user_timeout={}",
2517                    tcp_user_timeout.as_millis()
2518                ));
2519            }
2520        }
2521        config.options(options.join(" ").as_str());
2522
2523        let tunnel = match &self.tunnel {
2524            Tunnel::Direct => {
2525                // Ensure any host we connect to is resolved to an external address.
2526                let resolved = resolve_address(
2527                    &self.host,
2528                    ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
2529                )
2530                .await?;
2531                mz_postgres_util::TunnelConfig::Direct {
2532                    resolved_ips: Some(resolved),
2533                }
2534            }
2535            Tunnel::Ssh(SshTunnel {
2536                connection_id,
2537                connection,
2538            }) => {
2539                let secret = secrets_reader
2540                    .read_in_task_if(in_task, *connection_id)
2541                    .await?;
2542                let key_pair = SshKeyPair::from_bytes(&secret)?;
2543                // Ensure any ssh-bastion host we connect to is resolved to an external address.
2544                let resolved = resolve_address(
2545                    &connection.host,
2546                    ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
2547                )
2548                .await?;
2549                mz_postgres_util::TunnelConfig::Ssh {
2550                    config: SshTunnelConfig {
2551                        host: resolved
2552                            .iter()
2553                            .map(|a| a.to_string())
2554                            .collect::<BTreeSet<_>>(),
2555                        port: connection.port,
2556                        user: connection.user.clone(),
2557                        key_pair,
2558                    },
2559                }
2560            }
2561            Tunnel::AwsPrivatelink(connection) => {
2562                assert_none!(connection.port);
2563                mz_postgres_util::TunnelConfig::AwsPrivatelink {
2564                    connection_id: connection.connection_id,
2565                }
2566            }
2567            Tunnel::AwsPrivatelinks(_) => {
2568                unreachable!("MATCHING broker rules are only available for Kafka connections.");
2569            }
2570        };
2571
2572        Ok(mz_postgres_util::Config::new(
2573            config,
2574            tunnel,
2575            params.ssh_timeout_config,
2576            in_task,
2577        )?)
2578    }
2579
2580    pub async fn validate(
2581        &self,
2582        _id: CatalogItemId,
2583        storage_configuration: &StorageConfiguration,
2584    ) -> Result<mz_postgres_util::Client, anyhow::Error> {
2585        let config = self
2586            .config(
2587                &storage_configuration.connection_context.secrets_reader,
2588                storage_configuration,
2589                // We are in a normal tokio context during validation, already.
2590                InTask::No,
2591            )
2592            .await?;
2593        let client = config
2594            .connect(
2595                "connection validation",
2596                &storage_configuration.connection_context.ssh_tunnel_manager,
2597            )
2598            .await?;
2599
2600        let wal_level = mz_postgres_util::get_wal_level(&client).await?;
2601
2602        if wal_level < mz_postgres_util::replication::WalLevel::Logical {
2603            Err(PostgresConnectionValidationError::InsufficientWalLevel { wal_level })?;
2604        }
2605
2606        let max_wal_senders = mz_postgres_util::get_max_wal_senders(&client).await?;
2607
2608        if max_wal_senders < 1 {
2609            Err(PostgresConnectionValidationError::ReplicationDisabled)?;
2610        }
2611
2612        let available_replication_slots =
2613            mz_postgres_util::available_replication_slots(&client).await?;
2614
2615        // We need 1 replication slot for the snapshots and 1 for the continuing replication
2616        if available_replication_slots < 2 {
2617            Err(
2618                PostgresConnectionValidationError::InsufficientReplicationSlotsAvailable {
2619                    count: 2,
2620                },
2621            )?;
2622        }
2623
2624        Ok(client)
2625    }
2626}
2627
2628#[derive(Debug, Clone, thiserror::Error)]
2629pub enum PostgresConnectionValidationError {
2630    #[error("PostgreSQL server has insufficient number of replication slots available")]
2631    InsufficientReplicationSlotsAvailable { count: usize },
2632    #[error("server must have wal_level >= logical, but has {wal_level}")]
2633    InsufficientWalLevel {
2634        wal_level: mz_postgres_util::replication::WalLevel,
2635    },
2636    #[error("replication disabled on server")]
2637    ReplicationDisabled,
2638}
2639
2640impl PostgresConnectionValidationError {
2641    pub fn detail(&self) -> Option<String> {
2642        match self {
2643            Self::InsufficientReplicationSlotsAvailable { count } => Some(format!(
2644                "executing this statement requires {} replication slot{}",
2645                count,
2646                if *count == 1 { "" } else { "s" }
2647            )),
2648            _ => None,
2649        }
2650    }
2651
2652    pub fn hint(&self) -> Option<String> {
2653        match self {
2654            Self::InsufficientReplicationSlotsAvailable { .. } => Some(
2655                "you might be able to wait for other sources to finish snapshotting and try again"
2656                    .into(),
2657            ),
2658            Self::ReplicationDisabled => Some("set max_wal_senders to a value > 0".into()),
2659            Self::InsufficientWalLevel { .. } => None,
2660        }
2661    }
2662}
2663
2664impl<C: ConnectionAccess> AlterCompatible for PostgresConnection<C> {
2665    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
2666        let PostgresConnection {
2667            tunnel,
2668            // All non-tunnel options may change arbitrarily
2669            host: _,
2670            port: _,
2671            database: _,
2672            user: _,
2673            password: _,
2674            tls_mode: _,
2675            tls_root_cert: _,
2676            tls_identity: _,
2677        } = self;
2678
2679        let compatibility_checks = [(tunnel.alter_compatible(id, &other.tunnel).is_ok(), "tunnel")];
2680
2681        for (compatible, field) in compatibility_checks {
2682            if !compatible {
2683                tracing::warn!(
2684                    "PostgresConnection incompatible at {field}:\nself:\n{:#?}\n\nother\n{:#?}",
2685                    self,
2686                    other
2687                );
2688
2689                return Err(AlterError { id });
2690            }
2691        }
2692        Ok(())
2693    }
2694}
2695
2696/// Specifies how to tunnel a connection.
2697#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
2698pub enum Tunnel<C: ConnectionAccess = InlinedConnection> {
2699    /// No tunneling.
2700    Direct,
2701    /// Via the specified SSH tunnel connection.
2702    Ssh(SshTunnel<C>),
2703    /// Via the specified AWS PrivateLink connection.
2704    AwsPrivatelink(AwsPrivatelink),
2705    AwsPrivatelinks(AwsPrivatelinks),
2706}
2707
2708impl<R: ConnectionResolver> IntoInlineConnection<Tunnel, R> for Tunnel<ReferencedConnection> {
2709    fn into_inline_connection(self, r: R) -> Tunnel {
2710        match self {
2711            Tunnel::Direct => Tunnel::Direct,
2712            Tunnel::Ssh(ssh) => Tunnel::Ssh(ssh.into_inline_connection(r)),
2713            Tunnel::AwsPrivatelink(awspl) => Tunnel::AwsPrivatelink(awspl),
2714            Tunnel::AwsPrivatelinks(x) => Tunnel::AwsPrivatelinks(x),
2715        }
2716    }
2717}
2718
2719impl<C: ConnectionAccess> AlterCompatible for Tunnel<C> {
2720    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
2721        let compatible = match (self, other) {
2722            (Self::Ssh(s), Self::Ssh(o)) => s.alter_compatible(id, o).is_ok(),
2723            (s, o) => s == o,
2724        };
2725
2726        if !compatible {
2727            tracing::warn!(
2728                "Tunnel incompatible:\nself:\n{:#?}\n\nother\n{:#?}",
2729                self,
2730                other
2731            );
2732
2733            return Err(AlterError { id });
2734        }
2735
2736        Ok(())
2737    }
2738}
2739
2740/// Specifies which MySQL SSL Mode to use:
2741/// <https://dev.mysql.com/doc/refman/8.0/en/connection-options.html#option_general_ssl-mode>
2742/// This is not available as an enum in the mysql-async crate, so we define our own.
2743#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
2744pub enum MySqlSslMode {
2745    Disabled,
2746    Required,
2747    VerifyCa,
2748    VerifyIdentity,
2749}
2750
2751/// A connection to a MySQL server.
2752#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
2753pub struct MySqlConnection<C: ConnectionAccess = InlinedConnection> {
2754    /// The hostname of the server.
2755    pub host: String,
2756    /// The port of the server.
2757    pub port: u16,
2758    /// The username to authenticate as.
2759    pub user: StringOrSecret,
2760    /// An optional password for authentication.
2761    pub password: Option<CatalogItemId>,
2762    /// A tunnel through which to route traffic.
2763    pub tunnel: Tunnel<C>,
2764    /// Whether to use TLS for encryption, verify the server's certificate, and identity.
2765    pub tls_mode: MySqlSslMode,
2766    /// An optional root TLS certificate in PEM format, to verify the server's
2767    /// identity.
2768    pub tls_root_cert: Option<StringOrSecret>,
2769    /// An optional TLS client certificate for authentication.
2770    pub tls_identity: Option<TlsIdentity>,
2771    /// Reference to the AWS connection information to be used for IAM authenitcation and
2772    /// assuming AWS roles.
2773    pub aws_connection: Option<AwsConnectionReference<C>>,
2774}
2775
2776impl<R: ConnectionResolver> IntoInlineConnection<MySqlConnection, R>
2777    for MySqlConnection<ReferencedConnection>
2778{
2779    fn into_inline_connection(self, r: R) -> MySqlConnection {
2780        let MySqlConnection {
2781            host,
2782            port,
2783            user,
2784            password,
2785            tunnel,
2786            tls_mode,
2787            tls_root_cert,
2788            tls_identity,
2789            aws_connection,
2790        } = self;
2791
2792        MySqlConnection {
2793            host,
2794            port,
2795            user,
2796            password,
2797            tunnel: tunnel.into_inline_connection(&r),
2798            tls_mode,
2799            tls_root_cert,
2800            tls_identity,
2801            aws_connection: aws_connection.map(|aws| aws.into_inline_connection(&r)),
2802        }
2803    }
2804}
2805
2806impl<C: ConnectionAccess> MySqlConnection<C> {
2807    fn validate_by_default(&self) -> bool {
2808        true
2809    }
2810}
2811
2812impl MySqlConnection<InlinedConnection> {
2813    pub async fn config(
2814        &self,
2815        secrets_reader: &Arc<dyn mz_secrets::SecretsReader>,
2816        storage_configuration: &StorageConfiguration,
2817        in_task: InTask,
2818    ) -> Result<mz_mysql_util::Config, anyhow::Error> {
2819        // TODO(roshan): Set appropriate connection timeouts
2820        let mut opts = mysql_async::OptsBuilder::default()
2821            .ip_or_hostname(&self.host)
2822            .tcp_port(self.port)
2823            .user(Some(&self.user.get_string(in_task, secrets_reader).await?));
2824
2825        if let Some(password) = self.password {
2826            let password = secrets_reader
2827                .read_string_in_task_if(in_task, password)
2828                .await?;
2829            opts = opts.pass(Some(password));
2830        }
2831
2832        // Our `MySqlSslMode` enum matches the official MySQL Client `--ssl-mode` parameter values
2833        // which uses opt-in security features (SSL, CA verification, & Identity verification).
2834        // The mysql_async crate `SslOpts` struct uses an opt-out mechanism for each of these, so
2835        // we need to appropriately disable features to match the intent of each enum value.
2836        let mut ssl_opts = match self.tls_mode {
2837            MySqlSslMode::Disabled => None,
2838            MySqlSslMode::Required => Some(
2839                mysql_async::SslOpts::default()
2840                    .with_danger_accept_invalid_certs(true)
2841                    .with_danger_skip_domain_validation(true),
2842            ),
2843            MySqlSslMode::VerifyCa => {
2844                Some(mysql_async::SslOpts::default().with_danger_skip_domain_validation(true))
2845            }
2846            MySqlSslMode::VerifyIdentity => Some(mysql_async::SslOpts::default()),
2847        };
2848
2849        if matches!(
2850            self.tls_mode,
2851            MySqlSslMode::VerifyCa | MySqlSslMode::VerifyIdentity
2852        ) {
2853            if let Some(tls_root_cert) = &self.tls_root_cert {
2854                let tls_root_cert = tls_root_cert.get_string(in_task, secrets_reader).await?;
2855                ssl_opts = ssl_opts.map(|opts| {
2856                    opts.with_root_certs(vec![tls_root_cert.as_bytes().to_vec().into()])
2857                });
2858            }
2859        }
2860
2861        if let Some(identity) = &self.tls_identity {
2862            let key = secrets_reader
2863                .read_string_in_task_if(in_task, identity.key)
2864                .await?;
2865            let cert = identity.cert.get_string(in_task, secrets_reader).await?;
2866            let (der, pass) =
2867                mz_tls_util::pkcs12der_from_pem(key.as_bytes(), cert.as_bytes())?.into_parts();
2868
2869            // Add client identity to SSLOpts
2870            ssl_opts = ssl_opts.map(|opts| {
2871                opts.with_client_identity(Some(
2872                    mysql_async::ClientIdentity::new(der.into()).with_password(pass),
2873                ))
2874            });
2875        }
2876
2877        opts = opts.ssl_opts(ssl_opts);
2878
2879        let tunnel = match &self.tunnel {
2880            Tunnel::Direct => {
2881                // Ensure any host we connect to is resolved to an external address.
2882                let resolved = resolve_address(
2883                    &self.host,
2884                    ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
2885                )
2886                .await?;
2887                mz_mysql_util::TunnelConfig::Direct {
2888                    resolved_ips: Some(resolved),
2889                }
2890            }
2891            Tunnel::Ssh(SshTunnel {
2892                connection_id,
2893                connection,
2894            }) => {
2895                let secret = secrets_reader
2896                    .read_in_task_if(in_task, *connection_id)
2897                    .await?;
2898                let key_pair = SshKeyPair::from_bytes(&secret)?;
2899                // Ensure any ssh-bastion host we connect to is resolved to an external address.
2900                let resolved = resolve_address(
2901                    &connection.host,
2902                    ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
2903                )
2904                .await?;
2905                mz_mysql_util::TunnelConfig::Ssh {
2906                    config: SshTunnelConfig {
2907                        host: resolved
2908                            .iter()
2909                            .map(|a| a.to_string())
2910                            .collect::<BTreeSet<_>>(),
2911                        port: connection.port,
2912                        user: connection.user.clone(),
2913                        key_pair,
2914                    },
2915                }
2916            }
2917            Tunnel::AwsPrivatelink(connection) => {
2918                assert_none!(connection.port);
2919                mz_mysql_util::TunnelConfig::AwsPrivatelink {
2920                    connection_id: connection.connection_id,
2921                }
2922            }
2923            Tunnel::AwsPrivatelinks(_) => {
2924                unreachable!("MATCHING broker rules are only available for Kafka connections.");
2925            }
2926        };
2927
2928        let aws_config = match self.aws_connection.as_ref() {
2929            None => None,
2930            Some(aws_ref) => Some(
2931                aws_ref
2932                    .connection
2933                    .load_sdk_config(
2934                        &storage_configuration.connection_context,
2935                        aws_ref.connection_id,
2936                        in_task,
2937                        ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
2938                    )
2939                    .await?,
2940            ),
2941        };
2942
2943        Ok(mz_mysql_util::Config::new(
2944            opts,
2945            tunnel,
2946            storage_configuration.parameters.ssh_timeout_config,
2947            in_task,
2948            storage_configuration
2949                .parameters
2950                .mysql_source_timeouts
2951                .clone(),
2952            aws_config,
2953        )?)
2954    }
2955
2956    pub async fn validate(
2957        &self,
2958        _id: CatalogItemId,
2959        storage_configuration: &StorageConfiguration,
2960    ) -> Result<MySqlConn, MySqlConnectionValidationError> {
2961        let config = self
2962            .config(
2963                &storage_configuration.connection_context.secrets_reader,
2964                storage_configuration,
2965                // We are in a normal tokio context during validation, already.
2966                InTask::No,
2967            )
2968            .await?;
2969        let mut conn = config
2970            .connect(
2971                "connection validation",
2972                &storage_configuration.connection_context.ssh_tunnel_manager,
2973            )
2974            .await?;
2975
2976        // Check if the MySQL database is configured to allow row-based consistent GTID replication
2977        let mut setting_errors = vec![];
2978        let gtid_res = mz_mysql_util::ensure_gtid_consistency(&mut conn).await;
2979        let binlog_res = mz_mysql_util::ensure_full_row_binlog_format(&mut conn).await;
2980        let order_res = mz_mysql_util::ensure_replication_commit_order(&mut conn).await;
2981        for res in [gtid_res, binlog_res, order_res] {
2982            match res {
2983                Err(MySqlError::InvalidSystemSetting {
2984                    setting,
2985                    expected,
2986                    actual,
2987                }) => {
2988                    setting_errors.push((setting, expected, actual));
2989                }
2990                Err(err) => Err(err)?,
2991                Ok(()) => {}
2992            }
2993        }
2994        if !setting_errors.is_empty() {
2995            Err(MySqlConnectionValidationError::ReplicationSettingsError(
2996                setting_errors,
2997            ))?;
2998        }
2999
3000        Ok(conn)
3001    }
3002}
3003
3004#[derive(Debug, thiserror::Error)]
3005pub enum MySqlConnectionValidationError {
3006    #[error("Invalid MySQL system replication settings")]
3007    ReplicationSettingsError(Vec<(String, String, String)>),
3008    #[error(transparent)]
3009    Client(#[from] MySqlError),
3010    #[error("{}", .0.display_with_causes())]
3011    Other(#[from] anyhow::Error),
3012}
3013
3014impl MySqlConnectionValidationError {
3015    pub fn detail(&self) -> Option<String> {
3016        match self {
3017            Self::ReplicationSettingsError(settings) => Some(format!(
3018                "Invalid MySQL system replication settings: {}",
3019                itertools::join(
3020                    settings.iter().map(|(setting, expected, actual)| format!(
3021                        "{}: expected {}, got {}",
3022                        setting, expected, actual
3023                    )),
3024                    "; "
3025                )
3026            )),
3027            _ => None,
3028        }
3029    }
3030
3031    pub fn hint(&self) -> Option<String> {
3032        match self {
3033            Self::ReplicationSettingsError(_) => {
3034                Some("Set the necessary MySQL database system settings.".into())
3035            }
3036            _ => None,
3037        }
3038    }
3039}
3040
3041impl<C: ConnectionAccess> AlterCompatible for MySqlConnection<C> {
3042    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
3043        let MySqlConnection {
3044            tunnel,
3045            // All non-tunnel options may change arbitrarily
3046            host: _,
3047            port: _,
3048            user: _,
3049            password: _,
3050            tls_mode: _,
3051            tls_root_cert: _,
3052            tls_identity: _,
3053            aws_connection: _,
3054        } = self;
3055
3056        let compatibility_checks = [(tunnel.alter_compatible(id, &other.tunnel).is_ok(), "tunnel")];
3057
3058        for (compatible, field) in compatibility_checks {
3059            if !compatible {
3060                tracing::warn!(
3061                    "MySqlConnection incompatible at {field}:\nself:\n{:#?}\n\nother\n{:#?}",
3062                    self,
3063                    other
3064                );
3065
3066                return Err(AlterError { id });
3067            }
3068        }
3069        Ok(())
3070    }
3071}
3072
3073/// Details how to connect to an instance of Microsoft SQL Server.
3074///
3075/// For specifics of connecting to SQL Server for purposes of creating a
3076/// Materialize Source, see [`SqlServerSourceConnection`] which wraps this type.
3077///
3078/// [`SqlServerSourceConnection`]: crate::sources::SqlServerSourceConnection
3079#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
3080pub struct SqlServerConnectionDetails<C: ConnectionAccess = InlinedConnection> {
3081    /// The hostname of the server.
3082    pub host: String,
3083    /// The port of the server.
3084    pub port: u16,
3085    /// Database we should connect to.
3086    pub database: String,
3087    /// The username to authenticate as.
3088    pub user: StringOrSecret,
3089    /// Password used for authentication.
3090    pub password: CatalogItemId,
3091    /// A tunnel through which to route traffic.
3092    pub tunnel: Tunnel<C>,
3093    /// Level of encryption to use for the connection.
3094    pub encryption: mz_sql_server_util::config::EncryptionLevel,
3095    /// Certificate validation policy
3096    pub certificate_validation_policy: mz_sql_server_util::config::CertificateValidationPolicy,
3097    /// TLS CA Certifiecate in PEM format
3098    pub tls_root_cert: Option<StringOrSecret>,
3099}
3100
3101impl<C: ConnectionAccess> SqlServerConnectionDetails<C> {
3102    fn validate_by_default(&self) -> bool {
3103        true
3104    }
3105}
3106
3107impl SqlServerConnectionDetails<InlinedConnection> {
3108    /// Attempts to open a connection to the upstream SQL Server instance.
3109    pub async fn validate(
3110        &self,
3111        _id: CatalogItemId,
3112        storage_configuration: &StorageConfiguration,
3113    ) -> Result<mz_sql_server_util::Client, anyhow::Error> {
3114        let config = self
3115            .resolve_config(
3116                &storage_configuration.connection_context.secrets_reader,
3117                storage_configuration,
3118                InTask::No,
3119            )
3120            .await?;
3121        tracing::debug!(?config, "Validating SQL Server connection");
3122
3123        let mut client = mz_sql_server_util::Client::connect(config).await?;
3124
3125        // Ensure the upstream SQL Server instance is configured to allow CDC.
3126        //
3127        // Run all of the checks necessary and collect the errors to provide the best
3128        // guidance as to which system settings need to be enabled.
3129        let mut replication_errors = vec![];
3130        for error in [
3131            mz_sql_server_util::inspect::ensure_database_cdc_enabled(&mut client).await,
3132            mz_sql_server_util::inspect::ensure_snapshot_isolation_enabled(&mut client).await,
3133            mz_sql_server_util::inspect::ensure_sql_server_agent_running(&mut client).await,
3134        ] {
3135            match error {
3136                Err(mz_sql_server_util::SqlServerError::InvalidSystemSetting {
3137                    name,
3138                    expected,
3139                    actual,
3140                }) => replication_errors.push((name, expected, actual)),
3141                Err(other) => Err(other)?,
3142                Ok(()) => (),
3143            }
3144        }
3145        if !replication_errors.is_empty() {
3146            Err(SqlServerConnectionValidationError::ReplicationSettingsError(replication_errors))?;
3147        }
3148
3149        Ok(client)
3150    }
3151
3152    /// Resolve all of the connection details (e.g. read from the [`SecretsReader`])
3153    /// so the returned [`Config`] can be used to open a connection with the
3154    /// upstream system.
3155    ///
3156    /// The provided [`InTask`] argument determines whether any I/O is run in an
3157    /// [`mz_ore::task`] (i.e. a different thread) or directly in the returned
3158    /// future. The main goal here is to prevent running I/O in timely threads.
3159    ///
3160    /// [`Config`]: mz_sql_server_util::Config
3161    pub async fn resolve_config(
3162        &self,
3163        secrets_reader: &Arc<dyn mz_secrets::SecretsReader>,
3164        storage_configuration: &StorageConfiguration,
3165        in_task: InTask,
3166    ) -> Result<mz_sql_server_util::Config, anyhow::Error> {
3167        let dyncfg = storage_configuration.config_set();
3168        let mut inner_config = tiberius::Config::new();
3169
3170        // Setup default connection params.
3171        inner_config.host(&self.host);
3172        inner_config.port(self.port);
3173        inner_config.database(self.database.clone());
3174        inner_config.encryption(self.encryption.into());
3175        match self.certificate_validation_policy {
3176            mz_sql_server_util::config::CertificateValidationPolicy::TrustAll => {
3177                inner_config.trust_cert()
3178            }
3179            mz_sql_server_util::config::CertificateValidationPolicy::VerifyCA => {
3180                inner_config.trust_cert_ca_pem(
3181                    self.tls_root_cert
3182                        .as_ref()
3183                        .unwrap()
3184                        .get_string(in_task, secrets_reader)
3185                        .await
3186                        .context("ca certificate")?,
3187                );
3188            }
3189            mz_sql_server_util::config::CertificateValidationPolicy::VerifySystem => (), // no-op
3190        }
3191
3192        inner_config.application_name("materialize");
3193
3194        // Read our auth settings from
3195        let user = self
3196            .user
3197            .get_string(in_task, secrets_reader)
3198            .await
3199            .context("username")?;
3200        let password = secrets_reader
3201            .read_string_in_task_if(in_task, self.password)
3202            .await
3203            .context("password")?;
3204        // TODO(sql_server3): Support other methods of authentication besides
3205        // username and password.
3206        inner_config.authentication(tiberius::AuthMethod::sql_server(user, password));
3207
3208        // Prevent users from probing our internal network ports by trying to
3209        // connect to localhost, or another non-external IP.
3210        let enforce_external_addresses = ENFORCE_EXTERNAL_ADDRESSES.get(dyncfg);
3211
3212        let tunnel = match &self.tunnel {
3213            Tunnel::Direct => {
3214                let resolved_addresses: Vec<SocketAddr> =
3215                    resolve_address(&self.host, enforce_external_addresses)
3216                        .await?
3217                        .into_iter()
3218                        .map(|ip| SocketAddr::new(ip, self.port))
3219                        .collect();
3220                mz_sql_server_util::config::TunnelConfig::Direct {
3221                    resolved_addresses: resolved_addresses.into_boxed_slice(),
3222                }
3223            }
3224            Tunnel::Ssh(SshTunnel {
3225                connection_id,
3226                connection: ssh_connection,
3227            }) => {
3228                let secret = secrets_reader
3229                    .read_in_task_if(in_task, *connection_id)
3230                    .await
3231                    .context("ssh secret")?;
3232                let key_pair = SshKeyPair::from_bytes(&secret).context("ssh key pair")?;
3233                // Ensure any SSH-bastion host we connect to is resolved to an
3234                // external address.
3235                let addresses = resolve_address(&ssh_connection.host, enforce_external_addresses)
3236                    .await
3237                    .context("ssh tunnel")?;
3238
3239                let config = SshTunnelConfig {
3240                    host: addresses.into_iter().map(|a| a.to_string()).collect(),
3241                    port: ssh_connection.port,
3242                    user: ssh_connection.user.clone(),
3243                    key_pair,
3244                };
3245                mz_sql_server_util::config::TunnelConfig::Ssh {
3246                    config,
3247                    manager: storage_configuration
3248                        .connection_context
3249                        .ssh_tunnel_manager
3250                        .clone(),
3251                    timeout: storage_configuration.parameters.ssh_timeout_config.clone(),
3252                    host: self.host.clone(),
3253                    port: self.port,
3254                }
3255            }
3256            Tunnel::AwsPrivatelink(private_link_connection) => {
3257                assert_none!(private_link_connection.port);
3258                mz_sql_server_util::config::TunnelConfig::AwsPrivatelink {
3259                    connection_id: private_link_connection.connection_id,
3260                    port: self.port,
3261                }
3262            }
3263            Tunnel::AwsPrivatelinks(_) => {
3264                unreachable!("MATCHING broker rules are only available for Kafka connections.");
3265            }
3266        };
3267
3268        Ok(mz_sql_server_util::Config::new(
3269            inner_config,
3270            tunnel,
3271            in_task,
3272        ))
3273    }
3274}
3275
3276#[derive(Debug, Clone, thiserror::Error)]
3277pub enum SqlServerConnectionValidationError {
3278    #[error("Invalid SQL Server system replication settings")]
3279    ReplicationSettingsError(Vec<(String, String, String)>),
3280}
3281
3282impl SqlServerConnectionValidationError {
3283    pub fn detail(&self) -> Option<String> {
3284        match self {
3285            Self::ReplicationSettingsError(settings) => Some(format!(
3286                "Invalid SQL Server system replication settings: {}",
3287                itertools::join(
3288                    settings.iter().map(|(setting, expected, actual)| format!(
3289                        "{}: expected {}, got {}",
3290                        setting, expected, actual
3291                    )),
3292                    "; "
3293                )
3294            )),
3295        }
3296    }
3297
3298    pub fn hint(&self) -> Option<String> {
3299        match self {
3300            _ => None,
3301        }
3302    }
3303}
3304
3305impl<R: ConnectionResolver> IntoInlineConnection<SqlServerConnectionDetails, R>
3306    for SqlServerConnectionDetails<ReferencedConnection>
3307{
3308    fn into_inline_connection(self, r: R) -> SqlServerConnectionDetails {
3309        let SqlServerConnectionDetails {
3310            host,
3311            port,
3312            database,
3313            user,
3314            password,
3315            tunnel,
3316            encryption,
3317            certificate_validation_policy,
3318            tls_root_cert,
3319        } = self;
3320
3321        SqlServerConnectionDetails {
3322            host,
3323            port,
3324            database,
3325            user,
3326            password,
3327            tunnel: tunnel.into_inline_connection(&r),
3328            encryption,
3329            certificate_validation_policy,
3330            tls_root_cert,
3331        }
3332    }
3333}
3334
3335impl<C: ConnectionAccess> AlterCompatible for SqlServerConnectionDetails<C> {
3336    fn alter_compatible(
3337        &self,
3338        id: mz_repr::GlobalId,
3339        other: &Self,
3340    ) -> Result<(), crate::controller::AlterError> {
3341        let SqlServerConnectionDetails {
3342            tunnel,
3343            // TODO(sql_server2): Figure out how these variables are allowed to change.
3344            host: _,
3345            port: _,
3346            database: _,
3347            user: _,
3348            password: _,
3349            encryption: _,
3350            certificate_validation_policy: _,
3351            tls_root_cert: _,
3352        } = self;
3353
3354        let compatibility_checks = [(tunnel.alter_compatible(id, &other.tunnel).is_ok(), "tunnel")];
3355
3356        for (compatible, field) in compatibility_checks {
3357            if !compatible {
3358                tracing::warn!(
3359                    "SqlServerConnectionDetails incompatible at {field}:\nself:\n{:#?}\n\nother\n{:#?}",
3360                    self,
3361                    other
3362                );
3363
3364                return Err(AlterError { id });
3365            }
3366        }
3367        Ok(())
3368    }
3369}
3370
3371/// A connection to an SSH tunnel.
3372#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
3373pub struct SshConnection {
3374    pub host: String,
3375    pub port: u16,
3376    pub user: String,
3377}
3378
3379use self::inline::{
3380    ConnectionAccess, ConnectionResolver, InlinedConnection, IntoInlineConnection,
3381    ReferencedConnection,
3382};
3383
3384impl AlterCompatible for SshConnection {
3385    fn alter_compatible(&self, _id: GlobalId, _other: &Self) -> Result<(), AlterError> {
3386        // Every element of the SSH connection is configurable.
3387        Ok(())
3388    }
3389}
3390
3391/// Specifies an AWS PrivateLink service for a [`Tunnel`].
3392#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
3393pub struct AwsPrivatelink {
3394    /// The ID of the connection to the AWS PrivateLink service.
3395    pub connection_id: CatalogItemId,
3396    // The availability zone to use when connecting to the AWS PrivateLink service.
3397    pub availability_zone: Option<String>,
3398    /// The port to use when connecting to the AWS PrivateLink service, if
3399    /// different from the port in [`KafkaBroker::address`].
3400    pub port: Option<u16>,
3401}
3402
3403impl AlterCompatible for AwsPrivatelink {
3404    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
3405        let AwsPrivatelink {
3406            connection_id,
3407            availability_zone: _,
3408            port: _,
3409        } = self;
3410
3411        let compatibility_checks = [(connection_id == &other.connection_id, "connection_id")];
3412
3413        for (compatible, field) in compatibility_checks {
3414            if !compatible {
3415                tracing::warn!(
3416                    "AwsPrivatelink incompatible at {field}:\nself:\n{:#?}\n\nother\n{:#?}",
3417                    self,
3418                    other
3419                );
3420
3421                return Err(AlterError { id });
3422            }
3423        }
3424
3425        Ok(())
3426    }
3427}
3428
3429#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
3430pub struct AwsPrivatelinks {
3431    /// Route to brokers through PrivateLink connections according to these rules.
3432    /// Exact-match rules (no wildcards) are used as bootstrap brokers.
3433    /// Wildcard rules are applied dynamically to discovered brokers.
3434    pub rules: Vec<AwsPrivatelinkRule>,
3435}
3436
3437#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
3438pub struct AwsPrivatelinkRule {
3439    /// Given a broker's host:port, should we use this route?
3440    pub pattern: ConnectionRulePattern,
3441    /// Route to the broker through this PrivateLink connection.
3442    pub to: AwsPrivatelink,
3443}
3444
3445/// Specifies an SSH tunnel connection.
3446#[derive(Clone, Debug, Eq, PartialEq, Hash, Serialize, Deserialize)]
3447pub struct SshTunnel<C: ConnectionAccess = InlinedConnection> {
3448    /// id of the ssh connection
3449    pub connection_id: CatalogItemId,
3450    /// ssh connection object
3451    pub connection: C::Ssh,
3452}
3453
3454impl<R: ConnectionResolver> IntoInlineConnection<SshTunnel, R> for SshTunnel<ReferencedConnection> {
3455    fn into_inline_connection(self, r: R) -> SshTunnel {
3456        let SshTunnel {
3457            connection,
3458            connection_id,
3459        } = self;
3460
3461        SshTunnel {
3462            connection: r.resolve_connection(connection).unwrap_ssh(),
3463            connection_id,
3464        }
3465    }
3466}
3467
3468impl SshTunnel<InlinedConnection> {
3469    /// Like [`SshTunnelConfig::connect`], but the SSH key is loaded from a
3470    /// secret.
3471    async fn connect(
3472        &self,
3473        storage_configuration: &StorageConfiguration,
3474        remote_host: &str,
3475        remote_port: u16,
3476        in_task: InTask,
3477    ) -> Result<ManagedSshTunnelHandle, anyhow::Error> {
3478        // Ensure any ssh-bastion host we connect to is resolved to an external address.
3479        let resolved = resolve_address(
3480            &self.connection.host,
3481            ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
3482        )
3483        .await?;
3484        storage_configuration
3485            .connection_context
3486            .ssh_tunnel_manager
3487            .connect(
3488                SshTunnelConfig {
3489                    host: resolved
3490                        .iter()
3491                        .map(|a| a.to_string())
3492                        .collect::<BTreeSet<_>>(),
3493                    port: self.connection.port,
3494                    user: self.connection.user.clone(),
3495                    key_pair: SshKeyPair::from_bytes(
3496                        &storage_configuration
3497                            .connection_context
3498                            .secrets_reader
3499                            .read_in_task_if(in_task, self.connection_id)
3500                            .await?,
3501                    )?,
3502                },
3503                remote_host,
3504                remote_port,
3505                storage_configuration.parameters.ssh_timeout_config,
3506                in_task,
3507            )
3508            .await
3509    }
3510}
3511
3512impl<C: ConnectionAccess> AlterCompatible for SshTunnel<C> {
3513    fn alter_compatible(&self, id: GlobalId, other: &Self) -> Result<(), AlterError> {
3514        let SshTunnel {
3515            connection_id,
3516            connection,
3517        } = self;
3518
3519        let compatibility_checks = [
3520            (connection_id == &other.connection_id, "connection_id"),
3521            (
3522                connection.alter_compatible(id, &other.connection).is_ok(),
3523                "connection",
3524            ),
3525        ];
3526
3527        for (compatible, field) in compatibility_checks {
3528            if !compatible {
3529                tracing::warn!(
3530                    "SshTunnel incompatible at {field}:\nself:\n{:#?}\n\nother\n{:#?}",
3531                    self,
3532                    other
3533                );
3534
3535                return Err(AlterError { id });
3536            }
3537        }
3538
3539        Ok(())
3540    }
3541}
3542
3543impl SshConnection {
3544    #[allow(clippy::unused_async)]
3545    async fn validate(
3546        &self,
3547        id: CatalogItemId,
3548        storage_configuration: &StorageConfiguration,
3549    ) -> Result<(), anyhow::Error> {
3550        let secret = storage_configuration
3551            .connection_context
3552            .secrets_reader
3553            .read_in_task_if(
3554                // We are in a normal tokio context during validation, already.
3555                InTask::No,
3556                id,
3557            )
3558            .await?;
3559        let key_pair = SshKeyPair::from_bytes(&secret)?;
3560
3561        // Ensure any ssh-bastion host we connect to is resolved to an external address.
3562        let resolved = resolve_address(
3563            &self.host,
3564            ENFORCE_EXTERNAL_ADDRESSES.get(storage_configuration.config_set()),
3565        )
3566        .await?;
3567
3568        let config = SshTunnelConfig {
3569            host: resolved
3570                .iter()
3571                .map(|a| a.to_string())
3572                .collect::<BTreeSet<_>>(),
3573            port: self.port,
3574            user: self.user.clone(),
3575            key_pair,
3576        };
3577        // Note that we do NOT use the `SshTunnelManager` here, as we want to validate that we
3578        // can actually create a new connection to the ssh bastion, without tunneling.
3579        config
3580            .validate(storage_configuration.parameters.ssh_timeout_config)
3581            .await
3582    }
3583
3584    fn validate_by_default(&self) -> bool {
3585        false
3586    }
3587}
3588
3589impl AwsPrivatelinkConnection {
3590    #[allow(clippy::unused_async)]
3591    async fn validate(
3592        &self,
3593        id: CatalogItemId,
3594        storage_configuration: &StorageConfiguration,
3595    ) -> Result<(), ConnectionValidationError> {
3596        // An endpoint for an unusable service name reports a misleading
3597        // condition (missing availability zones), so check the name first.
3598        Self::check_service_name(&self.service_name)?;
3599
3600        let Some(ref cloud_resource_reader) = storage_configuration
3601            .connection_context
3602            .cloud_resource_reader
3603        else {
3604            return Err(anyhow!("AWS PrivateLink connections are unsupported").into());
3605        };
3606
3607        // No need to optionally run this in a task, as we are just validating from envd.
3608        let status = cloud_resource_reader.read(id).await?;
3609
3610        let availability = status
3611            .conditions
3612            .as_ref()
3613            .and_then(|conditions| conditions.iter().find(|c| c.type_ == "Available"));
3614
3615        match availability {
3616            Some(condition) if condition.status == "True" => Ok(()),
3617            Some(condition) => Err(anyhow!("{}", condition.message).into()),
3618            None => Err(anyhow!("Endpoint availability is unknown").into()),
3619        }
3620    }
3621
3622    fn validate_by_default(&self) -> bool {
3623        false
3624    }
3625}
3626
3627#[cfg(test)]
3628mod tests {
3629    use super::*;
3630
3631    #[mz_ore::test]
3632    fn test_catalog_headers() {
3633        let props = BTreeMap::from_iter(
3634            [
3635                (REST_CATALOG_PROP_URI, "https://catalog.example"),
3636                (REST_CATALOG_PROP_WAREHOUSE, "wh"),
3637                ("header.x-goog-user-project", "some-project"),
3638                (REST_CATALOG_PROP_ACCESS_DELEGATION, "vended-credentials"),
3639            ]
3640            .map(|(k, v)| (k.to_string(), v.to_string())),
3641        );
3642
3643        // Only `header.*` props become headers, and the prefix is stripped. Header names are
3644        // matched case-insensitively, so the delegation prop's mixed-case spelling still lands.
3645        let headers = IcebergCatalogConnection::catalog_headers(&props).expect("valid headers");
3646        assert_eq!(headers.len(), 2);
3647        assert_eq!(headers["x-goog-user-project"], "some-project");
3648        assert_eq!(headers["x-iceberg-access-delegation"], "vended-credentials");
3649
3650        // A prop whose name is not a legal header is an error rather than a dropped header: it
3651        // would otherwise mean silently talking to the catalog differently than asked.
3652        let props = BTreeMap::from([("header.bad name".to_string(), "v".to_string())]);
3653        assert!(IcebergCatalogConnection::catalog_headers(&props).is_err());
3654    }
3655
3656    #[mz_ore::test]
3657    fn test_check_service_name() {
3658        // Customer-owned and AWS-managed endpoint services are both accepted,
3659        // as is anything else that could be an endpoint service name.
3660        for name in [
3661            "com.amazonaws.vpce.us-east-1.vpce-svc-0e123abc123198abc",
3662            "com.amazonaws.vpce.test.vpce-svc-e2e-test",
3663            "com.amazonaws.us-east-1.s3",
3664            "com.amazonaws.anything",
3665        ] {
3666            assert_eq!(
3667                AwsPrivatelinkConnection::check_service_name(name),
3668                Ok(()),
3669                "expected {name} to be accepted"
3670            );
3671        }
3672
3673        for name in [
3674            "",
3675            "com.amazonaws",
3676            "vpce-svc-0e123abc123198abc",
3677            "my-db-lb-0123456789abcdef.elb.eu-central-1.amazonaws.com",
3678            "db.internal.example.org",
3679        ] {
3680            let err = AwsPrivatelinkConnection::check_service_name(name)
3681                .expect_err("service name should be rejected");
3682            assert_eq!(err.name, name);
3683        }
3684    }
3685}