use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::Duration;
use crate::transports::MessageType;
use crate::transports::tcp::TcpFrameCodec;
use anyhow::Result;
use dashmap::DashMap;
use futures::StreamExt;
use futures::future::BoxFuture;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio_util::codec::Framed;
use tokio_util::sync::CancellationToken;
use uuid::Uuid;
use crate::streaming::transport::FrameTransport;
const ACCEPT_TIMEOUT: Duration = Duration::from_secs(60);
const TOKEN_READ_TIMEOUT: Duration = Duration::from_secs(20);
struct PendingStream {
frame_tx: flume::Sender<Vec<u8>>,
}
type TokenRegistry = DashMap<Uuid, PendingStream>;
pub struct TcpFrameTransport {
advertise_addr: SocketAddr,
registry: Arc<TokenRegistry>,
cancel: CancellationToken,
}
impl TcpFrameTransport {
pub async fn new(bind_addr: IpAddr) -> Result<Arc<Self>> {
let listener = TcpListener::bind((bind_addr, 0u16)).await?;
let local_port = listener.local_addr()?.port();
let advertise_ip = crate::streaming::util::resolve_advertise_ip(bind_addr);
let advertise_addr = SocketAddr::new(advertise_ip, local_port);
let registry: Arc<TokenRegistry> = Arc::new(DashMap::new());
let cancel = CancellationToken::new();
tokio::spawn(run_accept_loop(
Arc::new(listener),
registry.clone(),
cancel.clone(),
));
Ok(Arc::new(Self {
advertise_addr,
registry,
cancel,
}))
}
pub async fn default_bound() -> Result<Arc<Self>> {
Self::new(IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED)).await
}
}
impl Drop for TcpFrameTransport {
fn drop(&mut self) {
self.cancel.cancel();
}
}
async fn run_accept_loop(
listener: Arc<TcpListener>,
registry: Arc<TokenRegistry>,
cancel: CancellationToken,
) {
loop {
tokio::select! {
biased;
_ = cancel.cancelled() => {
tracing::debug!("TCP streaming accept loop cancelled, exiting");
return;
}
res = listener.accept() => match res {
Ok((stream, _peer)) => {
tokio::spawn(handle_one_connection(stream, registry.clone()));
}
Err(e) => {
tracing::warn!("TCP streaming accept error: {}", e);
tokio::time::sleep(Duration::from_millis(50)).await;
}
},
}
}
}
async fn handle_one_connection(mut stream: TcpStream, registry: Arc<TokenRegistry>) {
let mut token_buf = [0u8; 16];
let read_result = tokio::time::timeout(
TOKEN_READ_TIMEOUT,
AsyncReadExt::read_exact(&mut stream, &mut token_buf),
)
.await;
let token = match read_result {
Ok(Ok(_)) => Uuid::from_bytes(token_buf),
Ok(Err(e)) => {
tracing::warn!("TCP streaming handshake read failed: {}", e);
return;
}
Err(_) => {
tracing::warn!(
"TCP streaming handshake timed out after {:?}",
TOKEN_READ_TIMEOUT
);
return;
}
};
let pending = match registry.remove(&token) {
Some((_, p)) => p,
None => {
tracing::warn!(
"TCP streaming: rejecting unknown token {} (expired or never registered)",
token
);
return;
}
};
configure_socket(&stream);
pump_frames(stream, pending.frame_tx).await;
}
async fn pump_frames(stream: TcpStream, frame_tx: flume::Sender<Vec<u8>>) {
let mut framed = Framed::new(stream, TcpFrameCodec::new());
let mut last_was_terminal = false;
let mut consumer_dropped = false;
while let Some(result) = framed.next().await {
match result {
Ok((_msg_type, _header, payload)) => {
let payload_vec = payload.to_vec();
last_was_terminal = is_terminal_sentinel(&payload_vec);
if frame_tx.send_async(payload_vec).await.is_err() {
consumer_dropped = true;
break;
}
}
Err(e) => {
tracing::warn!("TCP streaming read error: {}", e);
break;
}
}
}
if !last_was_terminal && !consumer_dropped {
let _ = frame_tx
.send_async(crate::streaming::sender::cached_dropped().clone())
.await;
}
let mut stream = framed.into_inner();
if let Err(e) = stream.shutdown().await {
tracing::debug!("TCP streaming receiver shutdown: {}", e);
}
}
fn configure_socket(stream: &TcpStream) {
if let Err(e) = stream.set_nodelay(true) {
tracing::warn!("Failed to set TCP_NODELAY: {}", e);
}
let sock = socket2::SockRef::from(stream);
if let Err(e) = sock.set_tcp_keepalive(
&socket2::TcpKeepalive::new()
.with_time(Duration::from_secs(60))
.with_interval(Duration::from_secs(10)),
) {
tracing::warn!("Failed to set TCP keepalive: {}", e);
}
if let Err(e) = sock.set_send_buffer_size(1_048_576) {
tracing::warn!("Failed to set send buffer size: {}", e);
}
if let Err(e) = sock.set_recv_buffer_size(1_048_576) {
tracing::warn!("Failed to set recv buffer size: {}", e);
}
if let Err(e) = sock.set_tcp_user_timeout(Some(Duration::from_secs(30))) {
tracing::warn!("Failed to set TCP_USER_TIMEOUT: {}", e);
}
}
pub fn parse_tcp_endpoint(endpoint: &str) -> Result<(std::net::SocketAddr, Uuid)> {
let stripped = endpoint
.strip_prefix("tcp://")
.ok_or_else(|| anyhow::anyhow!("missing tcp:// prefix: {}", endpoint))?;
let slash_pos = stripped
.rfind('/')
.ok_or_else(|| anyhow::anyhow!("missing token in endpoint: {}", endpoint))?;
let addr_str = &stripped[..slash_pos];
let token_str = &stripped[slash_pos + 1..];
let addr: std::net::SocketAddr = addr_str
.parse()
.map_err(|e| anyhow::anyhow!("invalid address '{}': {}", addr_str, e))?;
let token: Uuid = token_str
.parse()
.map_err(|e| anyhow::anyhow!("invalid token '{}': {}", token_str, e))?;
Ok((addr, token))
}
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
}
}
impl FrameTransport for TcpFrameTransport {
fn bind(
&self,
_anchor_id: u64,
_session_id: u64,
) -> BoxFuture<'_, Result<(String, flume::Receiver<Vec<u8>>)>> {
let advertise_addr = self.advertise_addr;
let registry = self.registry.clone();
Box::pin(async move {
let token = Uuid::new_v4();
let (frame_tx, frame_rx) = flume::bounded::<Vec<u8>>(256);
registry.insert(token, PendingStream { frame_tx });
let expiry_registry = registry.clone();
tokio::spawn(async move {
tokio::time::sleep(ACCEPT_TIMEOUT).await;
if expiry_registry.remove(&token).is_some() {
tracing::warn!(
"TCP streaming: token {} expired before peer connected",
token
);
}
});
let endpoint = format!("tcp://{}/{}", advertise_addr, token);
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, token) = parse_tcp_endpoint(&endpoint)?;
let mut stream = TcpStream::connect(addr).await?;
configure_socket(&stream);
stream.write_all(token.as_bytes()).await?;
let (tx, rx) = flume::bounded::<Vec<u8>>(256);
tokio::spawn(async move {
while let Ok(frame_bytes) = rx.recv_async().await {
if let Err(e) = TcpFrameCodec::encode_frame(
&mut stream,
MessageType::Message,
&[], &frame_bytes,
)
.await
{
tracing::error!("TCP streaming write error: {}", e);
break;
}
}
if let Err(e) = stream.flush().await {
tracing::debug!("TCP streaming flush on close: {}", e);
}
if let Err(e) = stream.shutdown().await {
tracing::debug!("TCP streaming shutdown on close: {}", e);
}
});
Ok(tx)
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test(flavor = "multi_thread")]
async fn test_bind_returns_tcp_endpoint() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (endpoint, _rx) = transport.bind(1, 0).await.unwrap();
assert!(
endpoint.starts_with("tcp://"),
"endpoint should start with tcp://: {}",
endpoint
);
let (addr, token) = parse_tcp_endpoint(&endpoint).unwrap();
assert!(addr.port() > 0, "port should be non-zero");
assert_ne!(token, Uuid::nil(), "token should not be nil UUID");
}
#[tokio::test(flavor = "multi_thread")]
async fn test_connect_handshake() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (endpoint, rx) = transport.bind(1, 0).await.unwrap();
let sender = transport.connect(&endpoint, 1, 1).await.unwrap();
let frame_bytes = rmp_serde::to_vec(&crate::streaming::frame::StreamFrame::<String>::Item(
"hello world".to_string(),
))
.unwrap();
sender.send_async(frame_bytes.clone()).await.unwrap();
let received = tokio::time::timeout(Duration::from_secs(5), rx.recv_async())
.await
.expect("timeout waiting for frame")
.expect("channel closed unexpectedly");
assert_eq!(
received, frame_bytes,
"received frame should match sent frame"
);
drop(sender);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_invalid_token_rejected() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (endpoint, rx) = transport.bind(1, 0).await.unwrap();
let (addr, _valid_token) = parse_tcp_endpoint(&endpoint).unwrap();
let mut stream = TcpStream::connect(addr).await.unwrap();
let wrong_token = Uuid::new_v4();
stream.write_all(wrong_token.as_bytes()).await.unwrap();
let result = tokio::time::timeout(Duration::from_millis(500), rx.recv_async()).await;
assert!(
result.is_err(),
"should not receive frames from invalid token connection"
);
drop(stream);
}
#[test]
fn test_parse_tcp_endpoint_valid() {
let endpoint = "tcp://127.0.0.1:8080/550e8400-e29b-41d4-a716-446655440000";
let (addr, token) = parse_tcp_endpoint(endpoint).unwrap();
assert_eq!(addr.port(), 8080);
assert_eq!(token.to_string(), "550e8400-e29b-41d4-a716-446655440000");
}
#[test]
fn test_parse_tcp_endpoint_invalid_no_prefix() {
assert!(parse_tcp_endpoint("http://127.0.0.1:8080/token").is_err());
}
#[test]
fn test_parse_tcp_endpoint_invalid_no_token() {
assert!(parse_tcp_endpoint("tcp://127.0.0.1:8080").is_err());
}
#[test]
fn test_parse_tcp_endpoint_invalid_bad_uuid() {
assert!(parse_tcp_endpoint("tcp://127.0.0.1:8080/not-a-uuid").is_err());
}
#[test]
fn test_parse_tcp_endpoint_invalid_bad_addr() {
assert!(
parse_tcp_endpoint("tcp://not-an-addr/550e8400-e29b-41d4-a716-446655440000").is_err()
);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_connect_round_trip_multiple_frames() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (endpoint, rx) = transport.bind(42, 0).await.unwrap();
let sender = transport.connect(&endpoint, 42, 1).await.unwrap();
let mut expected_frames = Vec::new();
for i in 0..10 {
let frame =
rmp_serde::to_vec(&crate::streaming::frame::StreamFrame::<u32>::Item(i)).unwrap();
expected_frames.push(frame.clone());
sender.send_async(frame).await.unwrap();
}
for (i, expected) in expected_frames.iter().enumerate() {
let received = tokio::time::timeout(Duration::from_secs(5), rx.recv_async())
.await
.unwrap_or_else(|_| panic!("timeout waiting for frame {}", i))
.unwrap_or_else(|_| panic!("channel closed at frame {}", i));
assert_eq!(&received, expected, "frame {} mismatch", i);
}
}
#[tokio::test(flavor = "multi_thread")]
async fn test_dropped_sentinel_injected_on_abrupt_close() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (endpoint, rx) = transport.bind(1, 0).await.unwrap();
let sender = transport.connect(&endpoint, 1, 1).await.unwrap();
let frame = rmp_serde::to_vec(&crate::streaming::frame::StreamFrame::<String>::Item(
"data".to_string(),
))
.unwrap();
sender.send_async(frame.clone()).await.unwrap();
drop(sender);
let received = tokio::time::timeout(Duration::from_secs(5), rx.recv_async())
.await
.expect("timeout")
.expect("channel closed");
assert_eq!(received, frame);
let sentinel = tokio::time::timeout(Duration::from_secs(5), rx.recv_async())
.await
.expect("timeout waiting for Dropped sentinel")
.expect("channel closed before Dropped sentinel");
assert_eq!(
sentinel.as_slice(),
crate::streaming::sender::cached_dropped().as_slice(),
"should receive Dropped sentinel after abrupt close"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn test_no_extra_dropped_after_finalized() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (endpoint, rx) = transport.bind(1, 0).await.unwrap();
let sender = transport.connect(&endpoint, 1, 1).await.unwrap();
let finalized =
rmp_serde::to_vec(&crate::streaming::frame::StreamFrame::<()>::Finalized).unwrap();
sender.send_async(finalized.clone()).await.unwrap();
drop(sender);
let received = tokio::time::timeout(Duration::from_secs(5), rx.recv_async())
.await
.expect("timeout")
.expect("channel closed");
assert_eq!(received, finalized);
let result = tokio::time::timeout(Duration::from_secs(2), rx.recv_async()).await;
match result {
Ok(Ok(extra)) => {
assert_ne!(
extra.as_slice(),
crate::streaming::sender::cached_dropped().as_slice(),
"should not inject Dropped after Finalized"
);
}
Ok(Err(_)) => {} Err(_) => {} }
}
#[test]
fn test_parse_tcp_endpoint_ipv6() {
let endpoint = "tcp://[::1]:8080/550e8400-e29b-41d4-a716-446655440000";
let (addr, token) = parse_tcp_endpoint(endpoint).unwrap();
assert_eq!(addr, "[::1]:8080".parse::<std::net::SocketAddr>().unwrap());
assert_eq!(token.to_string(), "550e8400-e29b-41d4-a716-446655440000");
}
#[tokio::test(flavor = "multi_thread")]
async fn test_bind_ipv6_endpoint_format() {
if std::net::TcpListener::bind(std::net::SocketAddr::from((
std::net::Ipv6Addr::LOCALHOST,
0,
)))
.is_err()
{
eprintln!("Skipping test_bind_ipv6_endpoint_format: IPv6 not available");
return;
}
let transport = TcpFrameTransport::new(std::net::Ipv6Addr::LOCALHOST.into())
.await
.unwrap();
let (endpoint, _rx) = transport.bind(1, 0).await.unwrap();
assert!(
endpoint.contains("[::1]"),
"IPv6 endpoint must bracket the address: {}",
endpoint
);
let (addr, _token) = parse_tcp_endpoint(&endpoint).unwrap();
assert!(addr.ip().is_loopback());
}
#[tokio::test(flavor = "multi_thread")]
async fn test_bind_unspecified_resolves_to_loopback() {
let transport = TcpFrameTransport::default_bound().await.unwrap(); let (endpoint, _rx) = transport.bind(1, 0).await.unwrap();
let (addr, _token) = parse_tcp_endpoint(&endpoint).unwrap();
assert!(
addr.ip().is_loopback(),
"unspecified bind should resolve to loopback in endpoint: {}",
endpoint
);
}
#[tokio::test(flavor = "multi_thread")]
async fn shared_listener_serves_many_binds_on_one_port() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (endpoint0, _rx0) = transport.bind(0, 0).await.unwrap();
let (addr0, _) = parse_tcp_endpoint(&endpoint0).unwrap();
let mut tokens = std::collections::HashSet::new();
for i in 1u64..50 {
let (endpoint, _rx) = transport.bind(i, 0).await.unwrap();
let (addr, token) = parse_tcp_endpoint(&endpoint).unwrap();
assert_eq!(addr, addr0, "bind #{i} advertised a different address");
assert!(tokens.insert(token), "duplicate token {token} on bind #{i}");
}
}
#[tokio::test(flavor = "multi_thread")]
async fn unknown_token_is_rejected_and_does_not_disrupt_others() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (endpoint, rx) = transport.bind(1, 0).await.unwrap();
let (addr, valid_token) = parse_tcp_endpoint(&endpoint).unwrap();
let mut bad = TcpStream::connect(addr).await.unwrap();
bad.write_all(Uuid::new_v4().as_bytes()).await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
let mut good = TcpStream::connect(addr).await.unwrap();
good.write_all(valid_token.as_bytes()).await.unwrap();
let payload =
rmp_serde::to_vec(&crate::streaming::frame::StreamFrame::<u32>::Item(7)).unwrap();
TcpFrameCodec::encode_frame(&mut good, MessageType::Message, &[], &payload)
.await
.unwrap();
let received = tokio::time::timeout(Duration::from_secs(2), rx.recv_async())
.await
.expect("recv timeout — bad token disrupted the registered stream")
.expect("channel closed unexpectedly");
assert_eq!(received, payload);
}
#[tokio::test(flavor = "current_thread", start_paused = true)]
async fn token_expires_when_no_peer_connects() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let (_endpoint, rx) = transport.bind(99, 0).await.unwrap();
assert_eq!(transport.registry.len(), 1, "token should be registered");
tokio::time::sleep(ACCEPT_TIMEOUT + Duration::from_secs(1)).await;
match tokio::time::timeout(Duration::from_secs(120), rx.recv_async()).await {
Ok(Err(_)) => {} Ok(Ok(b)) => panic!("expected closed channel, got frame: {b:?}"),
Err(_) => panic!(
"consumer channel did not close after expiry — an extra Sender \
clone is leaking past the timeout and pinning the channel open"
),
}
assert_eq!(
transport.registry.len(),
0,
"expired token should be removed from registry"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn concurrent_binds_all_succeed_through_one_listener() {
let transport = TcpFrameTransport::default_bound().await.unwrap();
let mut handles = Vec::new();
for i in 0u64..32 {
let t = transport.clone();
handles.push(tokio::spawn(async move {
let (endpoint, rx) = t.bind(i, 0).await.unwrap();
let sender = t.connect(&endpoint, i, 0).await.unwrap();
let payload =
rmp_serde::to_vec(&crate::streaming::frame::StreamFrame::<u64>::Item(i))
.unwrap();
sender.send_async(payload.clone()).await.unwrap();
let received = tokio::time::timeout(Duration::from_secs(5), rx.recv_async())
.await
.expect("recv timeout")
.expect("channel closed");
assert_eq!(received, payload);
}));
}
for h in handles {
h.await.unwrap();
}
}
}