1use 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
22pub fn defaults() -> ConfigLoader {
25 let behavior_version = BehaviorVersion::latest();
31
32 #[allow(clippy::disallowed_methods)]
34 let loader = aws_config::defaults(behavior_version);
35
36 let loader = loader.http_client(http_client());
38
39 loader
40}
41
42pub fn http_client() -> impl HttpClient {
45 aws_smithy_http_client::Builder::new()
46 .tls_provider(tls_provider())
47 .build_https()
48}
49
50pub 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
64fn tls_provider() -> tls::Provider {
66 tls::Provider::Rustls(CryptoMode::AwsLc)
69}
70
71#[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}