use std::collections::{BTreeMap, HashMap};
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use crate::messenger::{Context, Handler, Messenger};
use crate::observability::{StreamingTransportMetricsHandle, VeloMetrics};
use anyhow::Result;
use dashmap::DashMap;
use futures::future::BoxFuture;
use velo_ext::WorkerId;
use crate::streaming::transport::FrameTransport;
const ANCHOR_ID_HEADER: &str = "anchor_id";
const SESSION_ID_HEADER: &str = "session_id";
const STREAM_SEQ_HEADER: &str = "seq";
const REORDER_WINDOW: usize = 4096;
type Deposit = (Option<u64>, Vec<u8>);
type DispatchMap = DashMap<(u64, u64), DispatchEntry>;
struct DispatchEntry {
token: u64,
sender: flume::Sender<Deposit>,
}
static DISPATCH_TOKEN: AtomicU64 = AtomicU64::new(0);
fn next_dispatch_token() -> u64 {
DISPATCH_TOKEN.fetch_add(1, Ordering::Relaxed)
}
pub struct VeloFrameTransport {
messenger: Arc<Messenger>,
dispatch: Arc<DispatchMap>,
worker_id: WorkerId,
backpressure_count: Arc<AtomicU64>,
streaming_metrics: Option<StreamingTransportMetricsHandle>,
}
impl VeloFrameTransport {
pub fn new(
messenger: Arc<Messenger>,
worker_id: WorkerId,
metrics: Option<Arc<VeloMetrics>>,
) -> Result<Self> {
let dispatch: Arc<DispatchMap> = Arc::new(DashMap::new());
let backpressure_count: Arc<AtomicU64> = Arc::new(AtomicU64::new(0));
let handler_dispatch = dispatch.clone();
let streaming_metrics = metrics
.as_ref()
.map(|metrics| metrics.bind_streaming_transport("velo"));
let handler = Handler::am_handler_async("_stream_data", move |ctx: Context| {
let handler_dispatch = handler_dispatch.clone();
async move {
let headers = match ctx.headers.as_ref() {
Some(h) => h,
None => {
tracing::warn!("_stream_data: missing headers, dropping frame");
return Ok(());
}
};
let anchor_id = match headers
.get(ANCHOR_ID_HEADER)
.and_then(|v| v.parse::<u64>().ok())
{
Some(id) => id,
None => {
tracing::warn!(
"_stream_data: missing or invalid {} header, dropping frame",
ANCHOR_ID_HEADER
);
return Ok(());
}
};
let session_id = match headers
.get(SESSION_ID_HEADER)
.and_then(|v| v.parse::<u64>().ok())
{
Some(id) => id,
None => {
tracing::warn!(
anchor_id,
"_stream_data: missing or invalid {} header, dropping frame",
SESSION_ID_HEADER
);
return Ok(());
}
};
let seq = match headers.get(STREAM_SEQ_HEADER) {
None => None,
Some(raw) => match raw.parse::<u64>() {
Ok(v) => Some(v),
Err(e) => {
tracing::error!(
anchor_id,
session_id,
raw = %raw,
error = %e,
"_stream_data: malformed {} header, dropping frame",
STREAM_SEQ_HEADER
);
return Ok(());
}
},
};
let frame_bytes = ctx.payload.to_vec();
let deposit_tx = handler_dispatch
.get(&(anchor_id, session_id))
.map(|entry| entry.value().sender.clone());
if let Some(tx) = deposit_tx
&& tx.send((seq, frame_bytes)).is_err()
{
}
Ok(())
}
})
.build();
messenger.register_streaming_handler(handler)?;
Ok(Self {
messenger,
dispatch,
worker_id,
backpressure_count,
streaming_metrics,
})
}
pub fn backpressure_count(&self) -> u64 {
self.backpressure_count.load(Ordering::Relaxed)
}
pub fn unbind(&self, anchor_id: u64) {
self.dispatch.retain(|&(aid, _), _| aid != anchor_id);
}
}
struct DelivererCtx {
deposit_rx: flume::Receiver<Deposit>,
token: u64,
consumer_tx: flume::Sender<Vec<u8>>,
backpressure: Arc<AtomicU64>,
metrics: Option<StreamingTransportMetricsHandle>,
dispatch: Arc<DispatchMap>,
anchor_id: u64,
session_id: u64,
}
async fn run_deliverer(ctx: DelivererCtx) {
let DelivererCtx {
deposit_rx,
token,
consumer_tx,
backpressure,
metrics,
dispatch,
anchor_id,
session_id,
} = ctx;
let mut pending: BTreeMap<u64, Vec<u8>> = BTreeMap::new();
let mut next_expected: u64 = 0;
struct Cleanup {
dispatch: Arc<DispatchMap>,
token: u64,
anchor_id: u64,
session_id: u64,
}
impl Drop for Cleanup {
fn drop(&mut self) {
self.dispatch
.remove_if(&(self.anchor_id, self.session_id), |_, entry| {
entry.token == self.token
});
}
}
let _cleanup = Cleanup {
dispatch,
token,
anchor_id,
session_id,
};
loop {
let (seq_opt, bytes) = match deposit_rx.recv_async().await {
Ok(p) => p,
Err(_) => return, };
let mut ctx = IngestCtx {
pending: &mut pending,
next_expected: &mut next_expected,
consumer_tx: &consumer_tx,
backpressure: &backpressure,
metrics: metrics.as_ref(),
anchor_id,
session_id,
};
if !ingest(seq_opt, bytes, &mut ctx).await {
return;
}
while let Ok((seq_opt, bytes)) = deposit_rx.try_recv() {
if !ingest(seq_opt, bytes, &mut ctx).await {
return;
}
}
}
}
struct IngestCtx<'a> {
pending: &'a mut BTreeMap<u64, Vec<u8>>,
next_expected: &'a mut u64,
consumer_tx: &'a flume::Sender<Vec<u8>>,
backpressure: &'a AtomicU64,
metrics: Option<&'a StreamingTransportMetricsHandle>,
anchor_id: u64,
session_id: u64,
}
async fn ingest(seq_opt: Option<u64>, bytes: Vec<u8>, ctx: &mut IngestCtx<'_>) -> bool {
match seq_opt {
None => forward(bytes, ctx).await,
Some(seq) => {
if seq < *ctx.next_expected {
return true;
}
if seq == *ctx.next_expected {
if !forward(bytes, ctx).await {
return false;
}
*ctx.next_expected += 1;
while let Some(b) = ctx.pending.remove(ctx.next_expected) {
if !forward(b, ctx).await {
return false;
}
*ctx.next_expected += 1;
}
return true;
}
if ctx.pending.len() >= REORDER_WINDOW && !ctx.pending.contains_key(&seq) {
tracing::error!(
anchor_id = ctx.anchor_id,
session_id = ctx.session_id,
seq,
next_expected = *ctx.next_expected,
window = REORDER_WINDOW,
"_stream_data: reorder window exceeded; closing stream"
);
return false;
}
ctx.pending.insert(seq, bytes);
true
}
}
}
async fn forward(bytes: Vec<u8>, ctx: &IngestCtx<'_>) -> bool {
match ctx.consumer_tx.try_send(bytes) {
Ok(()) => true,
Err(flume::TrySendError::Full(bytes)) => {
ctx.backpressure.fetch_add(1, Ordering::Relaxed);
if let Some(metrics) = ctx.metrics {
metrics.record_backpressure();
}
ctx.consumer_tx.send_async(bytes).await.is_ok()
}
Err(flume::TrySendError::Disconnected(_)) => false,
}
}
impl FrameTransport for VeloFrameTransport {
fn bind(
&self,
anchor_id: u64,
session_id: u64,
) -> BoxFuture<'_, Result<(String, flume::Receiver<Vec<u8>>)>> {
let worker_id = self.worker_id;
let dispatch = self.dispatch.clone();
let backpressure = self.backpressure_count.clone();
let metrics = self.streaming_metrics.clone();
Box::pin(async move {
let (consumer_tx, consumer_rx) = flume::bounded::<Vec<u8>>(256);
dispatch.retain(|&(aid, _), entry| aid != anchor_id || !entry.sender.is_disconnected());
let (deposit_tx, deposit_rx) = flume::unbounded::<Deposit>();
let token = next_dispatch_token();
dispatch.insert(
(anchor_id, session_id),
DispatchEntry {
token,
sender: deposit_tx,
},
);
tokio::spawn(run_deliverer(DelivererCtx {
deposit_rx,
token,
consumer_tx,
backpressure,
metrics,
dispatch: dispatch.clone(),
anchor_id,
session_id,
}));
let endpoint = format!("velo://{}/stream/{}", worker_id.as_u64(), anchor_id);
Ok((endpoint, consumer_rx))
})
}
fn connect(
&self,
endpoint: &str,
_anchor_id: u64,
session_id: u64,
) -> BoxFuture<'_, Result<flume::Sender<Vec<u8>>>> {
let endpoint = endpoint.to_string();
let messenger = self.messenger.clone();
Box::pin(async move {
let (target_worker_id, target_anchor_id) = parse_velo_uri(&endpoint)?;
let (tx, rx) = flume::bounded::<Vec<u8>>(256);
tokio::spawn(async move {
let mut seq: u64 = 0;
while let Ok(frame_bytes) = rx.recv_async().await {
let mut headers = HashMap::with_capacity(3);
headers.insert(ANCHOR_ID_HEADER.to_string(), target_anchor_id.to_string());
headers.insert(SESSION_ID_HEADER.to_string(), session_id.to_string());
headers.insert(STREAM_SEQ_HEADER.to_string(), seq.to_string());
if let Err(e) = messenger
.am_send_streaming("_stream_data")
.expect("am_send_streaming builder")
.headers(headers)
.raw_payload(bytes::Bytes::from(frame_bytes))
.worker(WorkerId::from_u64(target_worker_id))
.send()
.await
{
tracing::error!("_stream_data am_send failed: {}", e);
break;
}
seq += 1;
}
});
Ok(tx)
})
}
}
pub fn parse_velo_uri(uri: &str) -> Result<(u64, u64)> {
let stripped = uri
.strip_prefix("velo://")
.ok_or_else(|| anyhow::anyhow!("invalid velo URI: missing velo:// prefix: {}", uri))?;
let parts: Vec<&str> = stripped.split('/').collect();
if parts.len() != 3 || parts[1] != "stream" {
anyhow::bail!(
"invalid velo URI format: expected velo://{{worker_id}}/stream/{{anchor_id}}, got: {}",
uri
);
}
let worker_id: u64 = parts[0]
.parse()
.map_err(|_| anyhow::anyhow!("invalid worker_id in URI: {}", parts[0]))?;
let anchor_id: u64 = parts[2]
.parse()
.map_err(|_| anyhow::anyhow!("invalid anchor_id in URI: {}", parts[2]))?;
Ok((worker_id, anchor_id))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_velo_uri_valid() {
let (wid, aid) = parse_velo_uri("velo://123/stream/456").unwrap();
assert_eq!(wid, 123);
assert_eq!(aid, 456);
}
#[test]
fn test_parse_velo_uri_missing_prefix() {
assert!(parse_velo_uri("http://123/stream/456").is_err());
}
#[test]
fn test_parse_velo_uri_non_numeric_worker() {
assert!(parse_velo_uri("velo://abc/stream/456").is_err());
}
#[test]
fn test_parse_velo_uri_non_numeric_anchor() {
assert!(parse_velo_uri("velo://123/stream/xyz").is_err());
}
#[test]
fn test_parse_velo_uri_wrong_path_segment() {
assert!(parse_velo_uri("velo://123/wrong/456").is_err());
}
#[test]
fn test_parse_velo_uri_too_few_segments() {
assert!(parse_velo_uri("velo://123/stream").is_err());
}
#[test]
fn test_parse_velo_uri_too_many_segments() {
assert!(parse_velo_uri("velo://123/stream/456/extra").is_err());
}
fn spawn_deliverer() -> (
flume::Sender<Deposit>,
flume::Receiver<Vec<u8>>,
Arc<AtomicU64>,
) {
let (deposit_tx, deposit_rx) = flume::unbounded::<Deposit>();
let (consumer_tx, consumer_rx) = flume::bounded::<Vec<u8>>(64);
let backpressure = Arc::new(AtomicU64::new(0));
let dispatch: Arc<DispatchMap> = Arc::new(DashMap::new());
tokio::spawn(super::run_deliverer(super::DelivererCtx {
deposit_rx,
token: super::next_dispatch_token(),
consumer_tx,
backpressure: backpressure.clone(),
metrics: None,
dispatch,
anchor_id: 42,
session_id: 7,
}));
(deposit_tx, consumer_rx, backpressure)
}
#[tokio::test(flavor = "multi_thread")]
async fn deliverer_reorders_shuffled_seqs_in_order() {
let (tx, rx, _) = spawn_deliverer();
let order = [3u64, 0, 5, 4, 2, 1, 8, 7, 6, 11, 9, 10, 13, 12, 15, 14];
for s in order {
tx.send_async((Some(s), vec![s as u8])).await.unwrap();
}
for expected in 0u8..16 {
let bytes = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv_async())
.await
.expect("recv timeout")
.expect("deliverer closed");
assert_eq!(bytes, vec![expected], "frames out of order");
}
}
#[tokio::test(flavor = "multi_thread")]
async fn deliverer_accepts_head_of_line_at_full_window() {
let (tx, rx, _) = spawn_deliverer();
for s in 1u64..=(REORDER_WINDOW as u64) {
tx.send_async((Some(s), vec![(s & 0xff) as u8]))
.await
.unwrap();
}
tx.send_async((Some(0), vec![0])).await.unwrap();
for expected in 0u64..=(REORDER_WINDOW as u64) {
let bytes = tokio::time::timeout(std::time::Duration::from_secs(2), rx.recv_async())
.await
.expect("recv timeout — stream was wrongly closed")
.expect("deliverer closed prematurely");
assert_eq!(
bytes,
vec![(expected & 0xff) as u8],
"frame at seq {expected}"
);
}
}
#[tokio::test(flavor = "multi_thread")]
async fn deliverer_window_overflow_closes_stream() {
let (tx, rx, _) = spawn_deliverer();
for s in 1u64..=(REORDER_WINDOW as u64) {
tx.send_async((Some(s), vec![s as u8])).await.unwrap();
}
tx.send_async((Some(REORDER_WINDOW as u64 + 1), vec![0]))
.await
.unwrap();
let res = tokio::time::timeout(std::time::Duration::from_secs(2), rx.recv_async()).await;
match res {
Ok(Err(_)) => { }
Ok(Ok(b)) => panic!("expected closed channel, got frame: {b:?}"),
Err(_) => panic!("timed out waiting for channel close"),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn cleanup_does_not_evict_newer_binding_for_same_key() {
let key = (1u64, 2u64);
let dispatch: Arc<DispatchMap> = Arc::new(DashMap::new());
let (a_tx, a_rx) = flume::unbounded::<Deposit>();
let (a_consumer_tx, _a_consumer_rx) = flume::bounded::<Vec<u8>>(64);
let a_token = super::next_dispatch_token();
dispatch.insert(
key,
super::DispatchEntry {
token: a_token,
sender: a_tx.clone(),
},
);
let a_deliverer = tokio::spawn(super::run_deliverer(super::DelivererCtx {
deposit_rx: a_rx,
token: a_token,
consumer_tx: a_consumer_tx,
backpressure: Arc::new(AtomicU64::new(0)),
metrics: None,
dispatch: dispatch.clone(),
anchor_id: key.0,
session_id: key.1,
}));
let (b_tx, _b_rx) = flume::unbounded::<Deposit>();
let b_token = super::next_dispatch_token();
dispatch.insert(
key,
super::DispatchEntry {
token: b_token,
sender: b_tx.clone(),
},
);
drop(a_tx);
tokio::time::timeout(std::time::Duration::from_secs(5), a_deliverer)
.await
.expect("A deliverer did not exit after channel close (5s timeout)")
.expect("A deliverer panicked");
let current = dispatch.get(&key).expect("B entry was clobbered");
assert_eq!(current.value().token, b_token);
assert!(current.value().sender.same_channel(&b_tx));
}
#[tokio::test(flavor = "multi_thread")]
async fn deliverer_forwards_legacy_no_seq_in_arrival_order() {
let (tx, rx, _) = spawn_deliverer();
for v in 0u8..8 {
tx.send_async((None, vec![v])).await.unwrap();
}
for expected in 0u8..8 {
let bytes = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv_async())
.await
.expect("recv timeout")
.expect("deliverer closed");
assert_eq!(bytes, vec![expected]);
}
}
}