use super::super::admission::ProtocolRequirement;
use super::super::registry::PartitionState;
use crate::client::connect::{AsyncConn, BoxConn};
use crate::client::timeout::{self, TimeoutKind};
use aws_smithy_async::rt::sleep::SharedAsyncSleep;
use aws_smithy_runtime_api::box_error::BoxError;
use http_1x::Uri;
use std::future::{poll_fn, Future};
use std::pin::Pin;
use std::sync::Arc as StdArc;
use std::time::Duration;
use tower::Service;
#[derive(Clone, Debug)]
pub(in crate::client::pool) struct TransportTimeout {
duration: Duration,
sleep: SharedAsyncSleep,
}
impl TransportTimeout {
pub(in crate::client::pool) fn new(duration: Duration, sleep: SharedAsyncSleep) -> Self {
Self { duration, sleep }
}
}
pub(in crate::client::pool) type AlpnProtocols = &'static [&'static [u8]];
const HTTP_ALPN_PROTOCOLS: AlpnProtocols = &[b"h2", b"http/1.1"];
const HTTP1_ALPN_PROTOCOLS: AlpnProtocols = &[b"http/1.1"];
pub(in crate::client::pool) struct TransportConnectContext<'a> {
partition: &'a PartitionState,
uri: Uri,
timeout: Option<TransportTimeout>,
alpn_protocols: AlpnProtocols,
}
impl<'a> TransportConnectContext<'a> {
pub(super) fn new(
partition: &'a PartitionState,
uri: Uri,
timeout: Option<TransportTimeout>,
requirement: ProtocolRequirement,
) -> Self {
Self {
partition,
uri,
timeout,
alpn_protocols: alpn_protocols(requirement),
}
}
}
type TransportFuture = Pin<Box<dyn Future<Output = Result<BoxConn, BoxError>> + Send + 'static>>;
pub(in crate::client::pool) trait TransportFactory:
Send + Sync + 'static
{
fn initialize_for_partition(&self, partition: &PartitionState) -> Result<(), BoxError> {
let _ = partition;
Ok(())
}
fn can_guarantee_http1(&self) -> bool;
fn connect(&self, context: TransportConnectContext<'_>) -> TransportFuture;
}
struct ServiceTransportFactory<F> {
connector_for_interface: F,
can_guarantee_http1: bool,
}
impl<F, C, IO> TransportFactory for ServiceTransportFactory<F>
where
F: Fn(Option<&str>) -> C + Send + Sync + 'static,
C: Service<Uri, Response = IO> + Send + 'static,
C::Error: Into<BoxError>,
C::Future: Send + 'static,
IO: AsyncConn,
{
fn can_guarantee_http1(&self) -> bool {
self.can_guarantee_http1
}
fn connect(&self, context: TransportConnectContext<'_>) -> TransportFuture {
let TransportConnectContext {
partition,
uri,
timeout,
alpn_protocols: _alpn_protocols,
} = context;
let interface = partition.interface().map(|interface| interface.as_ref());
let mut connector = (self.connector_for_interface)(interface);
Box::pin(async move {
poll_fn(|cx| connector.poll_ready(cx))
.await
.map_err(Into::into)?;
let connect = connector.call(uri);
let io = timeout::maybe_timeout_future(
connect,
timeout.as_ref().map(|timeout| timeout.duration),
timeout.as_ref().map(|timeout| &timeout.sleep),
TimeoutKind::Connect,
)
.await?;
Ok(Box::new(io) as BoxConn)
})
}
}
#[cfg(any(
all(feature = "test-util", aws_sdk_unstable),
all(test, feature = "rt-tokio")
))]
pub(in crate::client::pool) fn from_connector<C, IO>(connector: C) -> StdArc<dyn TransportFactory>
where
C: Service<Uri, Response = IO> + Clone + Send + Sync + 'static,
C::Error: Into<BoxError>,
C::Future: Send + 'static,
IO: AsyncConn,
{
service_factory(move |_| connector.clone(), false)
}
pub(in crate::client::pool) fn from_interface_connector<F, C, IO>(
connector_for_interface: F,
) -> StdArc<dyn TransportFactory>
where
F: Fn(Option<&str>) -> C + Send + Sync + 'static,
C: Service<Uri, Response = IO> + Send + 'static,
C::Error: Into<BoxError>,
C::Future: Send + 'static,
IO: AsyncConn,
{
service_factory(connector_for_interface, true)
}
fn service_factory<F, C, IO>(
connector_for_interface: F,
can_guarantee_http1: bool,
) -> StdArc<dyn TransportFactory>
where
F: Fn(Option<&str>) -> C + Send + Sync + 'static,
C: Service<Uri, Response = IO> + Send + 'static,
C::Error: Into<BoxError>,
C::Future: Send + 'static,
IO: AsyncConn,
{
StdArc::new(ServiceTransportFactory {
connector_for_interface,
can_guarantee_http1,
})
}
#[cfg(any(feature = "__rustls", feature = "s2n-tls"))]
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
struct ConnectorCacheKey {
interface: Option<StdArc<str>>,
alpn_protocols: AlpnProtocols,
}
#[cfg(any(feature = "__rustls", feature = "s2n-tls"))]
impl ConnectorCacheKey {
fn for_partition(partition: &PartitionState, alpn_protocols: AlpnProtocols) -> Self {
Self {
interface: partition.interface().cloned(),
alpn_protocols,
}
}
fn interface(&self) -> Option<&str> {
self.interface.as_deref()
}
}
#[cfg(any(feature = "__rustls", feature = "s2n-tls"))]
struct CachedTransportFactory<F, C> {
factory: F,
can_guarantee_http1: bool,
connectors: crate::sync::Mutex<std::collections::HashMap<ConnectorCacheKey, C>>,
}
#[cfg(any(feature = "__rustls", feature = "s2n-tls"))]
impl<F, C> CachedTransportFactory<F, C>
where
F: Fn(Option<&str>, AlpnProtocols) -> C,
C: Clone,
{
fn connector(&self, partition: &PartitionState, alpn_protocols: AlpnProtocols) -> C {
let key = ConnectorCacheKey::for_partition(partition, alpn_protocols);
match self.connectors.lock().entry(key) {
std::collections::hash_map::Entry::Occupied(entry) => entry.get().clone(),
std::collections::hash_map::Entry::Vacant(entry) => {
let connector = (self.factory)(entry.key().interface(), entry.key().alpn_protocols);
entry.insert(connector).clone()
}
}
}
fn initialize_connectors_for_partition(&self, partition: &PartitionState) {
drop(self.connector(partition, HTTP_ALPN_PROTOCOLS));
drop(self.connector(partition, HTTP1_ALPN_PROTOCOLS));
}
}
#[cfg(any(feature = "__rustls", feature = "s2n-tls"))]
impl<F, C, IO> TransportFactory for CachedTransportFactory<F, C>
where
F: Fn(Option<&str>, AlpnProtocols) -> C + Send + Sync + 'static,
C: Service<Uri, Response = IO> + Clone + Send + Sync + 'static,
C::Error: Into<BoxError>,
C::Future: Send + 'static,
IO: AsyncConn,
{
fn initialize_for_partition(&self, partition: &PartitionState) -> Result<(), BoxError> {
self.initialize_connectors_for_partition(partition);
Ok(())
}
fn can_guarantee_http1(&self) -> bool {
self.can_guarantee_http1
}
fn connect(&self, context: TransportConnectContext<'_>) -> TransportFuture {
let TransportConnectContext {
partition,
uri,
timeout,
alpn_protocols,
} = context;
let mut connector = self.connector(partition, alpn_protocols);
Box::pin(async move {
poll_fn(|cx| connector.poll_ready(cx))
.await
.map_err(Into::into)?;
let connect = connector.call(uri);
let io = timeout::maybe_timeout_future(
connect,
timeout.as_ref().map(|timeout| timeout.duration),
timeout.as_ref().map(|timeout| &timeout.sleep),
TimeoutKind::Connect,
)
.await?;
Ok(Box::new(io) as BoxConn)
})
}
}
#[cfg(any(feature = "__rustls", feature = "s2n-tls"))]
pub(in crate::client::pool) fn from_cached_interface_connector<F, C, IO>(
connector_for_interface: F,
can_guarantee_http1: bool,
) -> StdArc<dyn TransportFactory>
where
F: Fn(Option<&str>, AlpnProtocols) -> C + Send + Sync + 'static,
C: Service<Uri, Response = IO> + Clone + Send + Sync + 'static,
C::Error: Into<BoxError>,
C::Future: Send + 'static,
IO: AsyncConn,
{
StdArc::new(CachedTransportFactory {
factory: connector_for_interface,
can_guarantee_http1,
connectors: crate::sync::Mutex::new(std::collections::HashMap::new()),
})
}
fn alpn_protocols(requirement: ProtocolRequirement) -> AlpnProtocols {
match requirement {
ProtocolRequirement::H1Required => HTTP1_ALPN_PROTOCOLS,
ProtocolRequirement::H1Compatible | ProtocolRequirement::H2Required => HTTP_ALPN_PROTOCOLS,
}
}
#[cfg(all(test, any(feature = "__rustls", feature = "s2n-tls")))]
mod tests {
use super::*;
use crate::client::pool::maintenance::MaintenanceConfig;
use crate::client::pool::partition::{
ConnectionReuseScope, DriverSpawner, Partition, PartitionId, Spawn,
};
use crate::client::pool::registry::PartitionRegistry;
use aws_smithy_runtime_api::client::http::HttpClient;
use aws_smithy_runtime_api::client::runtime_components::RuntimeComponentsBuilder;
use aws_smithy_types::config_bag::ConfigBag;
use std::sync::{Arc, Mutex};
#[derive(Debug)]
struct TestSpawner;
impl Spawn for TestSpawner {
fn spawn(&self, _: Pin<Box<dyn Future<Output = ()> + Send + 'static>>) {}
}
struct InitializationRecorder {
partitions: Arc<Mutex<Vec<PartitionId>>>,
}
impl TransportFactory for InitializationRecorder {
fn initialize_for_partition(&self, partition: &PartitionState) -> Result<(), BoxError> {
self.partitions
.lock()
.expect("initialization log is not poisoned")
.push(partition.id());
Ok(())
}
fn can_guarantee_http1(&self) -> bool {
true
}
fn connect(&self, _: TransportConnectContext<'_>) -> TransportFuture {
Box::pin(async { panic!("client validation must not start a connection") })
}
}
type Construction = (Option<String>, AlpnProtocols);
struct CachedFactoryFixture<F> {
factory: CachedTransportFactory<F, usize>,
constructions: Arc<Mutex<Vec<Construction>>>,
}
fn recording_factory() -> CachedFactoryFixture<impl Fn(Option<&str>, AlpnProtocols) -> usize> {
let constructions = Arc::new(Mutex::new(Vec::new()));
let observed = constructions.clone();
let factory = move |interface: Option<&str>, alpn_protocols: AlpnProtocols| {
let mut constructions = observed.lock().expect("construction log is not poisoned");
constructions.push((interface.map(str::to_owned), alpn_protocols));
constructions.len()
};
CachedFactoryFixture {
factory: CachedTransportFactory {
factory,
can_guarantee_http1: true,
connectors: crate::sync::Mutex::new(std::collections::HashMap::new()),
},
constructions,
}
}
fn registry(partitions: Option<Vec<Partition>>) -> PartitionRegistry {
PartitionRegistry::new(
partitions,
ConnectionReuseScope::Partition,
None,
MaintenanceConfig::default(),
)
.expect("valid partition registry")
}
#[test]
fn client_validation_initializes_its_selected_partition() {
let selected = PartitionId::from_index(7);
let initialized = Arc::new(Mutex::new(Vec::new()));
let pool = crate::client::pool::builder::Builder::default()
.partitions([Partition::new(selected, DriverSpawner::new(TestSpawner))])
.build_with_transport_for_test(Arc::new(InitializationRecorder {
partitions: initialized.clone(),
}))
.expect("valid test pool");
let client = crate::client::pool::Client::from_partition(&pool, selected)
.expect("selected partition exists");
client
.validate_base_client_config(&RuntimeComponentsBuilder::for_tests(), &ConfigBag::base())
.expect("transport initialization succeeds");
assert_eq!(
vec![selected],
*initialized
.lock()
.expect("initialization log is not poisoned")
);
}
#[test]
fn initialization_constructs_each_alpn_variant_once() {
let partition = registry(None)
.partition(PartitionId::ANONYMOUS)
.expect("anonymous partition exists");
let fixture = recording_factory();
fixture
.factory
.initialize_connectors_for_partition(&partition);
fixture
.factory
.initialize_connectors_for_partition(&partition);
let constructions = fixture
.constructions
.lock()
.expect("construction log is not poisoned");
assert_eq!(2, constructions.len());
assert_eq!(None, constructions[0].0);
assert_eq!(HTTP_ALPN_PROTOCOLS, constructions[0].1);
assert_eq!(HTTP1_ALPN_PROTOCOLS, constructions[1].1);
assert_eq!(2, fixture.factory.connectors.lock().len());
}
#[cfg(any(
target_os = "android",
target_os = "fuchsia",
target_os = "illumos",
target_os = "ios",
target_os = "linux",
target_os = "macos",
target_os = "solaris",
target_os = "tvos",
target_os = "visionos",
target_os = "watchos",
))]
#[test]
fn equal_interface_bindings_share_connector_entries() {
let first = PartitionId::from_index(1);
let second = PartitionId::from_index(2);
let registry = registry(Some(vec![
Partition::new(first, DriverSpawner::new(TestSpawner)).interface("interface-a"),
Partition::new(second, DriverSpawner::new(TestSpawner)).interface("interface-a"),
]));
let fixture = recording_factory();
fixture.factory.initialize_connectors_for_partition(
®istry.partition(first).expect("first partition exists"),
);
fixture.factory.initialize_connectors_for_partition(
®istry.partition(second).expect("second partition exists"),
);
let constructions = fixture
.constructions
.lock()
.expect("construction log is not poisoned");
assert_eq!(
[
(Some("interface-a".to_string()), HTTP_ALPN_PROTOCOLS),
(Some("interface-a".to_string()), HTTP1_ALPN_PROTOCOLS),
],
constructions.as_slice()
);
assert_eq!(2, fixture.factory.connectors.lock().len());
}
#[cfg(any(
target_os = "android",
target_os = "fuchsia",
target_os = "illumos",
target_os = "ios",
target_os = "linux",
target_os = "macos",
target_os = "solaris",
target_os = "tvos",
target_os = "visionos",
target_os = "watchos",
))]
#[test]
fn distinct_interface_bindings_use_distinct_connector_entries() {
let first = PartitionId::from_index(1);
let second = PartitionId::from_index(2);
let registry = registry(Some(vec![
Partition::new(first, DriverSpawner::new(TestSpawner)).interface("interface-a"),
Partition::new(second, DriverSpawner::new(TestSpawner)).interface("interface-b"),
]));
let fixture = recording_factory();
fixture.factory.initialize_connectors_for_partition(
®istry.partition(first).expect("first partition exists"),
);
fixture.factory.initialize_connectors_for_partition(
®istry.partition(second).expect("second partition exists"),
);
let constructions = fixture
.constructions
.lock()
.expect("construction log is not poisoned");
assert_eq!(
[
(Some("interface-a".to_string()), HTTP_ALPN_PROTOCOLS),
(Some("interface-a".to_string()), HTTP1_ALPN_PROTOCOLS),
(Some("interface-b".to_string()), HTTP_ALPN_PROTOCOLS),
(Some("interface-b".to_string()), HTTP1_ALPN_PROTOCOLS),
],
constructions.as_slice()
);
assert_eq!(4, fixture.factory.connectors.lock().len());
}
}