use alloc::boxed::Box;
use alloc::format;
use alloc::string::String;
use alloc::vec::Vec;
use core::fmt;
use core::pin::Pin;
use core::task::{Context, Poll};
use futures_core::Stream;
use crate::SdkError;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ConnectionState {
Connecting,
Connected,
Reconnecting,
Disconnected {
reason: DisconnectReason,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DisconnectReason {
Normal,
Error,
Timeout,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ConnectionEvent {
pub previous: ConnectionState,
pub current: ConnectionState,
}
impl ConnectionEvent {
#[must_use]
pub const fn new(previous: ConnectionState, current: ConnectionState) -> Self {
Self { previous, current }
}
}
pub struct ConnectionEvents<S> {
inner: S,
}
impl<S> ConnectionEvents<S> {
#[must_use]
pub const fn new(inner: S) -> Self {
Self { inner }
}
#[must_use]
pub const fn inner(&self) -> &S {
&self.inner
}
#[must_use]
pub fn into_inner(self) -> S {
self.inner
}
}
impl<S: Clone> Clone for ConnectionEvents<S> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl<S> fmt::Debug for ConnectionEvents<S> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.debug_struct("ConnectionEvents").finish()
}
}
impl<S> Stream for ConnectionEvents<S>
where
S: Stream<Item = ConnectionEvent> + Unpin,
{
type Item = ConnectionEvent;
fn poll_next(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let stream = &mut self.as_mut().get_mut().inner;
Pin::new(stream).poll_next(context)
}
}
type ConnectionObserver = Box<dyn FnMut(&ConnectionEvent) + Send>;
pub struct ConnectionLifecycle {
state: ConnectionState,
observers: Vec<ConnectionObserver>,
}
impl ConnectionLifecycle {
#[must_use]
pub const fn new() -> Self {
Self {
state: ConnectionState::Connecting,
observers: Vec::new(),
}
}
#[must_use]
pub const fn state(&self) -> &ConnectionState {
&self.state
}
pub fn observe(&mut self, observer: impl FnMut(&ConnectionEvent) + Send + 'static) {
self.observers.push(Box::new(observer));
}
pub fn connect(&mut self) -> Result<(), SdkError> {
match self.state {
ConnectionState::Disconnected { .. } => {
self.transition(ConnectionState::Connecting);
Ok(())
}
_ => Err(invalid_transition(&self.state, "Connecting")),
}
}
pub fn connected(&mut self) -> Result<(), SdkError> {
match self.state {
ConnectionState::Connecting | ConnectionState::Reconnecting => {
self.transition(ConnectionState::Connected);
Ok(())
}
_ => Err(invalid_transition(&self.state, "Connected")),
}
}
pub fn reconnect_started(&mut self) -> Result<(), SdkError> {
match self.state {
ConnectionState::Connecting | ConnectionState::Connected => {
self.transition(ConnectionState::Reconnecting);
Ok(())
}
ConnectionState::Reconnecting => Err(invalid_transition(&self.state, "Reconnecting")),
ConnectionState::Disconnected { .. } => {
Err(invalid_transition(&self.state, "Reconnecting"))
}
}
}
pub fn disconnect(&mut self, reason: DisconnectReason) -> Result<(), SdkError> {
match self.state {
ConnectionState::Connecting
| ConnectionState::Connected
| ConnectionState::Reconnecting => {
self.transition(ConnectionState::Disconnected { reason });
Ok(())
}
ConnectionState::Disconnected { .. } => {
Err(invalid_transition(&self.state, "Disconnected"))
}
}
}
fn transition(&mut self, next: ConnectionState) {
let previous = core::mem::replace(&mut self.state, next);
let event = ConnectionEvent::new(previous, self.state.clone());
for observer in &mut self.observers {
observer(&event);
}
}
}
impl Default for ConnectionLifecycle {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for ConnectionLifecycle {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ConnectionLifecycle")
.field("state", &self.state)
.field("observers", &self.observers.len())
.finish()
}
}
fn invalid_transition(previous: &ConnectionState, requested: &str) -> SdkError {
connection_error(format!(
"invalid connection transition from {previous:?} to {requested}"
))
}
const fn connection_error(description: String) -> SdkError {
SdkError::Connection { description }
}
#[cfg(test)]
mod tests;