use std::sync::Arc;
use std::time::Duration;
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use tokio_util::sync::CancellationToken;
use crate::streaming::anchor::AnchorManager;
use crate::streaming::control::{DETECTION_MULTIPLIER, StreamCancelHandle};
use crate::streaming::handle::StreamAnchorHandle;
use super::anchor::{MpscAnchorEntry, MpscSenderSlot};
#[derive(Debug, Serialize, Deserialize)]
pub struct MpscAnchorAttachRequest {
pub handle: StreamAnchorHandle,
pub session_id: u64,
pub stream_cancel_handle: StreamCancelHandle,
#[serde(default)]
pub supported_transport_keys: Vec<velo_ext::TransportKey>,
}
#[derive(Debug, Serialize, Deserialize)]
pub enum MpscAnchorAttachResponse {
Ok {
streaming_transport_key: velo_ext::TransportKey,
heartbeat_interval_ms: u64,
sender_id: u64,
#[serde(default)]
routing_session_id: u64,
#[serde(default)]
initial_credit: u32,
#[serde(default)]
slot_byte_budget: u32,
},
Err {
reason: String,
},
}
#[derive(Debug, Serialize, Deserialize)]
pub struct MpscAnchorDetachRequest {
pub handle: StreamAnchorHandle,
pub sender_id: u64,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct MpscAnchorCancelRequest {
pub handle: StreamAnchorHandle,
}
pub(crate) async fn mpsc_reader_pump(
sender_id: u64,
transport_rx: flume::Receiver<Vec<u8>>,
frame_tx: flume::Sender<(u64, Vec<u8>)>,
cancel_token: CancellationToken,
mpsc_registry: Arc<DashMap<u64, MpscAnchorEntry>>,
local_id: u64,
heartbeat_deadline: Duration,
) {
let mut missed_heartbeats: u8 = 0;
loop {
tokio::select! {
_ = cancel_token.cancelled() => break,
result = tokio::time::timeout(heartbeat_deadline, transport_rx.recv_async()) => {
match result {
Ok(Ok(bytes)) => {
missed_heartbeats = 0;
let explicit_terminal = bytes == *crate::streaming::sender::cached_detached()
|| bytes == *crate::streaming::sender::cached_dropped()
|| bytes == *crate::streaming::sender::cached_finalized();
if frame_tx.send_async((sender_id, bytes)).await.is_err() {
break;
}
if explicit_terminal {
if let Some(slot) =
super::anchor::remove_sender_slot(&mpsc_registry, local_id, sender_id)
&& let Some(pt) = slot.pump_token
{
pt.cancel();
}
break;
}
}
Ok(Err(_)) => {
let dropped = crate::streaming::sender::cached_dropped().clone();
let _ = frame_tx.send_async((sender_id, dropped)).await;
super::anchor::remove_sender_slot(
&mpsc_registry,
local_id,
sender_id,
);
break;
}
Err(_timeout) => {
missed_heartbeats += 1;
if missed_heartbeats >= DETECTION_MULTIPLIER {
let dropped = crate::streaming::sender::cached_dropped().clone();
let _ = frame_tx.send_async((sender_id, dropped)).await;
super::anchor::remove_sender_slot(
&mpsc_registry,
local_id,
sender_id,
);
break;
}
}
}
}
}
}
cancel_token.cancel();
}
pub fn create_mpsc_anchor_attach_handler(manager: Arc<AnchorManager>) -> crate::messenger::Handler {
crate::messenger::Handler::typed_unary_async(
"_mpsc_anchor_attach",
move |ctx: crate::messenger::TypedContext<MpscAnchorAttachRequest>| {
let manager = manager.clone();
async move {
let req = ctx.input;
if req.handle.is_spsc_stream() {
return Ok(MpscAnchorAttachResponse::Err {
reason: format!("anchor {} is spsc; use _anchor_attach", req.handle),
});
}
let (_, local_id) = req.handle.unpack();
let heartbeat_interval = {
let entry = manager.mpsc_registry.get(&local_id);
match entry {
None => {
return Ok(MpscAnchorAttachResponse::Err {
reason: format!("mpsc anchor {} not found", req.handle),
});
}
Some(e) => {
if let Some(limit) = e.max_senders
&& e.senders.len() >= limit
{
return Ok(MpscAnchorAttachResponse::Err {
reason: format!(
"mpsc anchor {} reached max_senders limit {}",
req.handle, limit
),
});
}
e.heartbeat_interval
}
}
};
let routing_session_id = manager
.next_routing_session_id
.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
+ 1;
let selection = manager.select_streaming_transport(&req.supported_transport_keys);
let transport_rx =
match selection.transport.bind(local_id, routing_session_id).await {
Ok(rx) => rx,
Err(e) => {
return Ok(MpscAnchorAttachResponse::Err {
reason: format!("transport error: {}", e),
});
}
};
let streaming_transport_key = selection.key;
use dashmap::mapref::entry::Entry;
let (frame_tx, pump_cancel, sender_id) = match manager.mpsc_registry.entry(local_id)
{
Entry::Vacant(_) => {
return Ok(MpscAnchorAttachResponse::Err {
reason: format!("mpsc anchor {} removed during bind", req.handle),
});
}
Entry::Occupied(mut occ) => {
let entry = occ.get_mut();
if let Some(limit) = entry.max_senders
&& entry.senders.len() >= limit
{
return Ok(MpscAnchorAttachResponse::Err {
reason: format!(
"mpsc anchor {} reached max_senders limit {}",
req.handle, limit
),
});
}
let sender_id = entry.next_sender_id;
entry.next_sender_id += 1;
let pump_cancel = entry.cancel_token.child_token();
let slot = MpscSenderSlot {
pump_token: Some(pump_cancel.clone()),
stream_cancel_handle: Some(req.stream_cancel_handle),
};
entry.senders.insert(sender_id, slot);
if let Some(ref tc) = entry.timeout_cancel {
tc.cancel();
}
entry.timeout_cancel = None;
(entry.frame_tx.clone(), pump_cancel, sender_id)
}
};
let pump_registry = manager.mpsc_registry.clone();
tokio::spawn(mpsc_reader_pump(
sender_id,
transport_rx,
frame_tx,
pump_cancel,
pump_registry,
local_id,
heartbeat_interval,
));
Ok(MpscAnchorAttachResponse::Ok {
streaming_transport_key,
heartbeat_interval_ms: heartbeat_interval.as_millis() as u64,
sender_id,
routing_session_id,
initial_credit: selection.initial_credit,
slot_byte_budget: selection.slot_byte_budget,
})
}
},
)
.spawn()
.build()
}
pub fn create_mpsc_anchor_detach_handler(manager: Arc<AnchorManager>) -> crate::messenger::Handler {
crate::messenger::Handler::typed_unary_async(
"_mpsc_anchor_detach",
move |ctx: crate::messenger::TypedContext<MpscAnchorDetachRequest>| {
let manager = manager.clone();
async move {
let req = ctx.input;
let (_, local_id) = req.handle.unpack();
if let Some(slot) = super::anchor::remove_sender_slot(
&manager.mpsc_registry,
local_id,
req.sender_id,
) && let Some(pt) = slot.pump_token
{
pt.cancel();
}
Ok(())
}
},
)
.spawn()
.build()
}
pub fn create_mpsc_anchor_cancel_handler(manager: Arc<AnchorManager>) -> crate::messenger::Handler {
crate::messenger::Handler::typed_unary_async(
"_mpsc_anchor_cancel",
move |ctx: crate::messenger::TypedContext<MpscAnchorCancelRequest>| {
let manager = manager.clone();
async move {
let req = ctx.input;
let (_, local_id) = req.handle.unpack();
if let Some((_, entry)) = manager.mpsc_registry.remove(&local_id) {
entry.cancel_token.cancel();
if let Some(ref tc) = entry.timeout_cancel {
tc.cancel();
}
super::anchor::cancel_all_senders(
&entry,
&manager.sender_registry,
manager.messenger_lock.get(),
);
manager.update_active_anchor_gauge();
}
Ok(())
}
},
)
.spawn()
.build()
}
#[cfg(test)]
mod tests {
use super::*;
use velo_ext::{TransportKey, WorkerId};
#[test]
fn an_mpsc_attach_request_from_before_negotiation_advertises_nothing() {
let legacy_json = r#"{
"handle": {"hi": 1, "lo": 2},
"session_id": 3,
"stream_cancel_handle": {"hi": 4, "lo": 5}
}"#;
let decoded: MpscAnchorAttachRequest =
serde_json::from_str(legacy_json).expect("legacy mpsc request must deserialize");
assert!(decoded.supported_transport_keys.is_empty());
}
#[test]
fn an_mpsc_attach_response_from_before_negotiation_offers_no_mux() {
let legacy_json = r#"{"Ok":{
"streaming_transport_key": "tcp-stream",
"heartbeat_interval_ms": 5000,
"sender_id": 2
}}"#;
let decoded: MpscAnchorAttachResponse =
serde_json::from_str(legacy_json).expect("legacy mpsc response must deserialize");
match decoded {
MpscAnchorAttachResponse::Ok {
initial_credit,
slot_byte_budget,
..
} => {
assert_eq!(
initial_credit, 0,
"an absent credit window is a peer not offering the mux"
);
assert_eq!(
slot_byte_budget, 0,
"an absent byte cap means the default, not a refusal"
);
}
other => panic!("expected Ok, got {other:?}"),
}
}
#[test]
fn an_mpsc_attach_exchange_round_trips_its_negotiation_fields() {
let req = MpscAnchorAttachRequest {
handle: StreamAnchorHandle::pack_mpsc(WorkerId::from_u64(1), 2),
session_id: 3,
stream_cancel_handle: StreamCancelHandle::pack(WorkerId::from_u64(4), 5),
supported_transport_keys: vec![TransportKey::new("messenger-mux-v1")],
};
let decoded: MpscAnchorAttachRequest =
rmp_serde::from_slice(&rmp_serde::to_vec(&req).expect("encode")).expect("decode");
assert_eq!(
decoded
.supported_transport_keys
.iter()
.map(TransportKey::as_str)
.collect::<Vec<_>>(),
["messenger-mux-v1"],
);
let resp = MpscAnchorAttachResponse::Ok {
streaming_transport_key: TransportKey::new("messenger-mux-v1"),
heartbeat_interval_ms: 5000,
sender_id: 7,
routing_session_id: 8,
initial_credit: 256,
slot_byte_budget: 1024,
};
let decoded: MpscAnchorAttachResponse =
rmp_serde::from_slice(&rmp_serde::to_vec(&resp).expect("encode")).expect("decode");
match decoded {
MpscAnchorAttachResponse::Ok {
initial_credit,
slot_byte_budget,
..
} => {
assert_eq!(initial_credit, 256);
assert_eq!(slot_byte_budget, 1024);
}
other => panic!("expected Ok, got {other:?}"),
}
}
}