use std::sync::Arc;
use std::time::Duration;
use futures::StreamExt;
use velo::messenger::Messenger;
use velo::streaming::control::StreamCancelHandle;
use velo::streaming::mpsc::{
MpscAnchorAttachRequest, MpscAnchorAttachResponse, MpscAnchorCancelRequest,
MpscAnchorDetachRequest,
};
use velo::streaming::{
AnchorManager, AnchorManagerBuilder, AttachError, FrameTransport, MpscAnchorConfig, MpscFrame,
SenderId, StreamAnchorHandle, TcpFrameTransport,
};
use velo::transports::tcp::TcpTransportBuilder;
use velo_ext::{PeerInfo, WorkerId};
fn new_tcp_transport() -> Arc<velo::transports::tcp::TcpTransport> {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
Arc::new(
TcpTransportBuilder::new()
.from_listener(listener)
.unwrap()
.build()
.unwrap(),
)
}
async fn make_two_messengers() -> (Arc<Messenger>, Arc<Messenger>) {
let t1 = new_tcp_transport();
let t2 = new_tcp_transport();
let m1 = Messenger::builder()
.add_transport(t1)
.build()
.await
.expect("create messenger 1");
let m2 = Messenger::builder()
.add_transport(t2)
.build()
.await
.expect("create messenger 2");
let p1 = m1.peer_info();
let p2 = m2.peer_info();
m2.register_peer(p1).expect("register m1 on m2");
m1.register_peer(p2).expect("register m2 on m1");
tokio::time::sleep(Duration::from_millis(200)).await;
(m1, m2)
}
static STREAM_ENDPOINTS: std::sync::Mutex<Vec<(PeerInfo, Arc<TcpFrameTransport>)>> =
std::sync::Mutex::new(Vec::new());
async fn make_am(messenger: Arc<Messenger>) -> Arc<AnchorManager> {
let worker_id = messenger.instance_id().worker_id();
let stream = TcpFrameTransport::new(std::net::Ipv4Addr::LOCALHOST.into())
.await
.expect("bind streaming listener");
let peer_info = PeerInfo::new(messenger.instance_id(), stream.address());
{
let mut endpoints = STREAM_ENDPOINTS.lock().expect("endpoint registry poisoned");
for (other_peer, other_stream) in endpoints.iter() {
let _ = stream.register(other_peer);
let _ = other_stream.register(&peer_info);
}
endpoints.push((peer_info, Arc::clone(&stream)));
}
let am: Arc<AnchorManager> = Arc::new(
AnchorManagerBuilder::default()
.worker_id(worker_id)
.transport(Arc::clone(&stream) as Arc<dyn FrameTransport>)
.messenger(Some(Arc::clone(&messenger)))
.build()
.expect("AM build"),
);
am.register_handlers(Arc::clone(&messenger))
.expect("register_handlers");
am
}
fn roundtrip_handle(handle: StreamAnchorHandle) -> StreamAnchorHandle {
StreamAnchorHandle::from_u128(handle.as_u128())
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_multi_attach() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a).await;
let am_b = make_am(messenger_b).await;
let mut anchor = am_a.create_mpsc_anchor::<u32>();
let handle = anchor.handle();
let transferred = roundtrip_handle(handle);
let s1 = am_b
.attach_mpsc_stream_anchor::<u32>(transferred)
.await
.expect("remote attach s1");
let s2 = am_b
.attach_mpsc_stream_anchor::<u32>(transferred)
.await
.expect("remote attach s2");
assert_eq!(s1.sender_id(), SenderId(1));
assert_eq!(s2.sender_id(), SenderId(2));
for i in 0u32..5 {
s1.send(i).await.expect("s1 send");
s2.send(100 + i).await.expect("s2 send");
}
tokio::time::sleep(Duration::from_millis(50)).await;
let mut s1_count = 0;
let mut s2_count = 0;
let collect = async {
while s1_count < 5 || s2_count < 5 {
let Some(frame) = anchor.next().await else {
break;
};
match frame.expect("no stream error") {
(SenderId(1), MpscFrame::Item(_)) => s1_count += 1,
(SenderId(2), MpscFrame::Item(_)) => s2_count += 1,
(sid, MpscFrame::Item(_)) => panic!("unknown sender {sid}"),
(_, MpscFrame::SenderError(m)) => panic!("sender error: {m}"),
(_, MpscFrame::Detached | MpscFrame::Dropped(_)) => {}
}
}
};
tokio::time::timeout(Duration::from_secs(10), collect)
.await
.expect("consumer timed out");
assert_eq!(s1_count, 5);
assert_eq!(s2_count, 5);
let _ = s1.detach().await;
let _ = s2.detach().await;
anchor.cancel();
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_heartbeat_timeout_per_sender() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a).await;
let am_b = make_am(messenger_b).await;
let mut anchor = am_a.create_mpsc_anchor::<u32>();
let handle = anchor.handle();
let transferred = roundtrip_handle(handle);
let s1 = am_b
.attach_mpsc_stream_anchor::<u32>(transferred)
.await
.expect("attach s1");
let s2 = am_b
.attach_mpsc_stream_anchor::<u32>(transferred)
.await
.expect("attach s2");
s1.send(1).await.unwrap();
s2.send(2).await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
drop(s1);
s2.send(3).await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
let mut items_s1 = Vec::new();
let mut items_s2 = Vec::new();
let mut dropped_s1 = false;
let collect = async {
while let Some(frame) = anchor.next().await {
match frame.expect("stream error") {
(SenderId(1), MpscFrame::Item(v)) => items_s1.push(v),
(SenderId(2), MpscFrame::Item(v)) => items_s2.push(v),
(SenderId(1), MpscFrame::Dropped(_)) => {
dropped_s1 = true;
}
(SenderId(2), MpscFrame::Dropped(_) | MpscFrame::Detached) => break,
_ => {}
}
if dropped_s1 && items_s2.len() >= 2 {
break;
}
}
};
tokio::time::timeout(Duration::from_secs(10), collect)
.await
.expect("consumer timed out");
assert_eq!(items_s1, vec![1u32]);
assert!(items_s2.contains(&2u32));
assert!(items_s2.contains(&3u32));
assert!(dropped_s1, "must see Dropped for sender 1");
let _ = s2.detach().await;
anchor.cancel();
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_two_workers_no_routing_collision() {
let t_a = new_tcp_transport();
let t_b = new_tcp_transport();
let t_c = new_tcp_transport();
let messenger_a = Messenger::builder()
.add_transport(t_a)
.build()
.await
.unwrap();
let messenger_b = Messenger::builder()
.add_transport(t_b)
.build()
.await
.unwrap();
let messenger_c = Messenger::builder()
.add_transport(t_c)
.build()
.await
.unwrap();
let p_a = messenger_a.peer_info();
let p_b = messenger_b.peer_info();
let p_c = messenger_c.peer_info();
messenger_a.register_peer(p_b.clone()).unwrap();
messenger_a.register_peer(p_c.clone()).unwrap();
messenger_b.register_peer(p_a.clone()).unwrap();
messenger_c.register_peer(p_a.clone()).unwrap();
tokio::time::sleep(Duration::from_millis(200)).await;
let am_a = make_am(messenger_a).await;
let am_b = make_am(messenger_b).await;
let am_c = make_am(messenger_c).await;
let mut anchor = am_a.create_mpsc_anchor::<u32>();
let handle = anchor.handle();
let transferred = roundtrip_handle(handle);
let s_b = am_b
.attach_mpsc_stream_anchor::<u32>(transferred)
.await
.expect("worker B attach");
let s_c = am_c
.attach_mpsc_stream_anchor::<u32>(transferred)
.await
.expect("worker C attach");
const N: u32 = 50;
for i in 0..N {
s_b.send(i).await.expect("worker B send");
s_c.send(1000 + i).await.expect("worker C send");
}
let mut b_items: Vec<u32> = Vec::new();
let mut c_items: Vec<u32> = Vec::new();
let collect = async {
while b_items.len() < N as usize || c_items.len() < N as usize {
let Some(frame) = anchor.next().await else {
break;
};
match frame.expect("no stream error") {
(sid, MpscFrame::Item(v)) => {
if v < 1000 {
b_items.push(v);
assert_eq!(sid, SenderId(1), "B's items must all carry SenderId(1)");
} else {
c_items.push(v);
assert_eq!(sid, SenderId(2), "C's items must all carry SenderId(2)");
}
}
(_, MpscFrame::SenderError(m)) => panic!("sender error: {m}"),
(_, MpscFrame::Detached | MpscFrame::Dropped(_)) => {}
}
}
};
tokio::time::timeout(Duration::from_secs(10), collect)
.await
.expect("consumer timed out -- routing collision likely lost frames");
assert_eq!(b_items.len(), N as usize, "all B items must arrive");
assert_eq!(c_items.len(), N as usize, "all C items must arrive");
b_items.sort_unstable();
c_items.sort_unstable();
assert_eq!(b_items, (0..N).collect::<Vec<_>>());
assert_eq!(c_items, (1000..1000 + N).collect::<Vec<_>>());
let _ = s_b.detach().await;
let _ = s_c.detach().await;
anchor.cancel();
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_detach_is_exact_once() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a).await;
let am_b = make_am(messenger_b).await;
let config = MpscAnchorConfig {
max_senders: Some(1),
..Default::default()
};
let mut anchor = am_a.create_mpsc_anchor_with_config::<u32>(config);
let handle = roundtrip_handle(anchor.handle());
let sender = am_b.attach_mpsc_stream_anchor::<u32>(handle).await.unwrap();
assert_eq!(sender.sender_id(), SenderId(1));
let returned = sender.detach().await.expect("detach");
assert_eq!(returned, handle);
match tokio::time::timeout(Duration::from_secs(2), anchor.next())
.await
.expect("detach frame timeout")
{
Some(Ok((SenderId(1), MpscFrame::Detached))) => {}
other => panic!(
"expected exactly one Detached for sender 1, got {:?}",
other
),
}
let s2 = am_b.attach_mpsc_stream_anchor::<u32>(handle).await.unwrap();
assert_eq!(s2.sender_id(), SenderId(2));
s2.send(42).await.unwrap();
match tokio::time::timeout(Duration::from_secs(2), anchor.next())
.await
.expect("item frame timeout")
{
Some(Ok((SenderId(2), MpscFrame::Item(42)))) => {}
other => panic!("expected SenderId(2) Item(42), got {:?}", other),
}
match tokio::time::timeout(Duration::from_millis(200), anchor.next()).await {
Err(_) => {}
other => panic!(
"unexpected extra frame after clean remote detach: {:?}",
other
),
}
let _ = s2.detach().await;
anchor.cancel();
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_drop_is_exact_once() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a).await;
let am_b = make_am(messenger_b).await;
let config = MpscAnchorConfig {
max_senders: Some(1),
..Default::default()
};
let mut anchor = am_a.create_mpsc_anchor_with_config::<u32>(config);
let handle = roundtrip_handle(anchor.handle());
let sender = am_b.attach_mpsc_stream_anchor::<u32>(handle).await.unwrap();
assert_eq!(sender.sender_id(), SenderId(1));
drop(sender);
match tokio::time::timeout(Duration::from_secs(2), anchor.next())
.await
.expect("dropped frame timeout")
{
Some(Ok((SenderId(1), MpscFrame::Dropped(None)))) => {}
other => panic!(
"expected exactly one Dropped(None) for sender 1, got {:?}",
other
),
}
let s2 = am_b.attach_mpsc_stream_anchor::<u32>(handle).await.unwrap();
assert_eq!(s2.sender_id(), SenderId(2));
s2.send(99).await.unwrap();
match tokio::time::timeout(Duration::from_secs(2), anchor.next())
.await
.expect("item frame timeout")
{
Some(Ok((SenderId(2), MpscFrame::Item(99)))) => {}
other => panic!("expected SenderId(2) Item(99), got {:?}", other),
}
match tokio::time::timeout(Duration::from_millis(200), anchor.next()).await {
Err(_) => {}
other => panic!("unexpected extra frame after remote drop: {:?}", other),
}
let _ = s2.detach().await;
anchor.cancel();
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_handler_rejects_spsc_handle() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a.clone()).await;
let _am_b = make_am(messenger_b.clone()).await;
let spsc_anchor = am_a.create_anchor::<u32>();
let spsc_handle = StreamAnchorHandle::from_u128(spsc_anchor.handle().as_u128());
let worker_a = spsc_handle.unpack().0;
let req = MpscAnchorAttachRequest {
handle: spsc_handle,
session_id: 9999,
stream_cancel_handle: StreamCancelHandle::pack(WorkerId::from_u64(99), 9999),
supported_transport_keys: Vec::new(),
};
let response: MpscAnchorAttachResponse = messenger_b
.typed_unary_streaming::<MpscAnchorAttachResponse>("_mpsc_anchor_attach")
.payload(&req)
.expect("payload build")
.worker(worker_a)
.send()
.await
.expect("AM send");
match response {
MpscAnchorAttachResponse::Err { reason } => {
assert!(
reason.contains("spsc"),
"rejection must mention kind, got: {reason}"
);
}
MpscAnchorAttachResponse::Ok { .. } => panic!("handler must reject SPSC handle"),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_handler_anchor_not_found() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let _am_a = make_am(messenger_a.clone()).await;
let _am_b = make_am(messenger_b.clone()).await;
let worker_a = messenger_a.instance_id().worker_id();
let fake = StreamAnchorHandle::pack_mpsc(worker_a, 0xDEAD_BEEF);
let req = MpscAnchorAttachRequest {
handle: fake,
session_id: 1,
stream_cancel_handle: StreamCancelHandle::pack(WorkerId::from_u64(99), 1),
supported_transport_keys: Vec::new(),
};
let response: MpscAnchorAttachResponse = messenger_b
.typed_unary_streaming::<MpscAnchorAttachResponse>("_mpsc_anchor_attach")
.payload(&req)
.expect("payload build")
.worker(worker_a)
.send()
.await
.expect("AM send");
match response {
MpscAnchorAttachResponse::Err { reason } => {
assert!(reason.contains("not found"), "reason: {reason}");
}
MpscAnchorAttachResponse::Ok { .. } => panic!("handler must reject unknown anchor"),
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_max_senders_enforced_by_handler() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a.clone()).await;
let am_b = make_am(messenger_b.clone()).await;
let config = MpscAnchorConfig {
max_senders: Some(1),
..Default::default()
};
let anchor = am_a.create_mpsc_anchor_with_config::<u32>(config);
let handle = roundtrip_handle(anchor.handle());
let s1 = am_b
.attach_mpsc_stream_anchor::<u32>(handle)
.await
.expect("first attach");
let result = am_b.attach_mpsc_stream_anchor::<u32>(handle).await;
match result {
Err(AttachError::TransportError(e)) => {
assert!(
e.to_string().contains("max_senders"),
"error must mention max_senders, got: {e}"
);
}
other => panic!("expected TransportError(max_senders…), got {other:?}"),
}
drop(s1);
drop(anchor);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_detach_handler() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a.clone()).await;
let am_b = make_am(messenger_b.clone()).await;
let mut anchor = am_a.create_mpsc_anchor::<u32>();
let handle = roundtrip_handle(anchor.handle());
let sender = am_b
.attach_mpsc_stream_anchor::<u32>(handle)
.await
.expect("attach");
assert_eq!(sender.sender_id(), SenderId(1));
sender.send(7).await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
let req = MpscAnchorDetachRequest {
handle,
sender_id: 1,
};
let _: () = messenger_b
.typed_unary_streaming::<()>("_mpsc_anchor_detach")
.payload(&req)
.expect("payload build")
.worker(messenger_a.instance_id().worker_id())
.send()
.await
.expect("AM send");
tokio::time::sleep(Duration::from_millis(50)).await;
let s2 = am_b.attach_mpsc_stream_anchor::<u32>(handle).await.unwrap();
assert_eq!(s2.sender_id(), SenderId(2));
drop(s2);
while tokio::time::timeout(Duration::from_millis(50), anchor.next())
.await
.is_ok()
{}
drop(sender);
anchor.cancel();
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_remote_cancel_handler() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a.clone()).await;
let _am_b = make_am(messenger_b.clone()).await;
let anchor = am_a.create_mpsc_anchor::<u32>();
let handle = roundtrip_handle(anchor.handle());
let req = MpscAnchorCancelRequest { handle };
let _: () = messenger_b
.typed_unary_streaming::<()>("_mpsc_anchor_cancel")
.payload(&req)
.expect("payload build")
.worker(messenger_a.instance_id().worker_id())
.send()
.await
.expect("AM send");
tokio::time::sleep(Duration::from_millis(50)).await;
let result = am_a.attach_mpsc_stream_anchor::<u32>(handle).await;
assert!(
matches!(result, Err(AttachError::AnchorNotFound { .. })),
"anchor must be removed after cancel AM, got {result:?}"
);
drop(anchor);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_controller_cancel_propagates_cross_worker() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a.clone()).await;
let am_b = make_am(messenger_b.clone()).await;
let anchor = am_a.create_mpsc_anchor::<u32>();
let handle = roundtrip_handle(anchor.handle());
let controller = anchor.controller();
let sender = am_b.attach_mpsc_stream_anchor::<u32>(handle).await.unwrap();
sender.send(1).await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
controller.cancel();
tokio::time::sleep(Duration::from_millis(200)).await;
let result = sender.send(2).await;
assert!(
matches!(result, Err(velo::streaming::SendError::ChannelClosed)),
"remote sender must see ChannelClosed after cancel, got {result:?}"
);
drop(sender);
drop(anchor);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_mpsc_pending_next_wakes_on_cross_worker_cancel() {
let (messenger_a, messenger_b) = make_two_messengers().await;
let am_a = make_am(messenger_a.clone()).await;
let am_b = make_am(messenger_b.clone()).await;
let anchor = am_a.create_mpsc_anchor::<u32>();
let handle = roundtrip_handle(anchor.handle());
let controller = anchor.controller();
let sender = am_b.attach_mpsc_stream_anchor::<u32>(handle).await.unwrap();
let next_task = tokio::spawn(async move {
let mut anchor = anchor;
anchor.next().await
});
tokio::time::sleep(Duration::from_millis(20)).await;
controller.cancel();
let result = tokio::time::timeout(Duration::from_secs(2), next_task)
.await
.expect("pending next() must wake after cross-worker cancel")
.expect("next task join");
assert!(
result.is_none(),
"cancelled cross-worker anchor.next() must resolve to None, got {result:?}"
);
drop(sender);
}