use std::collections::VecDeque;
use std::fmt;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use parking_lot::Mutex;
use tokio::sync::broadcast::{Receiver, Sender, WeakSender};
use tokio_stream::wrappers::BroadcastStream;
use tokio_stream::wrappers::errors::BroadcastStreamRecvError;
use tokio_stream::{Stream, StreamExt as _};
use crate::types::{SessionEvent, SessionLifecycleEvent};
use crate::{Custom, Repr};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Lagged(pub(crate) u64);
impl Lagged {
pub fn skipped(&self) -> u64 {
self.0
}
}
impl fmt::Display for Lagged {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "subscription lagged behind by {} events", self.0)
}
}
impl std::error::Error for Lagged {}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum RecvErrorKind {
Closed,
Lagged(Lagged),
}
impl fmt::Display for RecvErrorKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
RecvErrorKind::Closed => write!(f, "subscription closed"),
RecvErrorKind::Lagged(l) => write!(f, "{l}"),
}
}
}
#[derive(Debug)]
pub struct RecvError {
repr: Repr<RecvErrorKind>,
}
impl RecvError {
pub fn kind(&self) -> &RecvErrorKind {
match &self.repr {
Repr::Simple(k) | Repr::SimpleMessage(k, ..) | Repr::Custom(Custom { kind: k, .. }) => {
k
}
}
}
}
impl fmt::Display for RecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.repr {
Repr::Simple(k) => write!(f, "{k}"),
Repr::SimpleMessage(_, m) => write!(f, "{m}"),
Repr::Custom(Custom { error, .. }) => write!(f, "{error}"),
}
}
}
impl std::error::Error for RecvError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match &self.repr {
Repr::Custom(Custom { error, .. }) => Some(&**error),
_ => None,
}
}
}
impl From<RecvErrorKind> for RecvError {
fn from(kind: RecvErrorKind) -> Self {
Self {
repr: Repr::Simple(kind),
}
}
}
impl From<Lagged> for RecvError {
fn from(lagged: Lagged) -> Self {
Self::from(RecvErrorKind::Lagged(lagged))
}
}
enum ResumeBootstrapState {
Unclaimed(VecDeque<SessionEvent>),
Claimed(VecDeque<SessionEvent>),
Disabled,
}
pub(crate) struct ResumeBootstrap {
state: Mutex<ResumeBootstrapState>,
live: WeakSender<SessionEvent>,
}
pub(crate) struct ResumeBootstrapCleanup(Arc<ResumeBootstrap>);
impl Drop for ResumeBootstrapCleanup {
fn drop(&mut self) {
self.0.release_unclaimed();
}
}
impl ResumeBootstrap {
pub(crate) fn new(event_tx: &Sender<SessionEvent>) -> Arc<Self> {
Arc::new(Self {
state: Mutex::new(ResumeBootstrapState::Unclaimed(VecDeque::new())),
live: event_tx.downgrade(),
})
}
pub(crate) fn cleanup_guard(self: &Arc<Self>) -> ResumeBootstrapCleanup {
ResumeBootstrapCleanup(self.clone())
}
pub(crate) fn publish(&self, event_tx: &Sender<SessionEvent>, event: SessionEvent) {
let mut state = self.state.lock();
match &mut *state {
ResumeBootstrapState::Unclaimed(events) | ResumeBootstrapState::Claimed(events) => {
events.push_back(event.clone());
}
ResumeBootstrapState::Disabled => {}
}
let _ = event_tx.send(event);
}
pub(crate) fn subscribe(
self: &Arc<Self>,
event_tx: &Sender<SessionEvent>,
) -> EventSubscription {
let mut state = self.state.lock();
match &mut *state {
ResumeBootstrapState::Unclaimed(events) => {
let events = std::mem::take(events);
*state = ResumeBootstrapState::Claimed(events);
EventSubscription {
inner: None,
bootstrap: Some(self.clone()),
}
}
ResumeBootstrapState::Claimed(_) | ResumeBootstrapState::Disabled => {
EventSubscription::new(event_tx.subscribe())
}
}
}
fn pop(&self, live: &mut Option<BroadcastStream<SessionEvent>>) -> Option<SessionEvent> {
let mut state = self.state.lock();
let ResumeBootstrapState::Claimed(events) = &mut *state else {
return None;
};
if let Some(event) = events.pop_front() {
return Some(event);
}
*live = self
.live
.upgrade()
.map(|sender| BroadcastStream::new(sender.subscribe()));
*state = ResumeBootstrapState::Disabled;
None
}
pub(crate) fn release_unclaimed(&self) {
let mut state = self.state.lock();
if matches!(*state, ResumeBootstrapState::Unclaimed(_)) {
*state = ResumeBootstrapState::Disabled;
}
}
fn abandon(&self) {
let mut state = self.state.lock();
if matches!(*state, ResumeBootstrapState::Claimed(_)) {
*state = ResumeBootstrapState::Disabled;
}
}
}
#[must_use = "dropping the subscription unsubscribes and discards any owned resume bootstrap backlog"]
pub struct EventSubscription {
inner: Option<BroadcastStream<SessionEvent>>,
bootstrap: Option<Arc<ResumeBootstrap>>,
}
impl EventSubscription {
pub(crate) fn new(rx: Receiver<SessionEvent>) -> Self {
Self {
inner: Some(BroadcastStream::new(rx)),
bootstrap: None,
}
}
fn next_bootstrap_event(&mut self) -> Option<SessionEvent> {
let event = self
.bootstrap
.as_ref()
.and_then(|bootstrap| bootstrap.pop(&mut self.inner));
if event.is_none() {
self.bootstrap = None;
}
event
}
pub async fn recv(&mut self) -> Result<SessionEvent, RecvError> {
match self.next().await {
Some(Ok(event)) => Ok(event),
Some(Err(lagged)) => Err(lagged.into()),
None => Err(RecvErrorKind::Closed.into()),
}
}
}
impl Stream for EventSubscription {
type Item = Result<SessionEvent, Lagged>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
if let Some(event) = self.next_bootstrap_event() {
return Poll::Ready(Some(Ok(event)));
}
let Some(inner) = self.inner.as_mut() else {
return Poll::Ready(None);
};
match Pin::new(inner).poll_next(cx) {
Poll::Ready(Some(Ok(event))) => Poll::Ready(Some(Ok(event))),
Poll::Ready(Some(Err(BroadcastStreamRecvError::Lagged(n)))) => {
Poll::Ready(Some(Err(Lagged(n))))
}
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
}
impl Drop for EventSubscription {
fn drop(&mut self) {
if let Some(bootstrap) = &self.bootstrap {
bootstrap.abandon();
}
}
}
#[must_use = "dropping the subscription unsubscribes"]
pub struct LifecycleSubscription {
inner: BroadcastStream<SessionLifecycleEvent>,
}
impl LifecycleSubscription {
pub(crate) fn new(rx: Receiver<SessionLifecycleEvent>) -> Self {
Self {
inner: BroadcastStream::new(rx),
}
}
pub async fn recv(&mut self) -> Result<SessionLifecycleEvent, RecvError> {
match self.next().await {
Some(Ok(event)) => Ok(event),
Some(Err(lagged)) => Err(lagged.into()),
None => Err(RecvErrorKind::Closed.into()),
}
}
}
impl Stream for LifecycleSubscription {
type Item = Result<SessionLifecycleEvent, Lagged>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
match Pin::new(&mut self.inner).poll_next(cx) {
Poll::Ready(Some(Ok(event))) => Poll::Ready(Some(Ok(event))),
Poll::Ready(Some(Err(BroadcastStreamRecvError::Lagged(n)))) => {
Poll::Ready(Some(Err(Lagged(n))))
}
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
}
#[cfg(test)]
mod tests;