use super::connection::{CloseReason, ConnectionInfo, ConnectionProtocol, ConnectionState};
use super::origin::OriginKey;
use super::partition::PartitionId;
use super::stats::CellConnectionStats;
use crate::sync::Arc as PoolArc;
use aws_smithy_async::time::SharedTimeSource;
use aws_smithy_runtime_api::client::result::ConnectorError;
use std::error::Error;
use std::fmt;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, SystemTime};
pub trait ConnectionEventListener: Send + Sync + 'static {
fn on_event(&self, event: &ConnectionEvent<'_>);
}
impl<F> ConnectionEventListener for F
where
F: for<'a> Fn(&ConnectionEvent<'a>) + Send + Sync + 'static,
{
fn on_event(&self, event: &ConnectionEvent<'_>) {
self(event);
}
}
#[derive(Clone)]
pub struct SharedConnectionEventListener(Arc<dyn ConnectionEventListener>);
impl SharedConnectionEventListener {
pub fn new(listener: impl ConnectionEventListener) -> Self {
Self(Arc::new(listener))
}
fn notify(&self, event: &ConnectionEvent<'_>) {
#[cfg(panic = "unwind")]
if std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| self.on_event(event))).is_err()
{
tracing::warn!("connection event listener panicked");
}
#[cfg(not(panic = "unwind"))]
self.on_event(event);
}
pub(super) fn logical_close(
&self,
connection: &PoolArc<ConnectionInfo>,
cause: LogicalCloseCause,
) {
self.notify(&ConnectionEvent::LogicalClose(ConnectionLogicalClose {
connection,
cause,
}));
}
pub(super) fn physical_close(&self, connection: &PoolArc<ConnectionInfo>, reason: CloseReason) {
self.notify(&ConnectionEvent::PhysicalClose(ConnectionPhysicalClose {
connection,
reason,
}));
}
}
impl ConnectionEventListener for SharedConnectionEventListener {
fn on_event(&self, event: &ConnectionEvent<'_>) {
self.0.on_event(event);
}
}
impl fmt::Debug for SharedConnectionEventListener {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("SharedConnectionEventListener")
}
}
#[derive(Debug)]
#[non_exhaustive]
pub enum ConnectionEvent<'a> {
EstablishmentFailed(ConnectionEstablishmentFailed<'a>),
Opened(ConnectionOpened<'a>),
LogicalClose(ConnectionLogicalClose<'a>),
PhysicalClose(ConnectionPhysicalClose<'a>),
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub struct ConnectionEstablishmentId(u64);
impl fmt::Display for ConnectionEstablishmentId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(f)
}
}
#[derive(Debug)]
pub struct ConnectionEstablishmentInfo {
id: ConnectionEstablishmentId,
origin: OriginKey,
partition: PartitionId,
}
impl ConnectionEstablishmentInfo {
pub fn id(&self) -> ConnectionEstablishmentId {
self.id
}
pub fn origin(&self) -> &OriginKey {
&self.origin
}
pub fn partition(&self) -> PartitionId {
self.partition
}
}
#[derive(Clone, Copy, Debug)]
#[non_exhaustive]
pub struct ConnectionEstablishmentStats {
total_duration: Duration,
transport_duration: Duration,
protocol_handshake_duration: Option<Duration>,
}
impl ConnectionEstablishmentStats {
pub fn total_duration(&self) -> Duration {
self.total_duration
}
pub fn transport_duration(&self) -> Duration {
self.transport_duration
}
pub fn protocol_handshake_duration(&self) -> Option<Duration> {
self.protocol_handshake_duration
}
fn connection_metadata(
&self,
) -> aws_smithy_runtime_api::client::connection::ConnectionEstablishmentMetadata {
let mut builder =
aws_smithy_runtime_api::client::connection::ConnectionEstablishmentMetadata::builder()
.total_duration(self.total_duration)
.transport_duration(self.transport_duration);
builder.set_protocol_handshake_duration(self.protocol_handshake_duration);
builder.build()
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum ConnectionEstablishmentStage {
Transport,
ProtocolSelection,
ProtocolHandshake,
PoolInstallation,
}
#[derive(Debug)]
pub struct ConnectionEstablishmentFailed<'a> {
establishment: &'a ConnectionEstablishmentInfo,
stats: &'a ConnectionEstablishmentStats,
stage: ConnectionEstablishmentStage,
remote_addr: Option<SocketAddr>,
protocol: Option<ConnectionProtocol>,
error: &'a (dyn Error + Send + Sync),
}
impl<'a> ConnectionEstablishmentFailed<'a> {
pub fn establishment(&self) -> &'a ConnectionEstablishmentInfo {
self.establishment
}
pub fn stats(&self) -> &'a ConnectionEstablishmentStats {
self.stats
}
pub fn stage(&self) -> ConnectionEstablishmentStage {
self.stage
}
pub fn remote_addr(&self) -> Option<SocketAddr> {
self.remote_addr
}
pub fn protocol(&self) -> Option<ConnectionProtocol> {
self.protocol
}
pub fn error(&self) -> &'a (dyn Error + Send + Sync) {
self.error
}
}
#[derive(Debug)]
pub struct ConnectionOpened<'a> {
establishment: &'a ConnectionEstablishmentInfo,
stats: &'a ConnectionEstablishmentStats,
connection: &'a PoolArc<ConnectionInfo>,
}
impl<'a> ConnectionOpened<'a> {
pub fn establishment(&self) -> &'a ConnectionEstablishmentInfo {
self.establishment
}
pub fn stats(&self) -> &'a ConnectionEstablishmentStats {
self.stats
}
pub fn connection(&self) -> &'a PoolArc<ConnectionInfo> {
self.connection
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[non_exhaustive]
pub enum LogicalCloseCause {
IdleTimeout,
Poisoned,
ProtocolEnded,
IncompleteH1Exchange,
Reclaimed,
PoolDropped,
OwnerRuntimeShutdown,
}
impl LogicalCloseCause {
pub(super) fn from_reason(reason: CloseReason) -> Self {
match reason {
CloseReason::IdleTimeout => Self::IdleTimeout,
CloseReason::Poisoned => Self::Poisoned,
CloseReason::ProtocolClosed | CloseReason::Upgraded => Self::ProtocolEnded,
CloseReason::IncompleteH1Exchange => Self::IncompleteH1Exchange,
CloseReason::Reclaimed => Self::Reclaimed,
CloseReason::PoolDropped => Self::PoolDropped,
CloseReason::OwnerRuntimeShutdown => Self::OwnerRuntimeShutdown,
}
}
}
#[derive(Debug)]
pub struct ConnectionLogicalClose<'a> {
connection: &'a PoolArc<ConnectionInfo>,
cause: LogicalCloseCause,
}
impl<'a> ConnectionLogicalClose<'a> {
pub fn connection(&self) -> &'a PoolArc<ConnectionInfo> {
self.connection
}
pub fn cause(&self) -> LogicalCloseCause {
self.cause
}
}
#[derive(Debug)]
pub struct ConnectionPhysicalClose<'a> {
connection: &'a PoolArc<ConnectionInfo>,
reason: CloseReason,
}
impl<'a> ConnectionPhysicalClose<'a> {
pub fn connection(&self) -> &'a PoolArc<ConnectionInfo> {
self.connection
}
pub fn reason(&self) -> CloseReason {
self.reason
}
}
#[derive(Debug)]
pub(super) struct ConnectionEvents {
listener: Option<SharedConnectionEventListener>,
time_source: SharedTimeSource,
next_establishment_id: AtomicU64,
}
impl ConnectionEvents {
pub(super) fn new(
listener: Option<SharedConnectionEventListener>,
time_source: SharedTimeSource,
) -> Self {
Self {
listener,
time_source,
next_establishment_id: AtomicU64::new(0),
}
}
pub(super) fn establishment_started(
&self,
origin: &OriginKey,
partition: PartitionId,
connection_stats: PoolArc<CellConnectionStats>,
) -> ConnectionEstablishment {
connection_stats.establishment_started();
let timing = EstablishmentTiming {
time_source: self.time_source.clone(),
started_at: self.time_source.now(),
protocol_handshake_started_at: None,
transport_duration: None,
protocol_handshake_duration: None,
};
let observation = self
.listener
.as_ref()
.map(|listener| EstablishmentObservation {
listener: listener.clone(),
info: ConnectionEstablishmentInfo {
id: ConnectionEstablishmentId(
self.next_establishment_id.fetch_add(1, Ordering::Relaxed),
),
origin: origin.clone(),
partition,
},
stage: ConnectionEstablishmentStage::Transport,
remote_addr: None,
protocol: None,
});
ConnectionEstablishment {
timing,
observation,
successful_stats: None,
connection_stats: Some(connection_stats),
}
}
}
pub(super) struct ConnectionEstablishment {
timing: EstablishmentTiming,
observation: Option<EstablishmentObservation>,
successful_stats: Option<ConnectionEstablishmentStats>,
connection_stats: Option<PoolArc<CellConnectionStats>>,
}
struct EstablishmentTiming {
time_source: SharedTimeSource,
started_at: SystemTime,
protocol_handshake_started_at: Option<SystemTime>,
transport_duration: Option<Duration>,
protocol_handshake_duration: Option<Duration>,
}
struct EstablishmentObservation {
listener: SharedConnectionEventListener,
info: ConnectionEstablishmentInfo,
stage: ConnectionEstablishmentStage,
remote_addr: Option<SocketAddr>,
protocol: Option<ConnectionProtocol>,
}
impl ConnectionEstablishment {
pub(super) fn transport_completed(&mut self, remote_addr: Option<SocketAddr>) {
self.timing.transport_duration = Some(self.timing.elapsed_since(self.timing.started_at));
if let Some(observation) = &mut self.observation {
observation.remote_addr = remote_addr;
observation.stage = ConnectionEstablishmentStage::ProtocolSelection;
}
}
pub(super) fn protocol_selected(&mut self, protocol: ConnectionProtocol) {
if let Some(observation) = &mut self.observation {
observation.protocol = Some(protocol);
}
}
pub(super) fn protocol_handshake_started(&mut self) {
self.timing.protocol_handshake_started_at = Some(self.timing.time_source.now());
if let Some(observation) = &mut self.observation {
observation.stage = ConnectionEstablishmentStage::ProtocolHandshake;
}
}
pub(super) fn protocol_handshake_failed(&mut self) {
self.timing.finish_protocol_handshake();
}
pub(super) fn protocol_handshake_completed(&mut self) {
self.timing.finish_protocol_handshake();
if let Some(observation) = &mut self.observation {
observation.stage = ConnectionEstablishmentStage::PoolInstallation;
}
}
pub(super) fn installed(&mut self, connection: &PoolArc<ConnectionState>) {
let stats = self.timing.stats();
connection
.info()
.set_establishment(stats.connection_metadata());
assert!(
self.successful_stats.replace(stats).is_none(),
"connection establishment was installed more than once"
);
}
pub(super) fn failed(mut self, error: &ConnectorError) {
let stats = self.timing.stats();
self.finish_connection_stats();
let Some(observation) = self.observation.take() else {
return;
};
observation.notify_failure(error, &stats);
}
pub(super) fn opened(mut self, connection: &PoolArc<ConnectionState>) {
let stats = self.successful_stats.take().unwrap_or_else(|| {
let stats = self.timing.stats();
connection
.info()
.set_establishment(stats.connection_metadata());
stats
});
self.finish_connection_stats();
let listener = self.observation.take().map(|observation| {
observation
.listener
.notify(&ConnectionEvent::Opened(ConnectionOpened {
establishment: &observation.info,
stats: &stats,
connection: connection.info(),
}));
observation.listener
});
connection.complete_opened_event(listener.as_ref());
}
pub(super) fn superseded(mut self) {
self.finish_connection_stats();
self.observation.take();
}
fn finish_connection_stats(&mut self) {
if let Some(stats) = self.connection_stats.take() {
stats.establishment_finished();
}
}
}
impl Drop for ConnectionEstablishment {
fn drop(&mut self) {
let stats = self.successful_stats.unwrap_or_else(|| self.timing.stats());
self.finish_connection_stats();
let Some(observation) = self.observation.take() else {
return;
};
let error = ConnectorError::io("connection establishment task was dropped".into());
observation.notify_failure(&error, &stats);
}
}
impl EstablishmentObservation {
fn notify_failure(&self, error: &ConnectorError, stats: &ConnectionEstablishmentStats) {
self.listener.notify(&ConnectionEvent::EstablishmentFailed(
ConnectionEstablishmentFailed {
establishment: &self.info,
stats,
stage: self.stage,
remote_addr: self.remote_addr,
protocol: self.protocol,
error,
},
));
}
}
impl EstablishmentTiming {
fn stats(&self) -> ConnectionEstablishmentStats {
let total_duration = self.elapsed_since(self.started_at);
ConnectionEstablishmentStats {
total_duration,
transport_duration: self.transport_duration.unwrap_or(total_duration),
protocol_handshake_duration: self.protocol_handshake_duration.or_else(|| {
self.protocol_handshake_started_at
.map(|started_at| self.elapsed_since(started_at))
}),
}
}
fn finish_protocol_handshake(&mut self) {
self.protocol_handshake_duration = self
.protocol_handshake_started_at
.take()
.map(|started_at| self.elapsed_since(started_at));
}
fn elapsed_since(&self, started_at: SystemTime) -> Duration {
match self.time_source.now().duration_since(started_at) {
Ok(duration) => duration,
Err(error) => {
tracing::warn!(?error, "connection establishment clock moved backwards");
Duration::ZERO
}
}
}
}
#[cfg(all(test, not(smithy_http_client_loom)))]
mod tests {
use super::*;
use aws_smithy_async::{test_util::ManualTimeSource, time::StaticTimeSource};
use std::sync::atomic::AtomicUsize;
use std::sync::Mutex;
use std::time::UNIX_EPOCH;
fn origin() -> OriginKey {
OriginKey::from_parts(http_1x::uri::Scheme::HTTPS, "example.com", None).unwrap()
}
fn events(listener: Option<SharedConnectionEventListener>) -> ConnectionEvents {
ConnectionEvents::new(listener, SharedTimeSource::default())
}
fn establishment(events: &ConnectionEvents) -> ConnectionEstablishment {
events.establishment_started(
&origin(),
PartitionId::from_index(1),
PoolArc::new(CellConnectionStats::default()),
)
}
#[test]
fn disabled_callbacks_do_not_allocate_establishment_ids() {
let events = events(None);
events.next_establishment_id.store(7, Ordering::Relaxed);
establishment(&events).superseded();
assert_eq!(7, events.next_establishment_id.load(Ordering::Relaxed));
}
#[test]
fn listener_panic_is_isolated() {
let events = events(Some(SharedConnectionEventListener::new(
|_: &ConnectionEvent<'_>| panic!("listener failed"),
)));
let error = ConnectorError::io("synthetic transport failure".into());
establishment(&events).failed(&error);
}
#[test]
fn dropping_active_establishment_emits_transport_failure() {
let observed = Arc::new(Mutex::new(None));
let events = events(Some(SharedConnectionEventListener::new({
let observed = observed.clone();
move |event: &ConnectionEvent<'_>| {
let ConnectionEvent::EstablishmentFailed(failed) = event else {
panic!("unexpected event: {event:?}");
};
*observed.lock().unwrap() = Some((failed.establishment().id(), failed.stage()));
}
})));
drop(establishment(&events));
assert_eq!(
Some((
ConnectionEstablishmentId(0),
ConnectionEstablishmentStage::Transport
)),
*observed.lock().unwrap()
);
}
#[test]
fn failed_callback_observes_finished_establishment_count() {
let stats = PoolArc::new(CellConnectionStats::default());
let observed = Arc::new(Mutex::new(None));
let events = events(Some(SharedConnectionEventListener::new({
let stats = stats.clone();
let observed = observed.clone();
move |event: &ConnectionEvent<'_>| {
let ConnectionEvent::EstablishmentFailed(_) = event else {
panic!("unexpected event: {event:?}");
};
*observed.lock().unwrap() =
Some(stats.snapshot(0, 0, 0, 0, 0).establishing_connections());
}
})));
let establishment =
events.establishment_started(&origin(), PartitionId::from_index(1), stats.clone());
assert_eq!(1, stats.snapshot(0, 0, 0, 0, 0).establishing_connections());
let error = ConnectorError::io("synthetic transport failure".into());
establishment.failed(&error);
assert_eq!(Some(0), *observed.lock().unwrap());
}
#[test]
fn failed_event_carries_recorded_stage_and_protocol() {
let observed = Arc::new(Mutex::new(None));
let events = events(Some(SharedConnectionEventListener::new({
let observed = observed.clone();
move |event: &ConnectionEvent<'_>| {
let ConnectionEvent::EstablishmentFailed(failed) = event else {
panic!("unexpected event: {event:?}");
};
*observed.lock().unwrap() = Some((
failed.establishment().id(),
failed.stage(),
failed.protocol(),
failed.stats().protocol_handshake_duration().is_some(),
));
}
})));
let mut establishment = establishment(&events);
establishment.transport_completed(None);
establishment.protocol_selected(ConnectionProtocol::Http2);
establishment.protocol_handshake_started();
establishment.protocol_handshake_failed();
let error = ConnectorError::io("synthetic handshake failure".into());
establishment.failed(&error);
assert_eq!(
Some((
ConnectionEstablishmentId(0),
ConnectionEstablishmentStage::ProtocolHandshake,
Some(ConnectionProtocol::Http2),
true,
)),
*observed.lock().unwrap()
);
}
#[test]
fn establishment_stats_measure_recorded_phases() {
let time = ManualTimeSource::new(UNIX_EPOCH);
let observed = Arc::new(Mutex::new(None));
let events = ConnectionEvents::new(
Some(SharedConnectionEventListener::new({
let observed = observed.clone();
move |event: &ConnectionEvent<'_>| {
let ConnectionEvent::EstablishmentFailed(failed) = event else {
panic!("unexpected event: {event:?}");
};
*observed.lock().unwrap() = Some(*failed.stats());
}
})),
SharedTimeSource::new(time.clone()),
);
let mut establishment = establishment(&events);
time.advance(Duration::from_secs(2));
establishment.transport_completed(None);
establishment.protocol_selected(ConnectionProtocol::Http2);
establishment.protocol_handshake_started();
time.advance(Duration::from_secs(3));
establishment.protocol_handshake_completed();
time.advance(Duration::from_secs(2));
establishment.failed(&ConnectorError::io("synthetic installation failure".into()));
let stats = observed.lock().unwrap().expect("establishment stats");
assert_eq!(stats.total_duration(), Duration::from_secs(7));
assert_eq!(stats.transport_duration(), Duration::from_secs(2));
assert_eq!(
stats.protocol_handshake_duration(),
Some(Duration::from_secs(3))
);
let metadata = stats.connection_metadata();
assert_eq!(metadata.total_duration(), Duration::from_secs(7));
assert_eq!(metadata.transport_duration(), Duration::from_secs(2));
assert_eq!(
metadata.protocol_handshake_duration(),
Some(Duration::from_secs(3))
);
}
#[test]
fn backwards_clock_saturates_establishment_durations() {
let timing = EstablishmentTiming {
time_source: StaticTimeSource::new(UNIX_EPOCH).into(),
started_at: UNIX_EPOCH + Duration::from_secs(1),
protocol_handshake_started_at: Some(UNIX_EPOCH + Duration::from_secs(1)),
transport_duration: None,
protocol_handshake_duration: None,
};
let stats = timing.stats();
assert_eq!(stats.total_duration(), Duration::ZERO);
assert_eq!(stats.transport_duration(), Duration::ZERO);
assert_eq!(stats.protocol_handshake_duration(), Some(Duration::ZERO));
}
#[test]
fn superseded_establishment_emits_no_public_event() {
let observed = Arc::new(AtomicUsize::new(0));
let events = events(Some(SharedConnectionEventListener::new({
let observed = observed.clone();
move |_: &ConnectionEvent<'_>| {
observed.fetch_add(1, Ordering::Relaxed);
}
})));
establishment(&events).superseded();
assert_eq!(0, observed.load(Ordering::Relaxed));
assert_eq!(1, events.next_establishment_id.load(Ordering::Relaxed));
}
#[test]
fn closure_listener_receives_one_terminal_event() {
let observed = Arc::new(AtomicUsize::new(0));
let events = events(Some(SharedConnectionEventListener::new({
let observed = observed.clone();
move |event: &ConnectionEvent<'_>| {
assert!(matches!(event, ConnectionEvent::EstablishmentFailed(_)));
observed.fetch_add(1, Ordering::Relaxed);
}
})));
let error = ConnectorError::io("synthetic transport failure".into());
establishment(&events).failed(&error);
assert_eq!(1, observed.load(Ordering::Relaxed));
}
#[test]
fn logical_close_causes_cover_every_close_reason() {
assert_eq!(
[
LogicalCloseCause::IdleTimeout,
LogicalCloseCause::Poisoned,
LogicalCloseCause::ProtocolEnded,
LogicalCloseCause::ProtocolEnded,
LogicalCloseCause::IncompleteH1Exchange,
LogicalCloseCause::Reclaimed,
LogicalCloseCause::PoolDropped,
LogicalCloseCause::OwnerRuntimeShutdown,
],
[
CloseReason::IdleTimeout,
CloseReason::Poisoned,
CloseReason::ProtocolClosed,
CloseReason::Upgraded,
CloseReason::IncompleteH1Exchange,
CloseReason::Reclaimed,
CloseReason::PoolDropped,
CloseReason::OwnerRuntimeShutdown,
]
.map(LogicalCloseCause::from_reason)
);
}
}