use std::net::SocketAddr;
use std::sync::Arc;
use anyhow::{Result, anyhow};
use dashmap::DashMap;
use futures::StreamExt;
use futures::future::BoxFuture;
use tokio_stream::wrappers::TcpListenerStream;
use tokio_util::sync::CancellationToken;
use tonic::{Request, Response, Status, Streaming};
use velo_ext::{PeerInfo, TransportKey, WorkerAddress, WorkerId};
use crate::streaming::transport::FrameTransport;
use crate::transports::address::WorkerAddressBuilder;
use crate::transports::utils::interfaces::{
InterfaceEndpoint, InterfaceFilter, parse_endpoints, resolve_advertise_endpoints,
select_best_endpoint,
};
pub(crate) mod proto {
tonic::include_proto!("velo.streaming.v1");
}
use proto::{
FramedData,
velo_streaming_client::VeloStreamingClient,
velo_streaming_server::{VeloStreaming, VeloStreamingServer},
};
pub const GRPC_STREAM_KEY: &str = "grpc-stream";
const ANCHOR_ID_META: &str = "x-anchor-id";
const SESSION_ID_META: &str = "x-session-id";
const TERMINAL_ACK_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(20);
use crate::streaming::sender::is_terminal_sentinel;
type SessionRouting = DashMap<(u64, u64), flume::Sender<Vec<u8>>>;
#[derive(Clone)]
struct GrpcStreamingService {
routing: Arc<SessionRouting>,
metrics: Arc<std::sync::OnceLock<Arc<crate::observability::VeloMetrics>>>,
}
#[tonic::async_trait]
impl VeloStreaming for GrpcStreamingService {
type StreamStream =
std::pin::Pin<Box<dyn futures::Stream<Item = Result<FramedData, Status>> + Send + 'static>>;
async fn stream(
&self,
request: Request<Streaming<FramedData>>,
) -> Result<Response<Self::StreamStream>, Status> {
let anchor_id = read_u64_meta(&request, ANCHOR_ID_META).map_err(|boxed| *boxed)?;
let session_id = read_u64_meta(&request, SESSION_ID_META).map_err(|boxed| *boxed)?;
let frame_tx = match self.routing.remove(&(anchor_id, session_id)) {
Some((_, tx)) => tx,
None => {
return Err(Status::not_found(format!(
"no routing slot for (anchor_id={}, session_id={})",
anchor_id, session_id
)));
}
};
let mut stream = request.into_inner();
let metrics = self.metrics.get().cloned();
let (done_tx, done_rx) = tokio::sync::oneshot::channel::<()>();
tokio::spawn(async move {
let mut last_was_terminal = false;
let mut frames_seen = 0u64;
let mut end_reason = "end-of-stream";
while let Some(result) = stream.next().await {
match result {
Ok(framed) => {
frames_seen += 1;
let payload = framed.payload;
last_was_terminal = is_terminal_sentinel(&payload);
match frame_tx.try_send(payload) {
Ok(()) => {}
Err(flume::TrySendError::Full(b)) => {
if let Some(m) = metrics.as_ref() {
m.record_server_pump_backpressure();
}
if frame_tx.send_async(b).await.is_err() {
return; }
}
Err(flume::TrySendError::Disconnected(_)) => {
return;
}
}
}
Err(e) => {
end_reason = "recv-error";
tracing::warn!(
"gRPC streaming recv error for anchor={} session={}: {}",
anchor_id,
session_id,
e
);
break;
}
}
}
if !last_was_terminal {
tracing::warn!(
anchor_id,
session_id,
frames_seen,
end_reason,
"gRPC server pump: last frame was not a terminal sentinel, injecting Dropped"
);
let _ = frame_tx
.send_async(crate::streaming::sender::cached_dropped().clone())
.await;
}
let _ = done_tx.send(());
});
let response_stream = futures::stream::once(async move {
let _ = done_rx.await;
})
.filter_map(|()| std::future::ready(None::<Result<FramedData, Status>>));
Ok(Response::new(Box::pin(response_stream)))
}
}
fn read_u64_meta<T>(request: &Request<T>, name: &str) -> Result<u64, Box<Status>> {
let meta = request.metadata().get(name).ok_or_else(|| {
Box::new(Status::invalid_argument(format!(
"missing {} metadata header",
name
)))
})?;
let s = meta.to_str().map_err(|_| {
Box::new(Status::invalid_argument(format!(
"{} metadata is not valid UTF-8",
name
)))
})?;
s.parse::<u64>().map_err(|_| {
Box::new(Status::invalid_argument(format!(
"{} is not a valid u64",
name
)))
})
}
pub struct GrpcFrameTransport {
key: TransportKey,
bind_addr: SocketAddr,
local_address: WorkerAddress,
local_interfaces: std::sync::OnceLock<Vec<InterfaceEndpoint>>,
interface_filter: InterfaceFilter,
numa_hint: Option<u32>,
peers: Arc<DashMap<WorkerId, SocketAddr>>,
routing: Arc<SessionRouting>,
cancel: CancellationToken,
metrics: Arc<std::sync::OnceLock<Arc<crate::observability::VeloMetrics>>>,
}
impl GrpcFrameTransport {
pub async fn with_config(
bind_addr: SocketAddr,
key: TransportKey,
interface_filter: InterfaceFilter,
numa_hint: Option<u32>,
) -> Result<Arc<Self>> {
let routing: Arc<SessionRouting> = Arc::new(DashMap::new());
let cancel = CancellationToken::new();
let metrics: Arc<std::sync::OnceLock<Arc<crate::observability::VeloMetrics>>> =
Arc::new(std::sync::OnceLock::new());
let listener = tokio::net::TcpListener::bind(bind_addr).await?;
let bound_addr = listener.local_addr()?;
let endpoints = resolve_advertise_endpoints(bound_addr, &interface_filter)?;
let encoded = rmp_serde::to_vec(&endpoints)
.map_err(|e| anyhow!("Failed to encode interface endpoints: {e}"))?;
let mut addr_builder = WorkerAddressBuilder::new();
addr_builder
.add_entry(key.as_str(), encoded)
.map_err(|e| anyhow!("Failed to build WorkerAddress entry: {e}"))?;
let local_address = addr_builder
.build()
.map_err(|e| anyhow!("Failed to build WorkerAddress: {e}"))?;
let service = GrpcStreamingService {
routing: routing.clone(),
metrics: metrics.clone(),
};
let cancel_clone = cancel.clone();
tokio::spawn(async move {
let server =
tonic::transport::Server::builder().add_service(VeloStreamingServer::new(service));
if let Err(e) = server
.serve_with_incoming_shutdown(
TcpListenerStream::new(listener),
cancel_clone.cancelled(),
)
.await
{
tracing::warn!("GrpcFrameTransport server error: {}", e);
}
});
Ok(Arc::new(Self {
key,
bind_addr: bound_addr,
local_address,
local_interfaces: std::sync::OnceLock::new(),
interface_filter,
numa_hint,
peers: Arc::new(DashMap::new()),
routing,
cancel,
metrics,
}))
}
pub(crate) fn set_metrics(&self, metrics: Arc<crate::observability::VeloMetrics>) {
let _ = self.metrics.set(metrics);
}
pub async fn new(bind_addr: SocketAddr) -> Result<Arc<Self>> {
Self::with_config(
bind_addr,
TransportKey::new(GRPC_STREAM_KEY),
InterfaceFilter::All,
None,
)
.await
}
pub async fn default_new() -> Result<Arc<Self>> {
Self::new("0.0.0.0:0".parse().unwrap()).await
}
pub fn bound_addr(&self) -> SocketAddr {
self.bind_addr
}
}
impl Drop for GrpcFrameTransport {
fn drop(&mut self) {
self.cancel.cancel();
}
}
impl FrameTransport for GrpcFrameTransport {
fn key(&self) -> TransportKey {
self.key.clone()
}
fn address(&self) -> WorkerAddress {
self.local_address.clone()
}
fn register(&self, peer_info: &PeerInfo) -> Result<()> {
let raw = peer_info
.worker_address()
.get_entry(self.key.as_str())
.map_err(|e| anyhow!("decoding peer WorkerAddress: {e}"))?
.ok_or_else(|| {
anyhow!(
"peer {} has no '{}' streaming endpoint entry",
peer_info.worker_id(),
self.key
)
})?;
let remote_endpoints =
parse_endpoints(&raw).map_err(|e| anyhow!("Failed to parse gRPC endpoints: {e}"))?;
let local = self.local_interfaces.get_or_init(|| {
resolve_advertise_endpoints(self.bind_addr, &self.interface_filter).unwrap_or_default()
});
let addr =
select_best_endpoint(&remote_endpoints, local, self.numa_hint).ok_or_else(|| {
anyhow!(
"no suitable endpoint for peer {} from {:?}",
peer_info.worker_id(),
remote_endpoints
)
})?;
self.peers.insert(peer_info.worker_id(), addr);
Ok(())
}
fn bind(
&self,
anchor_id: u64,
session_id: u64,
) -> BoxFuture<'_, Result<flume::Receiver<Vec<u8>>>> {
let routing = self.routing.clone();
Box::pin(async move {
let (frame_tx, frame_rx) = flume::bounded::<Vec<u8>>(4096);
if routing.insert((anchor_id, session_id), frame_tx).is_some() {
tracing::warn!(
anchor_id,
session_id,
"GrpcFrameTransport::bind overwrote an existing routing entry; \
previous frame_tx dropped (consumer will see channel close)"
);
}
Ok(frame_rx)
})
}
fn connect(
&self,
peer: WorkerId,
anchor_id: u64,
session_id: u64,
) -> BoxFuture<'_, Result<flume::Sender<Vec<u8>>>> {
let peers = self.peers.clone();
Box::pin(async move {
let addr = *peers.get(&peer).ok_or_else(|| {
anyhow!(
"gRPC streaming: peer {} not registered (call register_peer first)",
peer
)
})?;
let channel = tonic::transport::Channel::from_shared(format!("http://{}", addr))?
.connect()
.await?;
let mut client = VeloStreamingClient::new(channel);
let (mpsc_tx, mpsc_rx) = tokio::sync::mpsc::channel::<FramedData>(256);
let (frame_tx, frame_rx) = flume::bounded::<Vec<u8>>(4096);
let request_stream = tokio_stream::wrappers::ReceiverStream::new(mpsc_rx);
let mut request = Request::new(request_stream);
request.metadata_mut().insert(
ANCHOR_ID_META,
anchor_id
.to_string()
.parse()
.map_err(|_| anyhow!("failed to encode anchor_id as metadata"))?,
);
request.metadata_mut().insert(
SESSION_ID_META,
session_id
.to_string()
.parse()
.map_err(|_| anyhow!("failed to encode session_id as metadata"))?,
);
let response = client
.stream(request)
.await
.map_err(|status| anyhow!("gRPC stream rejected: {}", status))?;
tokio::spawn(async move {
let mut inbound = response.into_inner();
while let Ok(payload) = frame_rx.recv_async().await {
let is_terminal = is_terminal_sentinel(&payload);
let framed = FramedData {
preamble: vec![],
header: vec![],
payload,
};
if mpsc_tx.send(framed).await.is_err() {
break;
}
if is_terminal {
break;
}
}
drop(mpsc_tx);
let drain = async {
while let Some(next) = inbound.next().await {
if let Err(status) = next {
tracing::debug!(
anchor_id,
session_id,
%status,
"gRPC streaming: error draining response after terminal sentinel"
);
break;
}
}
};
if tokio::time::timeout(TERMINAL_ACK_TIMEOUT, drain)
.await
.is_err()
{
tracing::warn!(
anchor_id,
session_id,
timeout_ms = TERMINAL_ACK_TIMEOUT.as_millis() as u64,
"gRPC streaming: timed out waiting for the server to acknowledge the \
terminal sentinel"
);
}
});
Ok(frame_tx)
})
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use super::*;
use velo_ext::InstanceId;
fn fresh_peer(address: WorkerAddress) -> (WorkerId, PeerInfo) {
let inst = InstanceId::new_v4();
let wid = inst.worker_id();
(wid, PeerInfo::new(inst, address))
}
#[tokio::test(flavor = "multi_thread")]
async fn response_completes_only_after_request_stream_drains() {
let server = GrpcFrameTransport::default_new().await.unwrap();
let rx = server.bind(9, 4).await.unwrap();
let addr = SocketAddr::new(
std::net::Ipv4Addr::LOCALHOST.into(),
server.bound_addr().port(),
);
let channel = tonic::transport::Channel::from_shared(format!("http://{addr}"))
.unwrap()
.connect()
.await
.unwrap();
let mut client = VeloStreamingClient::new(channel);
let (tx, req_rx) = tokio::sync::mpsc::channel::<FramedData>(8);
let mut request = Request::new(tokio_stream::wrappers::ReceiverStream::new(req_rx));
request
.metadata_mut()
.insert(ANCHOR_ID_META, "9".parse().unwrap());
request
.metadata_mut()
.insert(SESSION_ID_META, "4".parse().unwrap());
let mut inbound = client.stream(request).await.unwrap().into_inner();
let framed = |payload: Vec<u8>| FramedData {
preamble: vec![],
header: vec![],
payload,
};
tx.send(framed(b"frame".to_vec())).await.unwrap();
assert_eq!(rx.recv_async().await.unwrap(), b"frame".to_vec());
assert!(
tokio::time::timeout(Duration::from_millis(250), inbound.next())
.await
.is_err(),
"response completed while the request stream was still open"
);
tx.send(framed(crate::streaming::sender::cached_finalized().clone()))
.await
.unwrap();
drop(tx);
let drained = tokio::time::timeout(Duration::from_secs(5), async {
while let Some(item) = inbound.next().await {
item.expect("response stream must end cleanly");
}
})
.await;
assert!(
drained.is_ok(),
"response never completed after the request stream ended"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn terminal_sentinel_survives_rapid_session_cycles() {
const CYCLES: u64 = 96;
const FRAMES: u64 = 8;
let server = GrpcFrameTransport::default_new().await.unwrap();
let client = GrpcFrameTransport::default_new().await.unwrap();
let (server_worker, server_peer) = fresh_peer(server.address());
client.register(&server_peer).unwrap();
let finalized = crate::streaming::sender::cached_finalized().clone();
let dropped = crate::streaming::sender::cached_dropped().clone();
for cycle in 1..=CYCLES {
let rx = server.bind(cycle, cycle).await.unwrap();
let tx = client.connect(server_worker, cycle, cycle).await.unwrap();
let fin = finalized.clone();
let producer = tokio::spawn(async move {
for i in 0..FRAMES {
tx.send_async(i.to_be_bytes().to_vec()).await.unwrap();
}
tx.send_async(fin).await.unwrap();
});
let mut items = 0u64;
let mut saw_finalized = false;
while let Ok(frame) = tokio::time::timeout(Duration::from_secs(10), rx.recv_async())
.await
.unwrap_or_else(|_| panic!("cycle {cycle}: timed out after {items} frames"))
{
assert_ne!(
frame, dropped,
"cycle {cycle}: server injected Dropped after {items} frames"
);
if frame == finalized {
saw_finalized = true;
break;
}
items += 1;
}
producer.await.unwrap();
assert!(
saw_finalized,
"cycle {cycle}: channel closed after {items} frames without Finalized"
);
assert_eq!(items, FRAMES, "cycle {cycle}: wrong frame count");
}
}
#[tokio::test(flavor = "multi_thread")]
async fn round_trip_via_register_and_connect() {
let server = GrpcFrameTransport::default_new().await.unwrap();
let client = GrpcFrameTransport::default_new().await.unwrap();
let (server_worker, server_peer) = fresh_peer(server.address());
client.register(&server_peer).unwrap();
let rx = server.bind(7, 1).await.unwrap();
let tx = client.connect(server_worker, 7, 1).await.unwrap();
let payload = b"hello".to_vec();
tx.send_async(payload.clone()).await.unwrap();
let received = tokio::time::timeout(std::time::Duration::from_secs(5), rx.recv_async())
.await
.expect("recv timeout")
.expect("channel closed");
assert_eq!(received, payload);
}
#[tokio::test(flavor = "multi_thread")]
async fn bind_overwrite_emits_warn_and_returns_new_rx() {
let server = GrpcFrameTransport::default_new().await.unwrap();
let rx1 = server.bind(42, 7).await.unwrap();
let _rx2 = server.bind(42, 7).await.unwrap();
let res = tokio::time::timeout(std::time::Duration::from_millis(200), rx1.recv_async())
.await
.expect("rx1 should have closed promptly");
assert!(res.is_err(), "rx1 must observe channel closed");
}
#[test]
fn read_u64_meta_missing_header_returns_invalid_argument() {
let req: Request<()> = Request::new(());
let err = read_u64_meta(&req, "x-missing").unwrap_err();
assert_eq!(err.code(), tonic::Code::InvalidArgument);
assert!(
err.message().contains("missing x-missing"),
"got: {}",
err.message()
);
}
#[test]
fn read_u64_meta_non_u64_value_returns_invalid_argument() {
let mut req: Request<()> = Request::new(());
req.metadata_mut()
.insert("x-bad", "not-a-number".parse().unwrap());
let err = read_u64_meta(&req, "x-bad").unwrap_err();
assert_eq!(err.code(), tonic::Code::InvalidArgument);
assert!(
err.message().contains("not a valid u64"),
"got: {}",
err.message()
);
}
#[tokio::test(flavor = "multi_thread")]
async fn register_rejects_peer_without_grpc_entry() {
let transport = GrpcFrameTransport::default_new().await.unwrap();
let (_, peer) = fresh_peer(WorkerAddress::empty());
let err = transport.register(&peer).unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("has no") && msg.contains("streaming endpoint entry"),
"got: {msg}"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn connect_to_unregistered_peer_errors() {
let transport = GrpcFrameTransport::default_new().await.unwrap();
let bogus = InstanceId::new_v4().worker_id();
let err = transport.connect(bogus, 1, 1).await.unwrap_err();
let msg = format!("{err}");
assert!(msg.contains("not registered"), "got: {msg}");
}
#[tokio::test(flavor = "multi_thread")]
async fn grpc_client_pump_breaks_after_terminal_sentinel() {
let server = GrpcFrameTransport::default_new().await.unwrap();
let client = GrpcFrameTransport::default_new().await.unwrap();
let (server_worker, server_peer) = fresh_peer(server.address());
client.register(&server_peer).unwrap();
let rx = server.bind(11, 3).await.unwrap();
let tx = client.connect(server_worker, 11, 3).await.unwrap();
let finalized =
rmp_serde::to_vec(&crate::streaming::frame::StreamFrame::<()>::Finalized).unwrap();
tx.send_async(finalized.clone()).await.unwrap();
let heartbeat =
rmp_serde::to_vec(&crate::streaming::frame::StreamFrame::<()>::Heartbeat).unwrap();
let _ = tx.send_async(heartbeat.clone()).await;
drop(tx);
let received = tokio::time::timeout(std::time::Duration::from_secs(5), rx.recv_async())
.await
.expect("recv timeout")
.expect("channel closed");
assert_eq!(
received, finalized,
"first frame on the server side must be the Finalized sentinel"
);
let next =
tokio::time::timeout(std::time::Duration::from_millis(500), rx.recv_async()).await;
if let Ok(Ok(extra)) = next {
assert_ne!(
extra, heartbeat,
"heartbeat queued behind Finalized must not reach the server"
);
assert_ne!(
extra.as_slice(),
crate::streaming::sender::cached_dropped().as_slice(),
"server must not inject Dropped after Finalized"
);
}
}
}