use std::sync::{
Arc, OnceLock,
atomic::{AtomicBool, AtomicU8, AtomicUsize, Ordering},
};
use strum::{AsRefStr, Display, EnumString};
use crate::sink::{SocketState, SocketStateSink};
#[derive(Clone, Copy, Debug, Default, Display, Hash, PartialEq, Eq, AsRefStr, EnumString)]
#[repr(u8)]
#[strum(serialize_all = "UPPERCASE")]
pub enum ConnectionMode {
#[default]
Active = 0,
Reconnect = 1,
Disconnect = 2,
Closed = 3,
}
impl ConnectionMode {
#[inline]
#[must_use]
pub fn from_u8(value: u8) -> Self {
match value {
0 => Self::Active,
1 => Self::Reconnect,
2 => Self::Disconnect,
3 => Self::Closed,
_ => panic!("Invalid `ConnectionMode` value: {value}"),
}
}
#[inline]
#[must_use]
pub fn from_atomic(value: &AtomicU8) -> Self {
Self::from_u8(value.load(Ordering::SeqCst))
}
pub fn request_reconnect(value: &AtomicU8) -> bool {
Self::request_reconnect_outcome(value) == ReconnectRequestOutcome::Accepted
}
pub(crate) fn request_reconnect_outcome(value: &AtomicU8) -> ReconnectRequestOutcome {
match value.compare_exchange(
Self::Active.as_u8(),
Self::Reconnect.as_u8(),
Ordering::SeqCst,
Ordering::SeqCst,
) {
Ok(_) => ReconnectRequestOutcome::Accepted,
Err(actual) => ReconnectRequestOutcome::from_rejected(Self::from_u8(actual)),
}
}
pub(crate) fn request_reconnect_with_sink(
value: &AtomicU8,
sink: Option<&SocketStateSink>,
) -> bool {
Self::request_reconnect_outcome_with_sink(value, sink) == ReconnectRequestOutcome::Accepted
}
pub(crate) fn request_reconnect_outcome_with_sink(
value: &AtomicU8,
sink: Option<&SocketStateSink>,
) -> ReconnectRequestOutcome {
sink.map_or_else(
|| Self::request_reconnect_outcome(value),
|sink| {
sink.transition_result(
value,
Self::Active,
Self::Reconnect,
SocketState::Disconnected,
)
.map_or_else(ReconnectRequestOutcome::from_rejected, |()| {
ReconnectRequestOutcome::Accepted
})
},
)
}
pub(crate) fn close_websocket_on_loss(
value: &AtomicU8,
sink: Option<&SocketStateSink>,
) -> bool {
sink.map_or_else(
|| {
value
.try_update(Ordering::SeqCst, Ordering::SeqCst, |mode| {
matches!(Self::from_u8(mode), Self::Active | Self::Reconnect)
.then_some(Self::Closed.as_u8())
})
.is_ok()
},
|sink| sink.close_on_loss(value),
)
}
pub fn request_disconnect(value: &AtomicU8) -> bool {
value
.try_update(Ordering::SeqCst, Ordering::SeqCst, |mode| {
(!Self::from_u8(mode).is_closed()).then_some(Self::Disconnect.as_u8())
})
.is_ok()
}
pub(crate) fn complete_reconnect(value: &AtomicU8) -> ReconnectOutcome {
if value
.compare_exchange(
Self::Reconnect.as_u8(),
Self::Active.as_u8(),
Ordering::SeqCst,
Ordering::SeqCst,
)
.is_ok()
{
ReconnectOutcome::Reconnected
} else {
ReconnectOutcome::Aborted
}
}
pub(crate) fn complete_reconnect_with_sink(
value: &AtomicU8,
sink: Option<&SocketStateSink>,
) -> ReconnectOutcome {
let reconnected = sink.map_or_else(
|| Self::complete_reconnect(value) == ReconnectOutcome::Reconnected,
|sink| sink.transition(value, Self::Reconnect, Self::Active, SocketState::Connected),
);
if reconnected {
ReconnectOutcome::Reconnected
} else {
ReconnectOutcome::Aborted
}
}
#[inline]
#[must_use]
pub const fn as_u8(self) -> u8 {
self as u8
}
#[inline]
#[must_use]
pub const fn is_active(&self) -> bool {
matches!(self, Self::Active)
}
#[inline]
#[must_use]
pub const fn is_reconnect(&self) -> bool {
matches!(self, Self::Reconnect)
}
#[inline]
#[must_use]
pub const fn is_disconnect(&self) -> bool {
matches!(self, Self::Disconnect)
}
#[inline]
#[must_use]
pub const fn is_closed(&self) -> bool {
matches!(self, Self::Closed)
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ReconnectRequestOutcome {
Accepted,
AlreadyReconnecting,
Disconnected,
Closed,
Unsupported,
}
impl ReconnectRequestOutcome {
fn from_rejected(mode: ConnectionMode) -> Self {
match mode {
ConnectionMode::Active | ConnectionMode::Reconnect => Self::AlreadyReconnecting,
ConnectionMode::Disconnect => Self::Disconnected,
ConnectionMode::Closed => Self::Closed,
}
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum ReconnectOutcome {
Reconnected,
Aborted,
}
#[derive(Clone, Debug)]
pub(crate) struct ReadSessionFence {
valid: Arc<AtomicBool>,
}
impl ReadSessionFence {
#[must_use]
pub(crate) fn new() -> Self {
Self {
valid: Arc::new(AtomicBool::new(true)),
}
}
pub(crate) fn invalidate(&self) {
self.valid.store(false, Ordering::SeqCst);
}
#[must_use]
pub(crate) fn is_valid(&self) -> bool {
self.valid.load(Ordering::SeqCst)
}
}
const CONTROLLER_CLOSED: usize = 1 << (usize::BITS - 1);
const CONTROLLER_REQUEST_MASK: usize = CONTROLLER_CLOSED - 1;
pub(crate) struct ControllerLifecycle {
state: AtomicUsize,
abort_handle: OnceLock<tokio::task::AbortHandle>,
}
impl ControllerLifecycle {
pub(crate) const fn new() -> Self {
Self {
state: AtomicUsize::new(0),
abort_handle: OnceLock::new(),
}
}
pub(crate) fn enter_request(&self) -> Option<ControllerRequest<'_>> {
self.state
.try_update(Ordering::SeqCst, Ordering::SeqCst, |state| {
if state & CONTROLLER_CLOSED != 0 {
None
} else {
assert_ne!(
state, CONTROLLER_REQUEST_MASK,
"too many reconnect requests"
);
Some(state + 1)
}
})
.ok()
.map(|_| ControllerRequest(self))
}
pub(crate) fn set_abort_handle(&self, abort_handle: tokio::task::AbortHandle) {
assert!(
self.abort_handle.set(abort_handle).is_ok(),
"controller abort handle already set"
);
}
pub(crate) fn close_and_abort(&self) {
let previous = self.state.fetch_or(CONTROLLER_CLOSED, Ordering::SeqCst);
if previous & CONTROLLER_REQUEST_MASK == 0 {
self.abort();
}
}
pub(crate) fn activity(self: &Arc<Self>) -> ControllerActivity {
ControllerActivity(Arc::clone(self))
}
fn close(&self) {
self.state.fetch_or(CONTROLLER_CLOSED, Ordering::SeqCst);
}
fn abort(&self) {
if let Some(abort_handle) = self.abort_handle.get() {
abort_handle.abort();
}
}
}
pub(crate) struct ControllerRequest<'a>(&'a ControllerLifecycle);
impl Drop for ControllerRequest<'_> {
fn drop(&mut self) {
let previous = self.0.state.fetch_sub(1, Ordering::SeqCst);
if previous == CONTROLLER_CLOSED | 1 {
self.0.abort();
}
}
}
pub(crate) struct ControllerActivity(Arc<ControllerLifecycle>);
impl Drop for ControllerActivity {
fn drop(&mut self) {
self.0.close();
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use super::*;
#[rstest]
#[case(ConnectionMode::Active, true, ConnectionMode::Reconnect)]
#[case(ConnectionMode::Reconnect, false, ConnectionMode::Reconnect)]
#[case(ConnectionMode::Disconnect, false, ConnectionMode::Disconnect)]
#[case(ConnectionMode::Closed, false, ConnectionMode::Closed)]
fn request_reconnect_transitions(
#[case] start: ConnectionMode,
#[case] expected_result: bool,
#[case] expected_mode: ConnectionMode,
) {
let mode = AtomicU8::new(start.as_u8());
assert_eq!(ConnectionMode::request_reconnect(&mode), expected_result);
assert_eq!(ConnectionMode::from_atomic(&mode), expected_mode);
}
#[rstest]
#[case(
ConnectionMode::Reconnect,
ReconnectOutcome::Reconnected,
ConnectionMode::Active
)]
#[case(
ConnectionMode::Disconnect,
ReconnectOutcome::Aborted,
ConnectionMode::Disconnect
)]
#[case(
ConnectionMode::Closed,
ReconnectOutcome::Aborted,
ConnectionMode::Closed
)]
#[case(
ConnectionMode::Active,
ReconnectOutcome::Aborted,
ConnectionMode::Active
)]
fn complete_reconnect_transitions(
#[case] start: ConnectionMode,
#[case] expected_outcome: ReconnectOutcome,
#[case] expected_mode: ConnectionMode,
) {
let mode = AtomicU8::new(start.as_u8());
assert_eq!(ConnectionMode::complete_reconnect(&mode), expected_outcome);
assert_eq!(ConnectionMode::from_atomic(&mode), expected_mode);
}
#[rstest]
#[case(ConnectionMode::Active, true, ConnectionMode::Disconnect)]
#[case(ConnectionMode::Reconnect, true, ConnectionMode::Disconnect)]
#[case(ConnectionMode::Disconnect, true, ConnectionMode::Disconnect)]
#[case(ConnectionMode::Closed, false, ConnectionMode::Closed)]
fn request_disconnect_transitions(
#[case] start: ConnectionMode,
#[case] expected_result: bool,
#[case] expected_mode: ConnectionMode,
) {
let mode = AtomicU8::new(start.as_u8());
assert_eq!(ConnectionMode::request_disconnect(&mode), expected_result);
assert_eq!(ConnectionMode::from_atomic(&mode), expected_mode);
}
}