use std::net::SocketAddr;
use std::sync::Arc;
use anyhow::Result;
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 crate::streaming::transport::FrameTransport;
pub(crate) mod proto {
tonic::include_proto!("velo.streaming.v1");
}
use proto::{
FramedData,
velo_streaming_client::VeloStreamingClient,
velo_streaming_server::{VeloStreaming, VeloStreamingServer},
};
fn is_terminal_sentinel(bytes: &[u8]) -> bool {
use crate::streaming::sender::{cached_detached, cached_dropped, cached_finalized};
if bytes == cached_dropped().as_slice()
|| bytes == cached_detached().as_slice()
|| bytes == cached_finalized().as_slice()
{
return true;
}
if let Ok(frame) = rmp_serde::from_slice::<crate::streaming::frame::StreamFrame<()>>(bytes) {
matches!(
frame,
crate::streaming::frame::StreamFrame::TransportError(_)
)
} else {
false
}
}
#[derive(Clone)]
struct GrpcStreamingService {
routing: Arc<DashMap<u64, flume::Sender<Vec<u8>>>>,
active: Arc<DashMap<u64, ()>>,
}
#[tonic::async_trait]
impl VeloStreaming for GrpcStreamingService {
type StreamStream = futures::stream::Empty<Result<FramedData, Status>>;
async fn stream(
&self,
request: Request<Streaming<FramedData>>,
) -> Result<Response<Self::StreamStream>, Status> {
let anchor_id: u64 = {
let meta = request
.metadata()
.get("x-anchor-id")
.ok_or_else(|| Status::invalid_argument("missing x-anchor-id metadata header"))?;
let s = meta
.to_str()
.map_err(|_| Status::invalid_argument("x-anchor-id metadata is not valid UTF-8"))?;
s.parse::<u64>()
.map_err(|_| Status::invalid_argument("x-anchor-id is not a valid u64"))?
};
let entry = self.active.entry(anchor_id);
match entry {
dashmap::Entry::Occupied(_) => {
return Err(Status::already_exists(format!(
"anchor_id {} is already attached",
anchor_id
)));
}
dashmap::Entry::Vacant(v) => {
v.insert(());
}
}
let frame_tx = match self.routing.get(&anchor_id) {
Some(tx) => tx.clone(),
None => {
self.active.remove(&anchor_id);
return Err(Status::not_found(format!(
"no routing slot for anchor_id {}",
anchor_id
)));
}
};
let active_ref = self.active.clone();
let mut stream = request.into_inner();
tokio::spawn(async move {
let mut last_was_terminal = false;
while let Some(result) = stream.next().await {
match result {
Ok(framed) => {
let payload = framed.payload;
last_was_terminal = is_terminal_sentinel(&payload);
if frame_tx.send_async(payload).await.is_err() {
active_ref.remove(&anchor_id);
return;
}
}
Err(e) => {
tracing::warn!("gRPC streaming recv error for anchor {}: {}", anchor_id, e);
break;
}
}
}
if !last_was_terminal {
let _ = frame_tx
.send_async(crate::streaming::sender::cached_dropped().clone())
.await;
}
active_ref.remove(&anchor_id);
});
Ok(Response::new(futures::stream::empty()))
}
}
pub struct GrpcFrameTransport {
routing: Arc<DashMap<u64, flume::Sender<Vec<u8>>>>,
#[allow(dead_code)]
active: Arc<DashMap<u64, ()>>,
bound_addr: SocketAddr,
cancel: CancellationToken,
}
impl GrpcFrameTransport {
pub async fn new(bind_addr: SocketAddr) -> Result<Self> {
let routing: Arc<DashMap<u64, flume::Sender<Vec<u8>>>> = Arc::new(DashMap::new());
let active: Arc<DashMap<u64, ()>> = Arc::new(DashMap::new());
let cancel = CancellationToken::new();
let listener = tokio::net::TcpListener::bind(bind_addr).await?;
let bound_addr = listener.local_addr()?;
let service = GrpcStreamingService {
routing: routing.clone(),
active: active.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(Self {
routing,
active,
bound_addr,
cancel,
})
}
pub async fn default_new() -> Result<Self> {
Self::new("0.0.0.0:0".parse().unwrap()).await
}
pub fn bound_addr(&self) -> SocketAddr {
self.bound_addr
}
}
impl Drop for GrpcFrameTransport {
fn drop(&mut self) {
self.cancel.cancel();
}
}
pub fn parse_grpc_endpoint(endpoint: &str) -> Result<(SocketAddr, u64)> {
let stripped = endpoint
.strip_prefix("grpc://")
.ok_or_else(|| anyhow::anyhow!("missing grpc:// prefix: {}", endpoint))?;
let slash_pos = stripped
.rfind('/')
.ok_or_else(|| anyhow::anyhow!("missing anchor_id in endpoint: {}", endpoint))?;
let addr_str = &stripped[..slash_pos];
let anchor_id_str = &stripped[slash_pos + 1..];
let addr: SocketAddr = addr_str
.parse()
.map_err(|e| anyhow::anyhow!("invalid address '{}': {}", addr_str, e))?;
let anchor_id: u64 = anchor_id_str
.parse()
.map_err(|e| anyhow::anyhow!("invalid anchor_id '{}': {}", anchor_id_str, e))?;
Ok((addr, anchor_id))
}
impl FrameTransport for GrpcFrameTransport {
fn bind(
&self,
anchor_id: u64,
_session_id: u64,
) -> BoxFuture<'_, Result<(String, flume::Receiver<Vec<u8>>)>> {
let addr = self.bound_addr;
let routing = self.routing.clone();
Box::pin(async move {
let (frame_tx, frame_rx) = flume::bounded::<Vec<u8>>(256);
routing.insert(anchor_id, frame_tx);
let advertise_addr = std::net::SocketAddr::new(
crate::streaming::util::resolve_advertise_ip(addr.ip()),
addr.port(),
);
let endpoint = format!("grpc://{}/{}", advertise_addr, anchor_id);
Ok((endpoint, frame_rx))
})
}
fn connect(
&self,
endpoint: &str,
_anchor_id: u64,
_session_id: u64,
) -> BoxFuture<'_, Result<flume::Sender<Vec<u8>>>> {
let endpoint = endpoint.to_string();
Box::pin(async move {
let (addr, anchor_id) = parse_grpc_endpoint(&endpoint)?;
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>>(256);
let request_stream = tokio_stream::wrappers::ReceiverStream::new(mpsc_rx);
let mut request = Request::new(request_stream);
request.metadata_mut().insert(
"x-anchor-id",
anchor_id
.to_string()
.parse()
.map_err(|_| anyhow::anyhow!("failed to encode anchor_id as metadata value"))?,
);
let _response = client
.stream(request)
.await
.map_err(|status| anyhow::anyhow!("gRPC stream rejected: {}", status))?;
tokio::spawn(async move {
while let Ok(payload) = frame_rx.recv_async().await {
let framed = FramedData {
preamble: vec![],
header: vec![],
payload,
};
if mpsc_tx.send(framed).await.is_err() {
break;
}
}
});
Ok(frame_tx)
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_grpc_endpoint_ipv6() {
let endpoint = "grpc://[::1]:50051/42";
let (addr, anchor_id) = parse_grpc_endpoint(endpoint).unwrap();
assert_eq!(addr, "[::1]:50051".parse::<SocketAddr>().unwrap());
assert_eq!(anchor_id, 42);
}
}