use std::collections::HashMap;
use std::marker::PhantomData;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll};
use std::time::Duration;
use dashmap::DashMap;
use futures::Stream;
use serde::de::DeserializeOwned;
use tokio_util::sync::CancellationToken;
use crate::observability::VeloMetrics;
use crate::streaming::anchor::{AnchorContext, AnchorEntry, set_active_anchor_gauge};
use crate::streaming::frame::{StreamError, StreamFrame};
use crate::streaming::handle::StreamAnchorHandle;
use super::frame::MpscFrame;
use super::types::SenderId;
pub(crate) struct MpscSenderSlot {
pub pump_token: Option<CancellationToken>,
pub stream_cancel_handle: Option<crate::streaming::control::StreamCancelHandle>,
}
#[allow(dead_code)] pub(crate) struct MpscAnchorEntry {
pub frame_tx: flume::Sender<(u64, Vec<u8>)>,
pub cancel_token: CancellationToken,
pub senders: HashMap<u64, MpscSenderSlot>,
pub next_sender_id: u64,
pub unattached_timeout: Option<Duration>,
pub timeout_cancel: Option<CancellationToken>,
pub heartbeat_interval: Duration,
pub max_senders: Option<usize>,
pub spsc_registry: Arc<DashMap<u64, AnchorEntry>>,
pub metrics: Option<Arc<VeloMetrics>>,
}
struct MpscStreamControllerInner {
local_id: u64,
registry: Arc<DashMap<u64, MpscAnchorEntry>>,
spsc_registry: Arc<DashMap<u64, AnchorEntry>>,
metrics: Option<Arc<VeloMetrics>>,
sender_registry: Arc<crate::streaming::control::SenderRegistry>,
messenger: Option<Arc<crate::messenger::Messenger>>,
cancel_wake: flume::Sender<()>,
cancelled: AtomicBool,
}
#[derive(Clone)]
pub struct MpscStreamController {
inner: Arc<MpscStreamControllerInner>,
}
impl MpscStreamController {
pub fn cancel(&self) {
if self
.inner
.cancelled
.compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
.is_err()
{
return;
}
let Some((_, entry)) = self.inner.registry.remove(&self.inner.local_id) else {
return;
};
entry.cancel_token.cancel();
if let Some(ref tc) = entry.timeout_cancel {
tc.cancel();
}
set_active_anchor_gauge(
self.inner.metrics.as_ref(),
&self.inner.spsc_registry,
&self.inner.registry,
);
let _ = self.inner.cancel_wake.try_send(());
cancel_all_senders(
&entry,
&self.inner.sender_registry,
self.inner.messenger.as_ref(),
);
}
}
pub struct MpscStreamAnchor<T> {
handle: StreamAnchorHandle,
inner_stream: flume::r#async::RecvStream<'static, (u64, Vec<u8>)>,
cancel_stream: flume::r#async::RecvStream<'static, ()>,
terminated: bool,
local_id: u64,
registry: Arc<DashMap<u64, MpscAnchorEntry>>,
controller: MpscStreamController,
_phantom: PhantomData<T>,
}
impl<T> MpscStreamAnchor<T> {
pub(crate) fn new(
handle: StreamAnchorHandle,
rx: flume::Receiver<(u64, Vec<u8>)>,
local_id: u64,
ctx: AnchorContext,
sender_registry: Arc<crate::streaming::control::SenderRegistry>,
messenger: Option<Arc<crate::messenger::Messenger>>,
) -> Self {
let AnchorContext {
registry: spsc_registry,
mpsc_registry: registry,
metrics,
} = ctx;
let (cancel_wake, cancel_rx) = flume::bounded::<()>(1);
let inner = Arc::new(MpscStreamControllerInner {
local_id,
registry: registry.clone(),
spsc_registry,
metrics,
sender_registry,
messenger,
cancel_wake,
cancelled: AtomicBool::new(false),
});
let controller = MpscStreamController { inner };
Self {
handle,
inner_stream: rx.into_stream(),
cancel_stream: cancel_rx.into_stream(),
terminated: false,
local_id,
registry,
controller,
_phantom: PhantomData,
}
}
pub fn handle(&self) -> StreamAnchorHandle {
self.handle
}
pub fn controller(&self) -> MpscStreamController {
self.controller.clone()
}
pub fn cancel(mut self) -> MpscStreamController {
self.terminated = true;
self.controller.cancel();
self.controller.clone()
}
}
impl<T> Unpin for MpscStreamAnchor<T> {}
impl<T> Drop for MpscStreamAnchor<T> {
fn drop(&mut self) {
if !self.terminated {
self.controller.cancel();
}
}
}
impl<T: DeserializeOwned> Stream for MpscStreamAnchor<T> {
type Item = Result<(SenderId, MpscFrame<T>), StreamError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if this.terminated {
return Poll::Ready(None);
}
loop {
match Pin::new(&mut this.cancel_stream).poll_next(cx) {
Poll::Ready(Some(_)) | Poll::Ready(None) => {
this.terminated = true;
return Poll::Ready(None);
}
Poll::Pending => {}
}
match Pin::new(&mut this.inner_stream).poll_next(cx) {
Poll::Ready(Some((sender_id, bytes))) => {
match rmp_serde::from_slice::<StreamFrame<T>>(&bytes) {
Ok(frame) => {
let Some(mpsc_frame) = MpscFrame::from_stream_frame(frame) else {
continue; };
if mpsc_frame.is_sender_exit() {
remove_sender_slot(&this.registry, this.local_id, sender_id);
}
return Poll::Ready(Some(Ok((SenderId(sender_id), mpsc_frame))));
}
Err(e) => {
return Poll::Ready(Some(Err(StreamError::DeserializationError(
format!("sender {}: {}", sender_id, e),
))));
}
}
}
Poll::Ready(None) => {
this.terminated = true;
return Poll::Ready(None);
}
Poll::Pending => return Poll::Pending,
}
}
}
}
pub(crate) fn cancel_all_senders(
entry: &MpscAnchorEntry,
sender_registry: &Arc<crate::streaming::control::SenderRegistry>,
messenger: Option<&Arc<crate::messenger::Messenger>>,
) {
for (_sender_id, slot) in entry.senders.iter() {
if let Some(pump_token) = &slot.pump_token {
pump_token.cancel();
}
let Some(handle) = &slot.stream_cancel_handle else {
continue;
};
let (sender_worker_id, sender_stream_id) = handle.unpack();
if let Some((_, sender_entry)) = sender_registry.senders.remove(&sender_stream_id) {
drop(sender_entry.rx_closer.lock().unwrap().take());
sender_entry.cancel_token.cancel();
}
if let Some(messenger) = messenger.cloned() {
let payload = serde_json::to_vec(&crate::streaming::control::StreamCancelRequest {
sender_stream_id,
})
.expect("serialize StreamCancelRequest");
if let Ok(rt) = tokio::runtime::Handle::try_current() {
rt.spawn(async move {
let _ = messenger
.am_send_streaming("_stream_cancel")
.expect("am_send_streaming builder")
.raw_payload(bytes::Bytes::from(payload))
.worker(sender_worker_id)
.send()
.await;
});
}
}
}
}
pub(crate) fn remove_sender_slot(
registry: &Arc<DashMap<u64, MpscAnchorEntry>>,
local_id: u64,
sender_id: u64,
) -> Option<MpscSenderSlot> {
let mut entry = registry.get_mut(&local_id)?;
let slot = entry.senders.remove(&sender_id)?;
if entry.senders.is_empty()
&& let Some(duration) = entry.unattached_timeout
{
let parent = entry.cancel_token.clone();
let spsc = entry.spsc_registry.clone();
let metrics = entry.metrics.clone();
let tc = spawn_mpsc_timeout_task_with_metrics(
registry.clone(),
Some(spsc),
metrics,
local_id,
duration,
&parent,
);
entry.timeout_cancel = Some(tc);
}
Some(slot)
}
pub(crate) fn spawn_mpsc_timeout_task_with_metrics(
registry: Arc<DashMap<u64, MpscAnchorEntry>>,
spsc_registry: Option<Arc<DashMap<u64, AnchorEntry>>>,
metrics: Option<Arc<VeloMetrics>>,
local_id: u64,
timeout: Duration,
parent_cancel: &CancellationToken,
) -> CancellationToken {
let tc = parent_cancel.child_token();
let tc_clone = tc.clone();
tokio::spawn(async move {
tokio::select! {
_ = tc_clone.cancelled() => {}
_ = tokio::time::sleep(timeout) => {
if let Some((_, entry)) = registry.remove(&local_id) {
entry.cancel_token.cancel();
if let Some(spsc) = spsc_registry.as_ref() {
set_active_anchor_gauge(metrics.as_ref(), spsc, ®istry);
}
}
}
}
});
tc
}