Skip to main content

mz_aws_util/
lib.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10use std::net::IpAddr;
11
12use aws_config::{BehaviorVersion, ConfigLoader};
13use aws_smithy_http_client::tls::{self, rustls_provider::CryptoMode};
14use aws_smithy_runtime_api::client::dns::{DnsFuture, ResolveDns, ResolveDnsError};
15use aws_smithy_runtime_api::client::http::{HttpClient, SharedHttpClient};
16
17#[cfg(feature = "s3")]
18pub mod s3;
19#[cfg(feature = "s3")]
20pub mod s3_uploader;
21
22/// Creates an AWS SDK configuration loader with the defaults for the latest
23/// behavior version plus some Materialize-specific overrides.
24pub fn defaults() -> ConfigLoader {
25    // Use the SDK's latest behavior version. We already pin the crate versions,
26    // and CI puts version upgrades through rigorous testing, so we're happy to
27    // take the latest behavior version. We can adjust this in the future as
28    // necessary, if the AWS SDK ships a new behavior version that causes
29    // trouble.
30    let behavior_version = BehaviorVersion::latest();
31
32    // This is the only method allowed to call `aws_config::defaults`.
33    #[allow(clippy::disallowed_methods)]
34    let loader = aws_config::defaults(behavior_version);
35
36    // Install our custom HTTP client.
37    let loader = loader.http_client(http_client());
38
39    loader
40}
41
42/// Returns an HTTP client for use with the AWS SDK that is appropriately
43/// configured for Materialize.
44pub fn http_client() -> impl HttpClient {
45    aws_smithy_http_client::Builder::new()
46        .tls_provider(tls_provider())
47        .build_https()
48}
49
50/// Returns an AWS SDK HTTP client whose DNS resolver delegates to
51/// [`mz_ore::netio::resolve_address`].
52///
53/// Only the IP resolution step is overridden. The SDK still uses the original
54/// hostname for SNI and TLS certificate validation, so HTTPS endpoints work
55/// unchanged.
56pub fn http_client_with_resolver(enforce_external_addresses: bool) -> SharedHttpClient {
57    aws_smithy_http_client::Builder::new()
58        .tls_provider(tls_provider())
59        .build_with_resolver(MzAwsResolver {
60            enforce_external_addresses,
61        })
62}
63
64/// rustls on the aws-lc-rs provider, trusting the system's root certificates.
65fn tls_provider() -> tls::Provider {
66    // TODO(SEC-218): the rest of the workspace is still migrating from OpenSSL
67    // to rustls on aws-lc-rs.
68    tls::Provider::Rustls(CryptoMode::AwsLc)
69}
70
71/// A [`ResolveDns`] implementation that delegates to
72/// [`mz_ore::netio::resolve_address`], used by [`http_client_with_resolver`].
73#[derive(Clone, Debug)]
74struct MzAwsResolver {
75    enforce_external_addresses: bool,
76}
77
78impl MzAwsResolver {
79    async fn resolve(&self, host: &str) -> Result<Vec<IpAddr>, mz_ore::netio::DnsResolutionError> {
80        let ips = mz_ore::netio::resolve_address(host, self.enforce_external_addresses).await?;
81        Ok(ips.into_iter().collect())
82    }
83}
84
85impl ResolveDns for MzAwsResolver {
86    fn resolve_dns<'a>(&'a self, name: &'a str) -> DnsFuture<'a> {
87        DnsFuture::new(async move { self.resolve(name).await.map_err(ResolveDnsError::new) })
88    }
89}
90
91#[cfg(test)]
92mod tests {
93    use mz_ore::netio::DnsResolutionError;
94
95    use super::*;
96
97    #[mz_ore::test(tokio::test)]
98    #[cfg_attr(miri, ignore)]
99    async fn resolver_rejects_loopback_when_enforced() {
100        let resolver = MzAwsResolver {
101            enforce_external_addresses: true,
102        };
103        let err = resolver
104            .resolve("127.0.0.1")
105            .await
106            .expect_err("must reject loopback");
107        assert!(
108            matches!(err, DnsResolutionError::PrivateAddress),
109            "got {err:?}"
110        );
111        let err = resolver
112            .resolve_dns("127.0.0.1")
113            .await
114            .expect_err("must reject loopback");
115        let source = std::error::Error::source(&err).expect("wraps the resolution error");
116        assert!(
117            matches!(
118                source.downcast_ref::<DnsResolutionError>(),
119                Some(DnsResolutionError::PrivateAddress)
120            ),
121            "got {source:?}"
122        );
123    }
124
125    #[mz_ore::test(tokio::test)]
126    #[cfg_attr(miri, ignore)]
127    async fn resolver_allows_loopback_when_not_enforced() {
128        let resolver = MzAwsResolver {
129            enforce_external_addresses: false,
130        };
131        let addrs = resolver
132            .resolve_dns("127.0.0.1")
133            .await
134            .expect("loopback should resolve when enforcement is off");
135        assert!(addrs.contains(&IpAddr::from([127, 0, 0, 1])));
136    }
137
138    #[mz_ore::test(tokio::test)]
139    #[cfg_attr(miri, ignore)]
140    async fn resolver_allows_public_when_enforced() {
141        let resolver = MzAwsResolver {
142            enforce_external_addresses: true,
143        };
144        let addrs = resolver
145            .resolve_dns("8.8.8.8")
146            .await
147            .expect("public IP should resolve");
148        assert!(addrs.contains(&IpAddr::from([8, 8, 8, 8])));
149    }
150}