//! Bounded one-control-frame dispatch for a bidirectional QUIC stream. use std::{net::SocketAddr, sync::Arc, time::Duration}; use futures::{SinkExt as _, StreamExt as _}; use lanspread_proto::{ ControlErrorCode, ControlMessage, MAX_CONTROL_FRAME_BYTES, MAX_REQUEST_FRAME_BYTES, Request, Response, }; use s2n_quic::{ application, stream::{BidirectionalStream, SendStream}, }; use tokio::sync::{OwnedSemaphorePermit, Semaphore}; use tokio_util::{ codec::{FramedRead, FramedWrite, LengthDelimitedCodec}, sync::CancellationToken, }; use crate::{ context::PeerCtx, services::{ remote_state, state_sync::StateDomain, transfer::{ChunkDispatch, handle_file_chunk_request, handle_stream_install_request}, }, }; type ResponseWriter = FramedWrite; const INBOUND_CONTROL_FRAME_TIMEOUT: Duration = Duration::from_secs(10); const OUTBOUND_CONTROL_IO_TIMEOUT: Duration = Duration::from_secs(10); /// Response-side codec: snapshots may approach the full control-frame bound. fn control_codec() -> LengthDelimitedCodec { LengthDelimitedCodec::builder() .max_frame_length(MAX_CONTROL_FRAME_BYTES) .new_codec() } /// Request-side codec for anonymous inbound streams. The length-delimited /// decoder reserves the declared frame length up front, so the public /// responder only ever grants the small request allowance per stream. fn request_codec() -> LengthDelimitedCodec { LengthDelimitedCodec::builder() .max_frame_length(MAX_REQUEST_FRAME_BYTES) .new_codec() } /// Reads exactly one bounded request frame, requires request-side EOF, sends at /// most one control response, and then closes the stream. Raw transfer requests /// consume the response side after the same single control-frame admission. pub(super) async fn handle_peer_stream( stream: BidirectionalStream, ctx: PeerCtx, remote_addr: Option, stream_shutdown: CancellationToken, control_permit: OwnedSemaphorePermit, bulk_transfer_permits: Arc, stream_install_permits: Arc, ) -> eyre::Result<()> { let (rx, tx) = stream.split(); let mut framed_rx = FramedRead::new(rx, request_codec()); let mut framed_tx = FramedWrite::new(tx, control_codec()); log::trace!("{remote_addr:?} peer stream opened"); let source_ip = remote_addr.map(|addr| addr.ip()); let first_frame = read_expected_frame(&mut framed_rx, &stream_shutdown).await; let mut control_permit = Some(control_permit); let mut _bulk_permit = None; let mut response_reset = false; match first_frame { FrameRead::Frame(data) => { let trailing = read_expected_eof(&mut framed_rx, &stream_shutdown).await; if trailing == TrailingRead::Eof { match Request::decode(data.freeze()) { Ok(request) => { log::debug!("{remote_addr:?} msg: {request:?}"); if request_is_bulk(&request) { let bulk_permit = Arc::clone(&bulk_transfer_permits).try_acquire_owned(); // Once the single bounded request is decoded, bulk // work moves to its smaller pool so it cannot hold // every control-plane permit during long egress. drop(control_permit.take()); if let Ok(permit) = bulk_permit { _bulk_permit = Some(permit); let dispatched = dispatch_request( &ctx, request, source_ip, framed_tx, &stream_shutdown, &stream_install_permits, ) .await; framed_tx = dispatched.writer; response_reset = dispatched.response_reset; } else { let mut tx = framed_tx.into_inner(); let _ = tx.reset(application::Error::UNKNOWN); framed_tx = FramedWrite::new(tx, control_codec()); response_reset = true; } } else { let dispatched = dispatch_request( &ctx, request, source_ip, framed_tx, &stream_shutdown, &stream_install_permits, ) .await; framed_tx = dispatched.writer; response_reset = dispatched.response_reset; } } Err(error) => { log::warn!( "Rejecting invalid control request from {remote_addr:?}: {error}" ); framed_tx = send_response( framed_tx, Response::Error(ControlErrorCode::InvalidRequest), "invalid-request", &stream_shutdown, ) .await; } } } else if trailing != TrailingRead::Cancelled { log::warn!("Rejecting non-singular control request from {remote_addr:?}"); framed_tx = send_response( framed_tx, Response::Error(ControlErrorCode::InvalidRequest), "invalid-request", &stream_shutdown, ) .await; } } FrameRead::Invalid(error) => { log::warn!("Rejecting malformed control frame from {remote_addr:?}: {error}"); framed_tx = send_response( framed_tx, Response::Error(ControlErrorCode::InvalidRequest), "invalid-request", &stream_shutdown, ) .await; } FrameRead::Eof => log::trace!("{remote_addr:?} peer stream closed without a request"), FrameRead::Cancelled => {} } close_or_reset_stream( framed_rx, framed_tx, remote_addr, &stream_shutdown, response_reset, ) .await; Ok(()) } const fn request_is_bulk(request: &Request) -> bool { matches!( request, Request::GetGameFileChunk { .. } | Request::StreamInstall { .. } ) } enum FrameRead { Frame(bytes::BytesMut), Invalid(std::io::Error), Eof, Cancelled, } async fn read_expected_frame( framed_rx: &mut FramedRead, cancellation: &CancellationToken, ) -> FrameRead { tokio::select! { biased; () = cancellation.cancelled() => FrameRead::Cancelled, () = tokio::time::sleep(INBOUND_CONTROL_FRAME_TIMEOUT) => FrameRead::Invalid( std::io::Error::new(std::io::ErrorKind::TimedOut, "control request timed out") ), frame = framed_rx.next() => match frame { Some(Ok(bytes)) => FrameRead::Frame(bytes), Some(Err(error)) => FrameRead::Invalid(error), None => FrameRead::Eof, } } } #[derive(Clone, Copy, Debug, Eq, PartialEq)] enum TrailingRead { Eof, ExtraFrame, Invalid, TimedOut, Cancelled, } async fn read_expected_eof( framed_rx: &mut FramedRead, cancellation: &CancellationToken, ) -> TrailingRead { tokio::select! { biased; () = cancellation.cancelled() => TrailingRead::Cancelled, () = tokio::time::sleep(INBOUND_CONTROL_FRAME_TIMEOUT) => TrailingRead::TimedOut, frame = framed_rx.next() => match frame { Some(Ok(_)) => TrailingRead::ExtraFrame, Some(Err(_)) => TrailingRead::Invalid, None => TrailingRead::Eof, } } } async fn dispatch_request( ctx: &PeerCtx, request: Request, source_ip: Option, framed_tx: ResponseWriter, stream_shutdown: &CancellationToken, stream_install_permits: &Arc, ) -> DispatchResult { match request { Request::Ping => { match control_io_with_deadline(remote_state::local_revisions(ctx), stream_shutdown) .await { Some(Ok(revisions)) => DispatchResult::close( send_response( framed_tx, Response::Pong(revisions), "pong", stream_shutdown, ) .await, ), Some(Err(error)) => { log::error!("Failed to build local revisions: {error:#}"); DispatchResult::close( send_response( framed_tx, Response::Error(ControlErrorCode::Internal), "pong-error", stream_shutdown, ) .await, ) } None => reset_response_writer(framed_tx, "pong-computation"), } } Request::Hello => { match control_io_with_deadline(remote_state::local_snapshot(ctx), stream_shutdown).await { Some(Ok(snapshot)) => DispatchResult::close( send_response( framed_tx, Response::HelloSnapshot(snapshot), "hello-snapshot", stream_shutdown, ) .await, ), Some(Err(error)) => { log::error!("Failed to build local peer snapshot: {error:#}"); DispatchResult::close( send_response( framed_tx, Response::Error(ControlErrorCode::Internal), "hello-error", stream_shutdown, ) .await, ) } None => reset_response_writer(framed_tx, "hello-computation"), } } Request::LibraryChanged(hint) => { ctx.state_sync .schedule_hint(StateDomain::Library, hint, source_ip); DispatchResult::close(framed_tx) } Request::CallToPlayChanged(hint) => { ctx.state_sync .schedule_hint(StateDomain::CallToPlay, hint, source_ip); DispatchResult::close(framed_tx) } Request::GetGameFileChunk { game_id, content_id, relative_path, offset, length, } => { match handle_file_chunk_request( ctx, game_id, content_id, relative_path, offset, length, framed_tx, stream_shutdown, ) .await { ChunkDispatch::Finished(writer) => DispatchResult::close(writer), ChunkDispatch::Reset(writer) => DispatchResult::reset(writer), } } Request::StreamInstall { game_id, content_id, } => { let Ok(stream_install_permit) = Arc::clone(stream_install_permits).try_acquire_owned() else { let mut tx = framed_tx.into_inner(); let _ = tx.reset(application::Error::UNKNOWN); return DispatchResult::reset(FramedWrite::new(tx, control_codec())); }; DispatchResult::close( handle_stream_install_request( ctx, game_id, content_id, framed_tx, stream_shutdown, stream_install_permit, ) .await, ) } } } fn reset_response_writer(framed_tx: ResponseWriter, label: &str) -> DispatchResult { let mut tx = framed_tx.into_inner(); if let Err(error) = tx.reset(application::Error::UNKNOWN) { log::debug!("Failed to reset timed-out {label} response: {error}"); } DispatchResult::reset(FramedWrite::new(tx, control_codec())) } struct DispatchResult { writer: ResponseWriter, response_reset: bool, } impl DispatchResult { const fn close(writer: ResponseWriter) -> Self { Self { writer, response_reset: false, } } const fn reset(writer: ResponseWriter) -> Self { Self { writer, response_reset: true, } } } async fn send_response( mut framed_tx: ResponseWriter, response: Response, label: &str, stream_shutdown: &CancellationToken, ) -> ResponseWriter { let encoded = match response.encode() { Ok(encoded) => encoded, Err(error) => { log::error!("Failed to encode {label} response: {error}"); let mut tx = framed_tx.into_inner(); if let Err(reset_error) = tx.reset(application::Error::UNKNOWN) { log::debug!("Failed to reset unencodable {label} response: {reset_error}"); } return FramedWrite::new(tx, control_codec()); } }; let send_result = control_io_with_deadline(framed_tx.send(encoded), stream_shutdown).await; let Some(send_result) = send_result else { let mut tx = framed_tx.into_inner(); let _ = tx.reset(application::Error::UNKNOWN); return FramedWrite::new(tx, control_codec()); }; if let Err(error) = send_result { log::debug!("Failed to send {label} response: {error}"); } framed_tx } async fn close_or_reset_stream( framed_rx: FramedRead, mut framed_tx: ResponseWriter, remote_addr: Option, cancellation: &CancellationToken, response_reset: bool, ) { if cancellation.is_cancelled() { let mut rx = framed_rx.into_inner(); let _ = rx.stop_sending(application::Error::UNKNOWN); let mut tx = framed_tx.into_inner(); let _ = tx.reset(application::Error::UNKNOWN); return; } if response_reset { // The transfer handler already sent RESET_STREAM. A later clean FIN // would make a rejected zero-byte or truncated raw chunk ambiguous to // the receiver. drop(framed_rx); drop(framed_tx); return; } let close_result = control_io_with_deadline(framed_tx.close(), cancellation).await; if close_result.is_none() { let mut rx = framed_rx.into_inner(); let _ = rx.stop_sending(application::Error::UNKNOWN); let mut tx = framed_tx.into_inner(); let _ = tx.reset(application::Error::UNKNOWN); return; } if let Some(Err(error)) = close_result { log::debug!("{remote_addr:?} failed to close peer response stream: {error}"); } } async fn control_io_with_deadline( operation: impl std::future::Future, cancellation: &CancellationToken, ) -> Option { tokio::select! { biased; () = cancellation.cancelled() => None, () = tokio::time::sleep(OUTBOUND_CONTROL_IO_TIMEOUT) => None, result = operation => Some(result), } } #[cfg(test)] mod tests { use lanspread_db::content_manifest::ContentId; use super::*; #[test] fn every_control_codec_enforces_the_protocol_frame_bound() { let codec = control_codec(); assert_eq!(codec.max_frame_length(), MAX_CONTROL_FRAME_BYTES); } #[test] fn inbound_request_decoders_use_the_small_request_bound() { let codec = request_codec(); assert_eq!(codec.max_frame_length(), MAX_REQUEST_FRAME_BYTES); } #[test] fn trailing_frame_outcomes_are_never_accepted_as_eof() { for outcome in [ TrailingRead::ExtraFrame, TrailingRead::Invalid, TrailingRead::TimedOut, ] { assert_ne!(outcome, TrailingRead::Eof); assert_ne!(outcome, TrailingRead::Cancelled); } } #[tokio::test(start_paused = true)] async fn public_control_computation_and_egress_have_an_absolute_deadline() { let cancellation = CancellationToken::new(); let start = tokio::time::Instant::now(); assert!( control_io_with_deadline(std::future::pending::<()>(), &cancellation) .await .is_none() ); assert_eq!(start.elapsed(), OUTBOUND_CONTROL_IO_TIMEOUT); } #[test] fn only_long_lived_payload_requests_move_to_the_reserved_bulk_pool() { assert!(!request_is_bulk(&Request::Ping)); assert!(request_is_bulk(&Request::StreamInstall { game_id: "game".to_owned(), content_id: ContentId::from_bytes([1; 32]), })); } }