use std::sync::{Arc, OnceLock};
use std::time::Duration;
use dashmap::DashMap;
use serde::Serialize;
use tokio_util::sync::CancellationToken;
use crate::streaming::anchor::AnchorEntry;
use crate::streaming::frame::{SendError, StreamFrame};
use crate::streaming::handle::StreamAnchorHandle;
pub(crate) struct StreamSenderCancelInfo {
pub cancel_token: CancellationToken,
pub sender_stream_id: u64,
pub sender_registry: Arc<crate::streaming::control::SenderRegistry>,
pub poison_tx: flume::Sender<()>,
}
pub(crate) fn cached_heartbeat() -> &'static Vec<u8> {
static HEARTBEAT: OnceLock<Vec<u8>> = OnceLock::new();
HEARTBEAT.get_or_init(|| {
rmp_serde::to_vec(&StreamFrame::<()>::Heartbeat).expect("Heartbeat serializes infallibly")
})
}
pub(crate) fn cached_dropped() -> &'static Vec<u8> {
static DROPPED: OnceLock<Vec<u8>> = OnceLock::new();
DROPPED.get_or_init(|| {
rmp_serde::to_vec(&StreamFrame::<()>::Dropped).expect("Dropped serializes infallibly")
})
}
pub(crate) fn cached_finalized() -> &'static Vec<u8> {
static FINALIZED: OnceLock<Vec<u8>> = OnceLock::new();
FINALIZED.get_or_init(|| {
rmp_serde::to_vec(&StreamFrame::<()>::Finalized).expect("Finalized serializes infallibly")
})
}
pub(crate) fn cached_detached() -> &'static Vec<u8> {
static DETACHED: OnceLock<Vec<u8>> = OnceLock::new();
DETACHED.get_or_init(|| {
rmp_serde::to_vec(&StreamFrame::<()>::Detached).expect("Detached serializes infallibly")
})
}
pub(crate) fn is_terminal_sentinel(bytes: &[u8]) -> bool {
if bytes == cached_dropped().as_slice()
|| bytes == cached_detached().as_slice()
|| bytes == cached_finalized().as_slice()
{
return true;
}
matches!(
rmp_serde::from_slice::<StreamFrame<()>>(bytes),
Ok(StreamFrame::TransportError(_))
)
}
pub struct StreamSender<T> {
tx: flume::Sender<Vec<u8>>,
handle: StreamAnchorHandle,
heartbeat_cancel: CancellationToken,
sent_terminal: bool,
registry: Arc<DashMap<u64, AnchorEntry>>,
cancel_token: CancellationToken,
sender_stream_id: u64,
sender_registry: Arc<crate::streaming::control::SenderRegistry>,
poison_tx: flume::Sender<()>,
metrics: Option<Arc<crate::observability::VeloMetrics>>,
negotiated_transport: Option<velo_ext::TransportKey>,
_phantom: std::marker::PhantomData<T>,
}
impl<T> std::fmt::Debug for StreamSender<T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StreamSender")
.field("handle", &self.handle)
.field("sent_terminal", &self.sent_terminal)
.finish_non_exhaustive()
}
}
impl<T: Serialize> StreamSender<T> {
pub(crate) fn new(
tx: flume::Sender<Vec<u8>>,
handle: StreamAnchorHandle,
registry: Arc<DashMap<u64, AnchorEntry>>,
cancel: StreamSenderCancelInfo,
heartbeat_interval: Duration,
metrics: Option<Arc<crate::observability::VeloMetrics>>,
negotiated_transport: Option<velo_ext::TransportKey>,
) -> Self {
let StreamSenderCancelInfo {
cancel_token,
sender_stream_id,
sender_registry,
poison_tx,
} = cancel;
let heartbeat_cancel = CancellationToken::new();
let cancel = heartbeat_cancel.clone();
let tx_clone = tx.clone();
tokio::spawn(async move {
let mut interval = tokio::time::interval(heartbeat_interval);
interval.tick().await;
loop {
tokio::select! {
_ = cancel.cancelled() => break,
_ = interval.tick() => {
let bytes = cached_heartbeat().clone();
let _ = tx_clone.try_send(bytes);
}
}
}
});
Self {
tx,
handle,
heartbeat_cancel,
sent_terminal: false,
registry,
cancel_token,
sender_stream_id,
sender_registry,
poison_tx,
metrics,
negotiated_transport,
_phantom: std::marker::PhantomData,
}
}
pub fn negotiated_transport(&self) -> Option<&velo_ext::TransportKey> {
self.negotiated_transport.as_ref()
}
pub fn cancellation_token(&self) -> CancellationToken {
self.cancel_token.clone()
}
pub async fn send(&self, item: T) -> Result<(), SendError> {
if self.poison_tx.is_disconnected() {
return Err(SendError::ChannelClosed);
}
let bytes = rmp_serde::to_vec(&StreamFrame::Item(item))
.map_err(|e| SendError::SerializationError(e.to_string()))?;
match self.tx.try_send(bytes) {
Ok(()) => Ok(()),
Err(flume::TrySendError::Full(b)) => {
if let Some(m) = self.metrics.as_ref() {
m.record_producer_send_backpressure();
}
self.tx
.send_async(b)
.await
.map_err(|_| SendError::ChannelClosed)
}
Err(flume::TrySendError::Disconnected(_)) => Err(SendError::ChannelClosed),
}
}
pub async fn send_err(&self, msg: impl ToString) -> Result<(), SendError> {
let bytes = rmp_serde::to_vec(&StreamFrame::<()>::SenderError(msg.to_string()))
.expect("SenderError serializes infallibly");
self.tx
.send_async(bytes)
.await
.map_err(|_| SendError::ChannelClosed)
}
pub fn finalize(mut self) -> Result<(), SendError> {
self.heartbeat_cancel.cancel();
let bytes = cached_finalized().clone();
self.sent_terminal = true;
self.sender_registry.senders.remove(&self.sender_stream_id);
self.tx.send(bytes).map_err(|_| SendError::ChannelClosed)
}
pub fn detach(mut self) -> Result<StreamAnchorHandle, SendError> {
self.heartbeat_cancel.cancel();
let bytes = cached_detached().clone();
self.sent_terminal = true;
self.tx.send(bytes).map_err(|_| SendError::ChannelClosed)?;
let (_, local_id) = self.handle.unpack();
if let Some(mut entry) = self.registry.get_mut(&local_id) {
entry.attachment = false;
}
self.sender_registry.senders.remove(&self.sender_stream_id);
Ok(self.handle)
}
}
impl<T> Drop for StreamSender<T> {
fn drop(&mut self) {
self.sender_registry.senders.remove(&self.sender_stream_id);
if !self.sent_terminal {
self.heartbeat_cancel.cancel();
let bytes = cached_dropped().clone();
let _ = self.tx.send(bytes);
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use dashmap::DashMap;
use crate::streaming::anchor::AnchorEntry;
use crate::streaming::frame::{SendError, StreamFrame};
use crate::streaming::handle::StreamAnchorHandle;
use super::{StreamSender, StreamSenderCancelInfo};
fn empty_registry() -> Arc<DashMap<u64, AnchorEntry>> {
Arc::new(DashMap::new())
}
fn make_sender() -> (
StreamSender<u32>,
flume::Receiver<Vec<u8>>,
flume::Receiver<()>,
) {
let (tx, rx) = flume::bounded::<Vec<u8>>(256);
let handle = StreamAnchorHandle::pack(velo_ext::WorkerId::from_u64(1), 1);
let cancel_token = tokio_util::sync::CancellationToken::new();
let sender_registry =
std::sync::Arc::new(crate::streaming::control::SenderRegistry::default());
let (poison_tx, poison_rx) = flume::bounded::<()>(1);
let sender = StreamSender::new(
tx,
handle,
empty_registry(),
StreamSenderCancelInfo {
cancel_token,
sender_stream_id: 1,
sender_registry,
poison_tx,
},
Duration::from_secs(5),
None,
None,
);
(sender, rx, poison_rx)
}
fn make_sender_with_registry(
tx: flume::Sender<Vec<u8>>,
handle: StreamAnchorHandle,
sender_stream_id: u64,
) -> (
StreamSender<u32>,
std::sync::Arc<crate::streaming::control::SenderRegistry>,
) {
let cancel_token = tokio_util::sync::CancellationToken::new();
let sender_registry =
std::sync::Arc::new(crate::streaming::control::SenderRegistry::default());
let (poison_tx, poison_rx) = flume::bounded::<()>(1);
let entry = crate::streaming::control::SenderEntry {
cancel_token: cancel_token.clone(),
rx_closer: std::sync::Mutex::new(Some(poison_rx)),
};
sender_registry.senders.insert(sender_stream_id, entry);
let sender = StreamSender::new(
tx,
handle,
empty_registry(),
StreamSenderCancelInfo {
cancel_token,
sender_stream_id,
sender_registry: sender_registry.clone(),
poison_tx,
},
Duration::from_secs(5),
None,
None,
);
(sender, sender_registry)
}
fn decode<T: serde::de::DeserializeOwned>(bytes: &[u8]) -> StreamFrame<T> {
rmp_serde::from_slice(bytes).expect("deserialize StreamFrame")
}
#[tokio::test]
async fn test_heartbeat_emits() {
tokio::time::pause();
let (sender, rx, _poison_rx) = make_sender();
tokio::time::sleep(Duration::from_secs(6)).await;
let bytes = rx.try_recv().expect("should receive heartbeat frame");
let frame: StreamFrame<u32> = decode(&bytes);
assert!(
matches!(frame, StreamFrame::Heartbeat),
"expected Heartbeat, got {:?}",
frame
);
drop(sender);
}
#[tokio::test]
async fn test_send_item() {
let (sender, rx, _poison_rx) = make_sender();
sender.send(42u32).await.expect("send should succeed");
let bytes = rx.recv_async().await.expect("should receive item");
let frame: StreamFrame<u32> = decode(&bytes);
match frame {
StreamFrame::Item(val) => assert_eq!(val, 42),
other => panic!("expected Item(42), got {:?}", other),
}
drop(sender);
}
#[tokio::test]
async fn test_send_err() {
let (sender, rx, _poison_rx) = make_sender();
sender
.send_err("something went wrong")
.await
.expect("send_err should succeed");
let bytes = rx.recv_async().await.expect("should receive error frame");
let frame: StreamFrame<u32> = decode(&bytes);
match frame {
StreamFrame::SenderError(msg) => assert_eq!(msg, "something went wrong"),
other => panic!("expected SenderError, got {:?}", other),
}
drop(sender);
}
#[tokio::test]
async fn test_finalize() {
let (sender, rx, _poison_rx) = make_sender();
sender.finalize().expect("finalize should succeed");
let bytes = rx.recv_async().await.expect("should receive Finalized");
let frame: StreamFrame<u32> = decode(&bytes);
assert!(
matches!(frame, StreamFrame::Finalized),
"expected Finalized, got {:?}",
frame
);
while let Ok(bytes) = rx.try_recv() {
let frame: StreamFrame<u32> = decode(&bytes);
assert!(
!matches!(frame, StreamFrame::Dropped),
"should NOT receive Dropped after finalize"
);
}
}
#[tokio::test]
async fn test_detach() {
let (sender, rx, _poison_rx) = make_sender();
let expected_handle = StreamAnchorHandle::pack(velo_ext::WorkerId::from_u64(1), 1);
let returned_handle = sender.detach().expect("detach should succeed");
assert_eq!(returned_handle, expected_handle);
let bytes = rx.recv_async().await.expect("should receive Detached");
let frame: StreamFrame<u32> = decode(&bytes);
assert!(
matches!(frame, StreamFrame::Detached),
"expected Detached, got {:?}",
frame
);
while let Ok(bytes) = rx.try_recv() {
let frame: StreamFrame<u32> = decode(&bytes);
assert!(
!matches!(frame, StreamFrame::Dropped),
"should NOT receive Dropped after detach"
);
}
}
#[tokio::test]
async fn test_drop_sends_dropped() {
let (sender, rx, _poison_rx) = make_sender();
drop(sender);
let bytes = rx.recv_async().await.expect("should receive Dropped");
let frame: StreamFrame<u32> = decode(&bytes);
assert!(
matches!(frame, StreamFrame::Dropped),
"expected Dropped, got {:?}",
frame
);
}
#[tokio::test]
async fn test_heartbeat_non_blocking_on_full_channel() {
tokio::time::pause();
let (tx, rx) = flume::bounded::<Vec<u8>>(1);
let handle = StreamAnchorHandle::pack(velo_ext::WorkerId::from_u64(1), 1);
let cancel_token = tokio_util::sync::CancellationToken::new();
let sender_registry =
std::sync::Arc::new(crate::streaming::control::SenderRegistry::default());
let (poison_tx, _poison_rx) = flume::bounded::<()>(1);
let sender = StreamSender::new(
tx,
handle,
empty_registry(),
StreamSenderCancelInfo {
cancel_token,
sender_stream_id: 1,
sender_registry,
poison_tx,
},
Duration::from_secs(5),
None,
None,
);
sender.send(99u32).await.expect("send should succeed");
tokio::time::sleep(Duration::from_secs(6)).await;
let bytes = rx.try_recv().expect("should have the sent item");
let frame: StreamFrame<u32> = decode(&bytes);
assert!(matches!(frame, StreamFrame::Item(99)));
drop(sender);
}
#[tokio::test]
async fn test_send_on_closed_channel() {
let (tx, rx) = flume::bounded::<Vec<u8>>(256);
let handle = StreamAnchorHandle::pack(velo_ext::WorkerId::from_u64(1), 1);
let cancel_token = tokio_util::sync::CancellationToken::new();
let sender_registry =
std::sync::Arc::new(crate::streaming::control::SenderRegistry::default());
let (poison_tx, _poison_rx) = flume::bounded::<()>(1);
let sender = StreamSender::new(
tx,
handle,
empty_registry(),
StreamSenderCancelInfo {
cancel_token,
sender_stream_id: 1,
sender_registry,
poison_tx,
},
Duration::from_secs(5),
None,
None,
);
drop(rx);
let result = sender.send(42u32).await;
assert!(
matches!(result, Err(SendError::ChannelClosed)),
"expected ChannelClosed, got {:?}",
result
);
drop(sender);
}
#[tokio::test]
async fn test_heartbeat_stops_after_cancel() {
tokio::time::pause();
let (sender, rx, _poison_rx) = make_sender();
sender.finalize().expect("finalize should succeed");
while rx.try_recv().is_ok() {}
tokio::time::sleep(Duration::from_secs(10)).await;
assert!(
rx.try_recv().is_err(),
"should NOT receive any frames after heartbeat cancel"
);
}
#[test]
fn test_cancellation_token() {
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let (sender, _rx, _poison_rx) = make_sender();
let token = sender.cancellation_token();
assert!(!token.is_cancelled(), "token should not start cancelled");
token.cancel();
assert!(
token.is_cancelled(),
"token should be cancelled after cancel()"
);
let clone = sender.cancellation_token();
assert!(
clone.is_cancelled(),
"cloned token should reflect cancellation"
);
drop(sender);
});
}
#[tokio::test]
async fn test_send_after_cancel() {
let (frame_tx, _frame_rx) = flume::bounded::<Vec<u8>>(256);
let handle = StreamAnchorHandle::pack(velo_ext::WorkerId::from_u64(1), 2);
let cancel_token = tokio_util::sync::CancellationToken::new();
let sender_registry =
std::sync::Arc::new(crate::streaming::control::SenderRegistry::default());
let (poison_tx, poison_rx) = flume::bounded::<()>(1);
let sender = StreamSender::<u32>::new(
frame_tx,
handle,
empty_registry(),
StreamSenderCancelInfo {
cancel_token,
sender_stream_id: 2,
sender_registry,
poison_tx,
},
Duration::from_secs(5),
None,
None,
);
drop(poison_rx);
let result = sender.send(42u32).await;
assert!(
matches!(result, Err(SendError::ChannelClosed)),
"expected ChannelClosed after poison_rx drop, got {:?}",
result
);
drop(sender);
}
#[tokio::test]
async fn test_registry_cleanup_on_finalize() {
let (tx, _rx) = flume::bounded::<Vec<u8>>(256);
let handle = StreamAnchorHandle::pack(velo_ext::WorkerId::from_u64(1), 1);
let (sender, registry) = make_sender_with_registry(tx, handle, 1);
assert!(
registry.senders.contains_key(&1),
"entry should be present before finalize"
);
sender.finalize().expect("finalize should succeed");
assert!(
!registry.senders.contains_key(&1),
"entry should be removed after finalize"
);
}
#[tokio::test]
async fn test_registry_cleanup_on_detach() {
let (tx, _rx) = flume::bounded::<Vec<u8>>(256);
let handle = StreamAnchorHandle::pack(velo_ext::WorkerId::from_u64(1), 1);
let (sender, registry) = make_sender_with_registry(tx, handle, 1);
assert!(
registry.senders.contains_key(&1),
"entry should be present before detach"
);
sender.detach().expect("detach should succeed");
assert!(
!registry.senders.contains_key(&1),
"entry should be removed after detach"
);
}
#[test]
fn test_registry_cleanup_on_drop() {
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let (tx, _rx) = flume::bounded::<Vec<u8>>(256);
let handle = StreamAnchorHandle::pack(velo_ext::WorkerId::from_u64(1), 1);
let (sender, registry) = make_sender_with_registry(tx, handle, 1);
assert!(
registry.senders.contains_key(&1),
"entry should be present before drop"
);
drop(sender);
assert!(
!registry.senders.contains_key(&1),
"entry should be removed after drop"
);
});
}
}