diff --git a/dev_bench/src/main.rs b/dev_bench/src/main.rs index a769af28d..d051f9818 100644 --- a/dev_bench/src/main.rs +++ b/dev_bench/src/main.rs @@ -712,15 +712,20 @@ fn run_rewritten_iperf3(ctx: BenchCtx<'_>) -> Result<()> { let client_sh = xshell::Shell::new()?; let max_server_attempts = 5; let max_client_attempts = 50; + let guest_port = "5201"; let mut last_failure = "sandboxed server did not become ready".to_owned(); for server_attempt in 1..=max_server_attempts { - let port = TcpListener::bind((Ipv4Addr::LOCALHOST, 0))? + let host_port = TcpListener::bind((Ipv4Addr::LOCALHOST, 0))? .local_addr()? .port() .to_string(); let mut server_command = std::process::Command::new(&broker); server_command + .arg("--broker-ipv4-address") + .arg("127.0.0.1") + .arg("--tcp-port-mapping") + .arg(format!("{host_port}:{guest_port}")) .arg("--runner") .arg(&runner) .arg("--") @@ -733,7 +738,7 @@ fn run_rewritten_iperf3(ctx: BenchCtx<'_>) -> Result<()> { ]) .arg(&tar_file) .arg(&iperf3_rewritten) - .args(["-s", "-1", "-B", "127.0.0.1", "-p", &port]); + .args(["-s", "-1", "-B", "127.0.0.1", "-p", guest_port]); if COMMAND_EXECUTION_IS_QUIET.load(Relaxed) { server_command .stdout(std::process::Stdio::null()) @@ -749,7 +754,7 @@ fn run_rewritten_iperf3(ctx: BenchCtx<'_>) -> Result<()> { for client_attempt in 1..=max_client_attempts { let result = cmd!( client_sh, - "{iperf3_host} -c 127.0.0.1 -p {port} --bytes 1G --connect-timeout 50" + "{iperf3_host} -c 127.0.0.1 -p {host_port} --bytes 1G --connect-timeout 50" ) .quiet() .ignore_stdout() diff --git a/litebox_broker_core/src/lib.rs b/litebox_broker_core/src/lib.rs index 6bf3567fe..277105e80 100644 --- a/litebox_broker_core/src/lib.rs +++ b/litebox_broker_core/src/lib.rs @@ -40,7 +40,7 @@ pub use policy::{ }; use session::ObjectReference; pub use session::{BrokerSession, CallerCredential, ObjectRights, SessionId}; -use socket::SocketProvider; +use socket::{BrokerSocketPorts, SocketProvider, TcpPortMappingConfig}; /// BrokerCore result type. pub type Result = core::result::Result; @@ -117,6 +117,8 @@ pub struct BrokerCore { pub(crate) reserved_pipe_capacity: Arc, pub(crate) reserved_sockets: Arc, pub(crate) socket_provider: Arc, + pub(crate) tcp_port_mapping_config: TcpPortMappingConfig, + pub(crate) socket_ports: BrokerSocketPorts, } static BROKER_CORE_CREATED: AtomicBool = AtomicBool::new(false); @@ -133,6 +135,8 @@ impl BrokerCore { limits: BrokerCoreLimits, socket_provider: Arc, ) -> Result { + let tcp_port_mapping_config = + TcpPortMappingConfig::new(socket_provider.tcp_port_mappings())?; BROKER_CORE_CREATED .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) .map_err(|_| BrokerError::BrokerCoreAlreadyExists)?; @@ -147,6 +151,8 @@ impl BrokerCore { reserved_pipe_capacity: Arc::new(AtomicUsize::new(0)), reserved_sockets: Arc::new(AtomicUsize::new(0)), socket_provider, + tcp_port_mapping_config, + socket_ports: BrokerSocketPorts::default(), }) } diff --git a/litebox_broker_core/src/policy.rs b/litebox_broker_core/src/policy.rs index f12a9b2b1..81656b2ad 100644 --- a/litebox_broker_core/src/policy.rs +++ b/litebox_broker_core/src/policy.rs @@ -466,16 +466,17 @@ impl PolicyEngine { address: SocketAddrV4, ) -> Result<(), BrokerError> { self.principal_object_rights(caller_credential)?; - // Egress rules do not describe local listener authority. Socket - // creation admission plus this fixed loopback boundary governs binds. - let supported_socket = matches!( - (request.socket_type, request.protocol), - (SocketType::Stream, IpProtocol::Tcp) | (SocketType::Datagram, IpProtocol::Udp) - ); - if request.address_family == AddressFamily::Ipv4 - && supported_socket - && address.ip().is_loopback() - { + // Egress rules do not describe local listener authority. TCP may bind + // loopback for a private endpoint or unspecified for the broker's + // configured external interface. UDP remains loopback-only. + let permitted_address = match (request.socket_type, request.protocol) { + (SocketType::Stream, IpProtocol::Tcp) => { + address.ip().is_loopback() || address.ip().is_unspecified() + } + (SocketType::Datagram, IpProtocol::Udp) => address.ip().is_loopback(), + _ => false, + }; + if request.address_family == AddressFamily::Ipv4 && permitted_address { Ok(()) } else { Err(BrokerError::PolicyDenied) @@ -787,5 +788,13 @@ mod tests { ), Err(BrokerError::PolicyDenied) ); + assert_eq!( + policy.authorize_socket_bind( + CallerCredential::Unauthenticated, + IPV4_TCP, + address([0, 0, 0, 0], 0), + ), + Ok(()) + ); } } diff --git a/litebox_broker_core/src/socket.rs b/litebox_broker_core/src/socket.rs index 74a0f827b..0f0cdff83 100644 --- a/litebox_broker_core/src/socket.rs +++ b/litebox_broker_core/src/socket.rs @@ -7,6 +7,7 @@ use alloc::{sync::Arc, vec::Vec}; use core::net::{Ipv4Addr, SocketAddrV4}; use core::sync::atomic::{AtomicUsize, Ordering}; +use hashbrown::{HashMap, HashSet}; use litebox_broker_protocol::ObjectHandle; use litebox_broker_protocol::readiness::ReadinessFlags; use litebox_broker_protocol::socket::{ @@ -15,20 +16,184 @@ use litebox_broker_protocol::socket::{ SocketConnectionStatus, SocketError, SocketOutcome, SocketStatusResponse, SocketType, TcpOptionName, TcpOptionValue, }; +use spin::Mutex; use spin::Once; use crate::readiness::{ReadinessRegistration, ReadinessSink}; use crate::session::{ObjectEntry, ObjectRights}; use crate::{BrokerError, BrokerSession, Result, SessionId}; -const DEFAULT_TCP_LISTEN_ADDRESS: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0); +const DEFAULT_TCP_LISTEN_ADDRESS: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0); +const DEFAULT_TCP_LOCAL_ADDRESS: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0); +/// First port in the IANA dynamic/private range used for guest-local allocation. +const FIRST_EPHEMERAL_PORT: u16 = 49152; + +/// Portable mapping from one broker TCP port to a guest-local TCP port. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct TcpPortMapping { + /// TCP port on the broker's configured host-facing IPv4 address. + pub broker_port: u16, + /// Guest-local TCP port eligible for the mapping. + pub guest_port: u16, +} + +/// Broker startup configuration overriding selected guest TCP ports. +/// +/// This table does not track active listeners or runtime port ownership. When +/// an external listener starts, an absent guest port uses the identity mapping +/// (`broker_port = guest_port`). [`BrokerSocketPorts`] owns live guest ports, +/// while the platform owns each realized broker endpoint. +#[derive(Clone)] +pub(crate) struct TcpPortMappingConfig { + broker_ports: Arc>, +} + +impl TcpPortMappingConfig { + pub(crate) fn new(mappings: &[TcpPortMapping]) -> Result { + let mut broker_ports: HashMap = HashMap::new(); + broker_ports + .try_reserve(mappings.len()) + .map_err(|_| BrokerError::OutOfMemory)?; + for mapping in mappings.iter().copied() { + if mapping.guest_port == 0 + || mapping.broker_port == 0 + || broker_ports.contains_key(&mapping.guest_port) + || broker_ports + .values() + .any(|existing| *existing == mapping.broker_port) + { + return Err(BrokerError::Internal); + } + if broker_ports + .insert(mapping.guest_port, mapping.broker_port) + .is_some() + { + return Err(BrokerError::Internal); + } + } + Ok(Self { + broker_ports: Arc::new(broker_ports), + }) + } + + fn mapping_for_guest_port(&self, guest_port: u16) -> Option { + self.broker_ports + .get(&guest_port) + .copied() + .map(|broker_port| TcpPortMapping { + broker_port, + guest_port, + }) + } + + fn contains_guest_port(&self, guest_port: u16) -> bool { + self.broker_ports.contains_key(&guest_port) + } +} + +#[derive(Default)] +struct BrokerSocketPortState { + guest_tcp_ports: HashSet, + next_guest_tcp_ephemeral: Option, +} + +/// Broker-wide authority for the guest-visible TCP port namespace. +/// +/// Sessions identify endpoint owners but do not define separate network +/// namespaces. Guest TCP ports remain independent of backend host ports. +#[derive(Clone, Default)] +pub(crate) struct BrokerSocketPorts { + state: Arc>, +} + +impl BrokerSocketPorts { + /// Reserves one guest TCP port for `requested_address`. + /// + /// A zero requested port is allocated from the ephemeral range, skipping + /// ports reserved for explicit TCP mappings. The + /// reservation is released when the returned guard is dropped. + fn reserve( + &self, + request: CreateSocketRequest, + requested_address: SocketAddrV4, + mut port_is_reserved: impl FnMut(u16) -> bool, + ) -> Result> { + if !is_tcp(request) { + return Err(BrokerError::Internal); + } + let mut state = self.state.lock(); + let port = if requested_address.port() == 0 { + state.allocate_ephemeral(&mut port_is_reserved)? + } else if state.guest_tcp_ports.contains(&requested_address.port()) { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } else { + requested_address.port() + }; + let local_address = SocketAddrV4::new(*requested_address.ip(), port); + state + .guest_tcp_ports + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory)?; + if !state.guest_tcp_ports.insert(port) { + return Err(BrokerError::Internal); + } + drop(state); + Ok(SocketOutcome::Completed(( + local_address, + GuestPortReservation { + ports: self.clone(), + port, + }, + ))) + } +} + +impl BrokerSocketPortState { + fn allocate_ephemeral( + &mut self, + port_is_reserved: &mut impl FnMut(u16) -> bool, + ) -> Result { + let start = self + .next_guest_tcp_ephemeral + .unwrap_or(FIRST_EPHEMERAL_PORT); + let mut port = start; + loop { + if !self.guest_tcp_ports.contains(&port) && !port_is_reserved(port) { + self.next_guest_tcp_ephemeral = Some(if port == u16::MAX { + FIRST_EPHEMERAL_PORT + } else { + port + 1 + }); + return Ok(port); + } + port = if port == u16::MAX { + FIRST_EPHEMERAL_PORT + } else { + port + 1 + }; + if port == start { + return Err(BrokerError::ResourceExhausted); + } + } + } +} + +/// Guard releasing one guest TCP port back to the broker-wide namespace. +struct GuestPortReservation { + ports: BrokerSocketPorts, + port: u16, +} + +impl Drop for GuestPortReservation { + fn drop(&mut self) { + self.ports.state.lock().guest_tcp_ports.remove(&self.port); + } +} /// Platform socket and endpoint metadata returned by an accept operation. pub struct AcceptedPlatformSocket { /// Accepted nonblocking platform socket. pub socket: Arc, - /// Local endpoint of the accepted connection. - pub local_address: SocketAddrV4, /// Remote endpoint of the accepted connection. pub remote_address: SocketAddrV4, } @@ -78,9 +243,18 @@ pub struct PlatformDatagramReceive { /// bookkeeping shared across the sockets of a broker session. Operations on an /// individual socket belong to [`PlatformSocket`], not this shared provider. pub trait SocketProvider: Send + Sync { + /// Returns broker startup overrides from guest TCP ports to broker ports. + /// + /// BrokerCore copies and validates this configuration during construction. + /// Active listeners and realized platform endpoints are tracked elsewhere. + fn tcp_port_mappings(&self) -> &[TcpPortMapping] { + &[] + } + /// Creates one nonblocking socket resource for an authenticated session. /// - /// The returned socket must not retain authority beyond its `Arc` lifetime. + /// Any provider-retained clones must become inert when + /// [`PlatformSocket::retire`] is called. fn create( &self, session_id: SessionId, @@ -95,14 +269,28 @@ pub trait SocketProvider: Send + Sync { /// One nonblocking socket resource created by [`SocketProvider`]. /// /// The broker retains this resource in an `Arc`, allowing an operation already -/// in flight to finish after its object handle closes. Dropping the final `Arc` -/// releases the platform socket. +/// in flight to finish after its object handle closes. Before releasing its +/// portable authority, core explicitly retires the platform socket. pub trait PlatformSocket: Send + Sync { - /// Binds this socket to a local address. + /// Binds this socket to a local address and echoes the assigned address. + /// + /// A stream socket receives a broker-reserved guest-local address with a + /// nonzero port and must echo it unchanged; the host endpoint backing the + /// socket is chosen privately by the platform. A datagram socket receives + /// the address requested by the guest and returns the host-assigned one. fn bind(&self, address: SocketAddrV4) -> Result>; /// Makes this socket listen for incoming connections. - fn listen(&self, backlog: u32) -> Result>; + /// + /// The returned address is the socket's guest-local address. `mapping` + /// identifies the broker-facing endpoint to realize. `None` requests an + /// ordinary private listener. The platform validates that the mapping's + /// guest port matches this socket's guest-local port. + fn listen( + &self, + backlog: u32, + mapping: Option, + ) -> Result>; /// Accepts one pending connection without waiting. fn accept( @@ -179,6 +367,9 @@ pub trait PlatformSocket: Send + Sync { /// `Connected` and may change between those states through peer updates. fn status(&self) -> Result; + /// Synchronously and idempotently ends this socket's platform authority. + fn retire(&self); + /// Returns the current readiness snapshot. fn readiness(&self) -> ReadinessFlags; } @@ -220,6 +411,7 @@ pub fn create( platform_socket: Once::new(), readiness, _quota: quota, + port_reservation: Mutex::new(None), }); let platform_socket = match session.core.socket_provider.create( session.session_id, @@ -277,7 +469,7 @@ pub fn connect( return connect_datagram(&object, address); } - let resource = { + let (resource, needs_bind) = { let mut object = object.write(); let ObjectEntry::Socket(socket) = &mut *object else { return Err(BrokerError::InvalidRights); @@ -292,8 +484,46 @@ pub fn connect( return Ok(SocketOutcome::Completed(socket.connection_status)); } socket.connect_in_flight = true; - Arc::clone(&socket.resource) + (Arc::clone(&socket.resource), socket.local_address.is_none()) }; + if needs_bind { + match session.core.policy.authorize_socket_bind( + session.caller_credential, + create_request, + DEFAULT_TCP_LOCAL_ADDRESS, + ) { + Ok(()) => {} + Err(BrokerError::PolicyDenied) => { + finish_connect(&object, SocketConnectionStatus::Unconnected); + return Ok(SocketOutcome::Failed(SocketError::PolicyDenied)); + } + Err(error) => { + finish_connect(&object, SocketConnectionStatus::Unconnected); + return Err(error); + } + } + let binding = match reserve_and_bind( + session, + create_request, + &resource, + DEFAULT_TCP_LOCAL_ADDRESS, + ) { + Ok(binding) => binding, + Err(error) => { + finish_connect(&object, SocketConnectionStatus::Unconnected); + return Err(error); + } + }; + match binding { + SocketOutcome::Completed((local_address, reservation)) => { + attach_binding(&object, local_address, reservation); + } + SocketOutcome::Failed(error) => { + finish_connect(&object, SocketConnectionStatus::Unconnected); + return Ok(SocketOutcome::Failed(error)); + } + } + } let status = match resource.connect(address) { Ok(SocketConnectionStatus::Unconnected) => { finish_connect(&object, SocketConnectionStatus::Failed(SocketError::Other)); @@ -313,7 +543,7 @@ pub fn connect( Ok(SocketOutcome::Completed(status)) } -/// Binds a socket to an authorized loopback address. +/// Binds a socket to an authorized guest-local address. pub fn bind( session: &BrokerSession, handle: ObjectHandle, @@ -343,32 +573,50 @@ pub fn bind( ) { Ok(()) => {} Err(BrokerError::PolicyDenied) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Ok(SocketOutcome::Failed(SocketError::PolicyDenied)); } Err(error) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Err(error); } } - let outcome = resource.bind(address); - match outcome { - Ok(SocketOutcome::Completed(local_address)) => { - finish_configuration(&object, Some(local_address), false); - Ok(SocketOutcome::Completed(local_address)) - } - Ok(SocketOutcome::Failed(error)) => { - finish_configuration(&object, None, false); - Ok(SocketOutcome::Failed(error)) - } + if !is_tcp(create_request) { + // Datagram sockets remain backed directly by a host endpoint. + return match resource.bind(address) { + Ok(SocketOutcome::Completed(local_address)) => { + finish_configuration(&object, Some(local_address), None, false); + Ok(SocketOutcome::Completed(local_address)) + } + Ok(SocketOutcome::Failed(error)) => { + finish_configuration(&object, None, None, false); + Ok(SocketOutcome::Failed(error)) + } + Err(error) => { + finish_configuration(&object, None, None, false); + Err(error) + } + }; + } + let binding = match reserve_and_bind(session, create_request, &resource, address) { + Ok(binding) => binding, Err(error) => { - finish_configuration(&object, None, false); - Err(error) + finish_configuration(&object, None, None, false); + return Err(error); } - } + }; + let (local_address, reservation) = match binding { + SocketOutcome::Completed(binding) => binding, + SocketOutcome::Failed(error) => { + finish_configuration(&object, None, None, false); + return Ok(SocketOutcome::Failed(error)); + } + }; + finish_configuration(&object, Some(local_address), Some(reservation), false); + Ok(SocketOutcome::Completed(local_address)) } -/// Makes a socket listen for incoming loopback connections. +/// Makes a socket listen for incoming connections. pub fn listen( session: &BrokerSession, handle: ObjectHandle, @@ -378,7 +626,7 @@ pub fn listen( return Err(BrokerError::UnsupportedOperation); } let object = session.authorized_object(handle, ObjectRights::WRITE)?; - let (resource, create_request, needs_bind) = { + let (resource, create_request, existing_local_address) = { let mut object = object.write(); let ObjectEntry::Socket(socket) = &mut *object else { return Err(BrokerError::InvalidRights); @@ -396,12 +644,13 @@ pub fn listen( ( Arc::clone(&socket.resource), socket.create_request, - socket.local_address.is_none(), + socket.local_address, ) }; - let mut local_address = None; - if needs_bind { + let mut local_address = existing_local_address; + let mut port_reservation = None; + if local_address.is_none() { match session.core.policy.authorize_socket_bind( session.caller_credential, create_request, @@ -409,39 +658,69 @@ pub fn listen( ) { Ok(()) => {} Err(BrokerError::PolicyDenied) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Ok(SocketOutcome::Failed(SocketError::PolicyDenied)); } Err(error) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Err(error); } } - match resource.bind(DEFAULT_TCP_LISTEN_ADDRESS) { - Ok(SocketOutcome::Completed(address)) => local_address = Some(address), - Ok(SocketOutcome::Failed(error)) => { - finish_configuration(&object, None, false); - return Ok(SocketOutcome::Failed(error)); - } + let binding = match reserve_and_bind( + session, + create_request, + &resource, + DEFAULT_TCP_LISTEN_ADDRESS, + ) { + Ok(binding) => binding, Err(error) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Err(error); } + }; + match binding { + SocketOutcome::Completed((address, reservation)) => { + local_address = Some(address); + port_reservation = Some(reservation); + } + SocketOutcome::Failed(error) => { + finish_configuration(&object, None, None, false); + return Ok(SocketOutcome::Failed(error)); + } } } - match resource.listen(backlog) { + let mapping = local_address + .filter(|address| address.ip().is_unspecified()) + .map(|address| { + session + .core + .tcp_port_mapping_config + .mapping_for_guest_port(address.port()) + .unwrap_or(TcpPortMapping { + broker_port: address.port(), + guest_port: address.port(), + }) + }); + + match resource.listen(backlog, mapping) { Ok(SocketOutcome::Completed(address)) => { - local_address = Some(address); - finish_configuration(&object, local_address, true); + // The guest-local address is broker-authoritative, so a platform + // that reports a different one is not trustworthy. + if local_address != Some(address) { + resource.retire(); + finish_retired_configuration(&object, local_address, port_reservation); + return Err(BrokerError::Internal); + } + finish_configuration(&object, local_address, port_reservation, true); Ok(SocketOutcome::Completed(address)) } Ok(SocketOutcome::Failed(error)) => { - finish_configuration(&object, local_address, false); + finish_configuration(&object, local_address, port_reservation, false); Ok(SocketOutcome::Failed(error)) } Err(error) => { - finish_configuration(&object, local_address, false); + finish_configuration(&object, local_address, port_reservation, false); Err(error) } } @@ -454,7 +733,7 @@ pub fn accept( readiness_sink: Arc, ) -> Result> { let listener = session.authorized_object(handle, ObjectRights::WAIT)?; - let (listener_resource, create_request) = { + let (listener_resource, create_request, local_address) = { let listener = listener.read(); let ObjectEntry::Socket(socket) = &*listener else { return Err(BrokerError::InvalidRights); @@ -465,7 +744,11 @@ pub fn accept( if !socket.listening { return Ok(SocketOutcome::Failed(SocketError::NotConnected)); } - (Arc::clone(&socket.resource), socket.create_request) + ( + Arc::clone(&socket.resource), + socket.create_request, + socket.local_address.ok_or(BrokerError::Internal)?, + ) }; let rights = session .core @@ -482,6 +765,7 @@ pub fn accept( platform_socket: Once::new(), readiness, _quota: quota, + port_reservation: Mutex::new(None), }); let accepted = match listener_resource.accept(resource.readiness.clone()) { Ok(SocketOutcome::Completed(accepted)) => accepted, @@ -495,12 +779,11 @@ pub fn accept( } }; resource.platform_socket.call_once(|| accepted.socket); - let accepted_socket = - SocketObject::new_connected(resource, create_request, accepted.local_address); + let accepted_socket = SocketObject::new_connected(resource, create_request, local_address); let handle = reference.commit(ObjectEntry::Socket(accepted_socket))?; Ok(SocketOutcome::Completed(AcceptedBrokerSocket { handle, - local_address: accepted.local_address, + local_address, remote_address: accepted.remote_address, })) } @@ -756,7 +1039,7 @@ pub fn status(session: &BrokerSession, handle: ObjectHandle) -> Result Result> { + let (local_address, reservation) = + match session + .core + .socket_ports + .reserve(create_request, requested_address, |port| { + session + .core + .tcp_port_mapping_config + .contains_guest_port(port) + })? { + SocketOutcome::Completed(binding) => binding, + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + }; + match resource.bind(local_address)? { + SocketOutcome::Completed(bound_address) if bound_address == local_address => { + Ok(SocketOutcome::Completed((local_address, reservation))) + } + // The platform must echo the broker-reserved guest-local address. + SocketOutcome::Completed(_) => Err(BrokerError::Internal), + SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), + } +} + fn connect_datagram( object: &spin::RwLock, address: SocketAddrV4, @@ -940,19 +1257,52 @@ fn finish_connect(object: &spin::RwLock, status: SocketConnectionSt } } +fn attach_binding( + object: &spin::RwLock, + local_address: SocketAddrV4, + port_reservation: GuestPortReservation, +) { + let mut object = object.write(); + if let ObjectEntry::Socket(socket) = &mut *object { + socket.local_address = Some(local_address); + socket.resource.set_port_reservation(port_reservation); + } +} + fn finish_configuration( object: &spin::RwLock, local_address: Option, + port_reservation: Option, listening: bool, ) { let mut object = object.write(); if let ObjectEntry::Socket(socket) = &mut *object { socket.configuration_in_flight = false; socket.local_address = socket.local_address.or(local_address); + if let Some(port_reservation) = port_reservation { + socket.resource.set_port_reservation(port_reservation); + } socket.listening |= listening; } } +fn finish_retired_configuration( + object: &spin::RwLock, + local_address: Option, + port_reservation: Option, +) { + let mut object = object.write(); + if let ObjectEntry::Socket(socket) = &mut *object { + socket.configuration_in_flight = false; + socket.local_address = socket.local_address.or(local_address); + if let Some(port_reservation) = port_reservation { + socket.resource.set_port_reservation(port_reservation); + } + socket.listening = false; + socket.connection_status = SocketConnectionStatus::Failed(SocketError::Other); + } +} + pub(crate) struct SocketObject { resource: Arc, create_request: CreateSocketRequest, @@ -1004,9 +1354,22 @@ pub(crate) struct SocketResource { platform_socket: Once>, readiness: ReadinessRegistration, _quota: Arc, + port_reservation: Mutex>, } impl SocketResource { + /// Attaches the guest port reservation owning this socket's local port. + /// + /// The resource outlives its object handle, so an operation already in + /// flight keeps the port reserved until it finishes. + fn set_port_reservation(&self, reservation: GuestPortReservation) { + let mut slot = self.port_reservation.lock(); + debug_assert!(slot.is_none()); + if slot.is_none() { + *slot = Some(reservation); + } + } + fn platform_socket(&self) -> &dyn PlatformSocket { self.platform_socket .get() @@ -1025,8 +1388,16 @@ impl SocketResource { self.platform_socket().bind(address) } - fn listen(&self, backlog: u32) -> Result> { - self.platform_socket().listen(backlog) + fn listen( + &self, + backlog: u32, + mapping: Option, + ) -> Result> { + self.platform_socket().listen(backlog, mapping) + } + + fn retire(&self) { + self.platform_socket().retire(); } fn accept( @@ -1105,6 +1476,9 @@ const fn is_udp(request: CreateSocketRequest) -> bool { impl Drop for SocketResource { fn drop(&mut self) { + if let Some(platform_socket) = self.platform_socket.get() { + platform_socket.retire(); + } self.readiness.retire(); } } @@ -1165,9 +1539,107 @@ pub(crate) mod tests { use std::time::Duration; use std::vec; - #[derive(Clone, Default)] + #[derive(Clone)] pub(crate) struct TestSocketProvider { state: Arc, + tcp_port_mappings: Arc>, + } + + impl Default for TestSocketProvider { + fn default() -> Self { + Self { + state: Arc::default(), + tcp_port_mappings: Arc::new(vec![TcpPortMapping { + broker_port: 8080, + guest_port: 80, + }]), + } + } + } + + #[test] + fn guest_tcp_port_namespace_is_broker_wide() { + let ports = BrokerSocketPorts::default(); + let address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + + let SocketOutcome::Completed((_, first_reservation)) = + ports.reserve(create_request(), address, |_| false).unwrap() + else { + panic!("first TCP reservation failed"); + }; + assert!(matches!( + ports.reserve(create_request(), address, |_| false), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + )); + assert!(matches!( + ports.reserve(create_udp_request(), address, |_| false), + Err(BrokerError::Internal) + )); + + drop(first_reservation); + assert!(matches!( + ports.reserve(create_request(), address, |_| false), + Ok(SocketOutcome::Completed(_)) + )); + } + + #[test] + fn implicit_guest_port_allocation_skips_provider_reservations() { + let ports = BrokerSocketPorts::default(); + let requested_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0); + let SocketOutcome::Completed((address, _reservation)) = ports + .reserve(create_request(), requested_address, |port| { + port == FIRST_EPHEMERAL_PORT + }) + .unwrap() + else { + panic!("ephemeral TCP reservation failed"); + }; + + assert_eq!(address.port(), FIRST_EPHEMERAL_PORT + 1); + } + + #[test] + fn tcp_port_mapping_configuration_is_bounded_and_unique() { + let mapping = TcpPortMapping { + broker_port: 8080, + guest_port: 80, + }; + assert!(TcpPortMappingConfig::new(&[mapping]).is_ok()); + assert!(matches!( + TcpPortMappingConfig::new(&[TcpPortMapping { + guest_port: 0, + ..mapping + }]), + Err(BrokerError::Internal) + )); + assert!(matches!( + TcpPortMappingConfig::new(&[TcpPortMapping { + broker_port: 0, + ..mapping + }]), + Err(BrokerError::Internal) + )); + assert!(matches!( + TcpPortMappingConfig::new(&[ + mapping, + TcpPortMapping { + broker_port: 8081, + ..mapping + }, + ]), + Err(BrokerError::Internal) + )); + assert!(matches!( + TcpPortMappingConfig::new(&[ + mapping, + TcpPortMapping { + guest_port: 81, + ..mapping + }, + ]), + Err(BrokerError::Internal) + )); } #[derive(Default)] @@ -1181,9 +1653,16 @@ pub(crate) mod tests { status_block: StdMutex, mpsc::Receiver<()>)>>, binds: StdMutex>, listens: StdMutex>, + listen_mappings: StdMutex>>, listen_block: StdMutex, mpsc::Receiver<()>)>>, + fail_listen: core::sync::atomic::AtomicBool, + error_listen: core::sync::atomic::AtomicBool, + invalid_listen_address: core::sync::atomic::AtomicBool, shutdown_calls: AtomicUsize, + retired_sockets: AtomicUsize, dropped_sockets: AtomicUsize, + retained_platform_sockets: StdMutex>>, + retain_next_socket: core::sync::atomic::AtomicBool, fail_create: core::sync::atomic::AtomicBool, fail_connect: core::sync::atomic::AtomicBool, fail_connect_indeterminate: core::sync::atomic::AtomicBool, @@ -1211,9 +1690,31 @@ pub(crate) mod tests { fn fail_next_shutdown(&self) { self.state.fail_shutdown.store(true, Ordering::Relaxed); } + + fn fail_next_listen(&self) { + self.state.fail_listen.store(true, Ordering::Relaxed); + } + + fn error_next_listen(&self) { + self.state.error_listen.store(true, Ordering::Relaxed); + } + + fn invalidate_next_listen_address(&self) { + self.state + .invalid_listen_address + .store(true, Ordering::Relaxed); + } + + fn retain_next_socket(&self) { + self.state.retain_next_socket.store(true, Ordering::Relaxed); + } } impl SocketProvider for TestSocketProvider { + fn tcp_port_mappings(&self) -> &[TcpPortMapping] { + &self.tcp_port_mappings + } + fn create( &self, session_id: SessionId, @@ -1230,12 +1731,22 @@ pub(crate) mod tests { return Err(BrokerError::OutOfMemory); } *self.state.live_readiness.lock().unwrap() = Some(readiness.clone()); - Ok(Arc::new(TestPlatformSocket { + let socket = Arc::new(TestPlatformSocket { state: Arc::clone(&self.state), readiness, create_request: request, tcp_options: StdMutex::new(TestTcpOptions::default()), - })) + guest_local_address: StdMutex::new(None), + active: core::sync::atomic::AtomicBool::new(true), + }); + if self.state.retain_next_socket.swap(false, Ordering::Relaxed) { + self.state + .retained_platform_sockets + .lock() + .unwrap() + .push(Arc::clone(&socket)); + } + Ok(socket) } fn close_session(&self, session_id: SessionId) { @@ -1248,6 +1759,8 @@ pub(crate) mod tests { readiness: ReadinessRegistration, create_request: CreateSocketRequest, tcp_options: StdMutex, + guest_local_address: StdMutex>, + active: core::sync::atomic::AtomicBool, } #[derive(Default)] @@ -1259,6 +1772,10 @@ pub(crate) mod tests { impl PlatformSocket for TestPlatformSocket { fn bind(&self, address: SocketAddrV4) -> Result> { self.state.binds.lock().unwrap().push(address); + if is_tcp(self.create_request) { + *self.guest_local_address.lock().unwrap() = Some(address); + return Ok(SocketOutcome::Completed(address)); + } let address = if address.port() == 0 { SocketAddrV4::new(*address.ip(), 49152) } else { @@ -1267,17 +1784,40 @@ pub(crate) mod tests { Ok(SocketOutcome::Completed(address)) } - fn listen(&self, backlog: u32) -> Result> { + fn listen( + &self, + backlog: u32, + mapping: Option, + ) -> Result> { self.state.listens.lock().unwrap().push(backlog); + self.state.listen_mappings.lock().unwrap().push(mapping); let listen_block = self.state.listen_block.lock().unwrap().take(); if let Some((started, release)) = listen_block { started.send(()).unwrap(); release.recv_timeout(Duration::from_secs(5)).unwrap(); } - Ok(SocketOutcome::Completed(SocketAddrV4::new( - Ipv4Addr::LOCALHOST, - 49152, - ))) + if self.state.fail_listen.swap(false, Ordering::Relaxed) { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } + if self.state.error_listen.swap(false, Ordering::Relaxed) { + return Err(BrokerError::Internal); + } + let local_address = self + .guest_local_address + .lock() + .unwrap() + .ok_or(BrokerError::Internal)?; + if self + .state + .invalid_listen_address + .swap(false, Ordering::Relaxed) + { + return Ok(SocketOutcome::Completed(SocketAddrV4::new( + *local_address.ip(), + local_address.port().wrapping_add(1), + ))); + } + Ok(SocketOutcome::Completed(local_address)) } fn accept( @@ -1404,6 +1944,12 @@ pub(crate) mod tests { Ok(response) } + fn retire(&self) { + if self.active.swap(false, Ordering::AcqRel) { + self.state.retired_sockets.fetch_add(1, Ordering::Relaxed); + } + } + fn readiness(&self) -> ReadinessFlags { ReadinessFlags::READ | ReadinessFlags::WRITE } @@ -1411,6 +1957,7 @@ pub(crate) mod tests { impl Drop for TestPlatformSocket { fn drop(&mut self) { + self.retire(); self.state.dropped_sockets.fetch_add(1, Ordering::Relaxed); } } @@ -1422,6 +1969,9 @@ pub(crate) mod tests { check_udp_socket_operations(broker, provider); check_concurrent_udp_status_does_not_regress_connection(broker, provider); check_server_socket_operations(broker, provider); + check_tcp_port_mapping_lifecycle(broker, provider); + check_tcp_port_mapping_retirement(broker, provider); + check_invalid_listen_address_retires_mapping(broker, provider); check_failed_listener_shutdown_preserves_state(broker, provider); check_listener_shutdown_does_not_race_listen(broker, provider); check_connect_errors_classify_peer_state(broker, provider); @@ -1493,6 +2043,13 @@ pub(crate) mod tests { connect(&session, handle, loopback_address()), Ok(SocketOutcome::Completed(SocketConnectionStatus::Connecting)) ); + let local_address = *provider + .state + .binds + .lock() + .unwrap() + .last() + .expect("automatic TCP bind was not recorded"); assert_eq!( readiness.published.lock().unwrap().as_slice(), [(handle, ReadinessFlags::WRITE)] @@ -1501,7 +2058,7 @@ pub(crate) mod tests { status(&session, handle), Ok(SocketStatusResponse { status: SocketConnectionStatus::Connected, - local_address: None, + local_address: Some(local_address), pending_error: None, }) ); @@ -1773,11 +2330,13 @@ pub(crate) mod tests { assert_eq!(provider.state.binds.lock().unwrap().len(), binds_before); let requested_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0); - let local_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152); - assert_eq!( - bind(&session, listener, requested_address), - Ok(SocketOutcome::Completed(local_address)) - ); + let SocketOutcome::Completed(local_address) = + bind(&session, listener, requested_address).unwrap() + else { + panic!("guest TCP bind failed"); + }; + assert!(local_address.ip().is_loopback()); + assert_ne!(local_address.port(), 0); assert_eq!( status(&session, listener), Ok(SocketStatusResponse { @@ -1790,6 +2349,10 @@ pub(crate) mod tests { listen(&session, listener, 128), Ok(SocketOutcome::Completed(local_address)) ); + assert_eq!( + provider.state.listen_mappings.lock().unwrap().last(), + Some(&None) + ); assert_eq!( listen(&session, listener, MAX_TCP_LISTEN_BACKLOG + 1), Err(BrokerError::UnsupportedOperation) @@ -1835,17 +2398,226 @@ pub(crate) mod tests { assert_eq!(broker.reserved_sockets.load(Ordering::Relaxed), 0); let auto_bound = create(&session, create_request(), readiness).unwrap(); + let SocketOutcome::Completed(auto_bound_address) = listen(&session, auto_bound, 0).unwrap() + else { + panic!("automatic TCP listen failed"); + }; + assert!(auto_bound_address.ip().is_unspecified()); + assert_ne!(auto_bound_address.port(), 0); assert_eq!( - listen(&session, auto_bound, 0), - Ok(SocketOutcome::Completed(local_address)) + provider.state.binds.lock().unwrap().last(), + Some(&auto_bound_address) ); assert_eq!( - provider.state.binds.lock().unwrap().last(), - Some(&DEFAULT_TCP_LISTEN_ADDRESS) + provider.state.listen_mappings.lock().unwrap().last(), + Some(&Some(TcpPortMapping { + broker_port: auto_bound_address.port(), + guest_port: auto_bound_address.port(), + })) ); session.close_object_reference(auto_bound).unwrap(); } + fn check_tcp_port_mapping_lifecycle(broker: &BrokerCore, provider: &TestSocketProvider) { + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let readiness = Arc::new(TestReadinessSink::default()); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80); + let first = create(&first_session, create_request(), readiness.clone()).unwrap(); + let second = create(&second_session, create_request(), readiness).unwrap(); + assert_eq!( + bind(&first_session, first, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + bind(&second_session, second, guest_address), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + ); + + let mappings_before = provider.state.listen_mappings.lock().unwrap().len(); + provider.fail_next_listen(); + assert_eq!( + listen(&first_session, first, 1), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + ); + assert_eq!( + bind(&second_session, second, guest_address), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + ); + assert_eq!( + listen(&first_session, first, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + let mappings = provider.state.listen_mappings.lock().unwrap(); + let mapping = provider.tcp_port_mappings[0]; + assert_eq!(mappings[mappings_before], Some(mapping)); + assert_eq!(mappings[mappings_before + 1], Some(mapping)); + drop(mappings); + + assert_eq!( + shutdown(&first_session, first, ShutdownMode::StopListening), + Ok(SocketOutcome::Completed(())) + ); + first_session.close_object_reference(first).unwrap(); + assert_eq!( + bind(&second_session, second, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + listen(&second_session, second, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + provider.state.listen_mappings.lock().unwrap().last(), + Some(&Some(mapping)) + ); + second_session.close_object_reference(second).unwrap(); + + let second = create( + &second_session, + create_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + assert_eq!( + bind(&second_session, second, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + + provider.error_next_listen(); + assert_eq!( + listen(&second_session, second, 1), + Err(BrokerError::Internal) + ); + second_session.close_object_reference(second).unwrap(); + + let first = create( + &first_session, + create_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + assert_eq!( + bind(&first_session, first, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + listen(&first_session, first, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + first_session.close_object_reference(first).unwrap(); + } + + fn check_tcp_port_mapping_retirement(broker: &BrokerCore, provider: &TestSocketProvider) { + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80); + let retired_before = provider.state.retired_sockets.load(Ordering::Relaxed); + let dropped_before = provider.state.dropped_sockets.load(Ordering::Relaxed); + + provider.retain_next_socket(); + let first = create( + &first_session, + create_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + assert_eq!( + bind(&first_session, first, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + listen(&first_session, first, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + first_session.close_object_reference(first).unwrap(); + assert_eq!( + provider.state.retired_sockets.load(Ordering::Relaxed), + retired_before + 1 + ); + assert_eq!( + provider.state.dropped_sockets.load(Ordering::Relaxed), + dropped_before + ); + + let second = create( + &second_session, + create_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + assert_eq!( + bind(&second_session, second, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + listen(&second_session, second, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + second_session.close_object_reference(second).unwrap(); + provider + .state + .retained_platform_sockets + .lock() + .unwrap() + .clear(); + } + + fn check_invalid_listen_address_retires_mapping( + broker: &BrokerCore, + provider: &TestSocketProvider, + ) { + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80); + let retired_before = provider.state.retired_sockets.load(Ordering::Relaxed); + let first = create( + &first_session, + create_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + assert_eq!( + bind(&first_session, first, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + provider.invalidate_next_listen_address(); + assert_eq!(listen(&first_session, first, 1), Err(BrokerError::Internal)); + assert_eq!( + provider.state.retired_sockets.load(Ordering::Relaxed), + retired_before + 1 + ); + first_session.close_object_reference(first).unwrap(); + + let second = create( + &second_session, + create_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + assert_eq!( + bind(&second_session, second, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + listen(&second_session, second, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + second_session.close_object_reference(second).unwrap(); + } + fn check_concurrent_udp_status_does_not_regress_connection( broker: &BrokerCore, provider: &TestSocketProvider, @@ -2165,6 +2937,13 @@ pub(crate) mod tests { connect(&session, poisoned, loopback_address()), Err(BrokerError::Internal) ); + let poisoned_local_address = *provider + .state + .binds + .lock() + .unwrap() + .last() + .expect("automatic TCP bind was not recorded"); assert_eq!( connect(&session, poisoned, loopback_address()), Ok(SocketOutcome::Completed(SocketConnectionStatus::Failed( @@ -2175,7 +2954,7 @@ pub(crate) mod tests { status(&session, poisoned), Ok(SocketStatusResponse { status: SocketConnectionStatus::Failed(SocketError::Other), - local_address: None, + local_address: Some(poisoned_local_address), pending_error: None, }) ); @@ -2205,7 +2984,13 @@ pub(crate) mod tests { Ok(SocketOutcome::Completed(SocketConnectionStatus::Connecting)) ); - let local_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152); + let local_address = *provider + .state + .binds + .lock() + .unwrap() + .last() + .expect("automatic TCP bind was not recorded"); *provider.state.status_responses.lock().unwrap() = std::collections::VecDeque::from([ SocketStatusResponse { status: SocketConnectionStatus::Connecting, @@ -2262,8 +3047,15 @@ pub(crate) mod tests { connect(&session, handle, loopback_address()), Ok(SocketOutcome::Completed(SocketConnectionStatus::Connecting)) ); + let guest_local_address = *provider + .state + .binds + .lock() + .unwrap() + .last() + .expect("automatic TCP bind was not recorded"); - let local_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49153); + let platform_local_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49153); provider .state .status_responses @@ -2271,13 +3063,15 @@ pub(crate) mod tests { .unwrap() .push_back(SocketStatusResponse { status: SocketConnectionStatus::Failed(SocketError::TimedOut), - local_address: Some(local_address), + local_address: Some(platform_local_address), pending_error: None, }); + // The guest-local address reserved by the broker outranks any address + // the platform reports for its private host endpoint. let expected = SocketStatusResponse { status: SocketConnectionStatus::Failed(SocketError::TimedOut), - local_address: Some(local_address), + local_address: Some(guest_local_address), pending_error: None, }; assert_eq!(status(&session, handle), Ok(expected)); diff --git a/litebox_broker_host/src/lib.rs b/litebox_broker_host/src/lib.rs index 575cad89e..8ea37931d 100644 --- a/litebox_broker_host/src/lib.rs +++ b/litebox_broker_host/src/lib.rs @@ -766,6 +766,7 @@ mod tests { Ok(Arc::new(TestPlatformSocket { readiness, create_request: request, + local_address: std::sync::Mutex::new(None), })) } @@ -775,6 +776,7 @@ mod tests { struct TestPlatformSocket { readiness: ReadinessRegistration, create_request: CreateSocketRequest, + local_address: std::sync::Mutex>, } impl PlatformSocket for TestPlatformSocket { @@ -787,17 +789,21 @@ mod tests { } else { address }; + *self.local_address.lock().unwrap() = Some(address); Ok(SocketOutcome::Completed(address)) } fn listen( &self, _backlog: u32, + _mapping: Option, ) -> litebox_broker_core::Result> { - Ok(SocketOutcome::Completed(SocketAddrV4::new( - Ipv4Addr::LOCALHOST, - 49152, - ))) + let local_address = self + .local_address + .lock() + .unwrap() + .ok_or(litebox_broker_core::BrokerError::Internal)?; + Ok(SocketOutcome::Completed(local_address)) } fn accept( @@ -891,6 +897,8 @@ mod tests { }) } + fn retire(&self) {} + fn readiness(&self) -> litebox_broker_protocol::readiness::ReadinessFlags { litebox_broker_protocol::readiness::ReadinessFlags::READ | litebox_broker_protocol::readiness::ReadinessFlags::WRITE @@ -1363,7 +1371,7 @@ mod tests { ), BrokerResult::Socket(SocketResponse::Status(SocketStatusResponse { status: SocketConnectionStatus::Connected, - local_address: None, + local_address: Some(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 49152)), pending_error: None, })) ); diff --git a/litebox_broker_platform_linux_userland/src/socket.rs b/litebox_broker_platform_linux_userland/src/socket.rs index 184929c06..96d35c5fc 100644 --- a/litebox_broker_platform_linux_userland/src/socket.rs +++ b/litebox_broker_platform_linux_userland/src/socket.rs @@ -7,25 +7,25 @@ use std::collections::HashMap; use std::fmt; use std::io::{Error, ErrorKind, Result as IoResult}; use std::mem::size_of; -use std::net::SocketAddrV4; +use std::net::{Ipv4Addr, SocketAddrV4}; use std::os::fd::OwnedFd; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::mpsc::{Receiver, SyncSender, TryRecvError, TrySendError, sync_channel}; use std::sync::{Arc, Mutex}; use std::thread::{self, JoinHandle}; -use std::time::Duration; +use std::time::{Duration, Instant}; use litebox_broker_core::socket::{ AcceptedPlatformSocket, PlatformConnectError, PlatformDatagramReceive, PlatformSocket, - PlatformStreamReceive, SocketProvider, + PlatformStreamReceive, SocketProvider, TcpPortMapping, }; use litebox_broker_core::{BrokerError, Result as BrokerResult, SessionId}; use litebox_broker_protocol::readiness::ReadinessFlags; use litebox_broker_protocol::socket::{ AddressFamily, CreateSocketRequest, IpProtocol, MAX_SOCKET_PEEK_SIZE, MAX_SOCKET_TRANSFER_SIZE, - MAX_UDP_DATAGRAM_SIZE, ReceiveFlags, ReceiveFromFlags, SendFlags, ShutdownMode, - SocketConnectionStatus, SocketError, SocketOutcome, SocketStatusResponse, SocketType, - TcpOptionName, TcpOptionValue, + MAX_TCP_LISTEN_BACKLOG, MAX_UDP_DATAGRAM_SIZE, ReceiveFlags, ReceiveFromFlags, SendFlags, + ShutdownMode, SocketConnectionStatus, SocketError, SocketOutcome, SocketStatusResponse, + SocketType, TcpOptionName, TcpOptionValue, }; use rustix::buffer::spare_capacity; use rustix::event::{EventfdFlags, PollFd, PollFlags, Timespec, epoll, eventfd, poll}; @@ -43,6 +43,9 @@ use litebox_broker_core::readiness::ReadinessRegistration; const WAKE_TOKEN: u64 = 0; const MAX_QUEUED_SOCKET_COMMANDS: usize = 64; const MAX_EPOLL_EVENTS: usize = 64; +const MAX_STALE_PORT_MAPPING_CONNECTIONS: usize = MAX_TCP_LISTEN_BACKLOG as usize + 3; +const MAX_TRACKED_GUEST_CONNECTIONS: usize = 1 << 14; +const PENDING_CONNECT_DISCARD_LIFETIME: Duration = Duration::from_mins(5); /// Linux-userland socket provider. /// @@ -51,21 +54,115 @@ const MAX_EPOLL_EVENTS: usize = 64; /// immediate nonblocking operation, never for network readiness. pub struct LinuxSocketProvider { reactor: Arc, + tcp_port_mappings: Vec, +} + +struct PortMappingState { + mapping: TcpPortMapping, + reservation: Option, + reservation_registered: bool, + owned_by: Option, +} + +struct StaleTcpConnection { + session_id: SessionId, + deadline: Option, + retained_connector: Option, } impl LinuxSocketProvider { - /// Starts a provider whose reactor tracks at most `max_sockets` resources. - pub fn new(max_sockets: usize) -> IoResult { + /// Starts a provider with matching global and per-session socket limits. + pub fn new(max_sockets: usize, max_sockets_per_session: usize) -> IoResult { + Self::new_with_tcp_port_mappings( + max_sockets, + max_sockets_per_session, + Ipv4Addr::UNSPECIFIED, + &[], + ) + } + + /// Starts a limited provider with broker-wide TCP port mapping overrides. + pub fn new_with_tcp_port_mappings( + max_sockets: usize, + max_sockets_per_session: usize, + broker_ipv4_address: Ipv4Addr, + tcp_port_mappings: &[TcpPortMapping], + ) -> IoResult { + for (index, mapping) in tcp_port_mappings.iter().enumerate() { + if mapping.guest_port == 0 || mapping.broker_port == 0 { + return Err(Error::new( + ErrorKind::InvalidInput, + "mapped broker and guest ports must be nonzero", + )); + } + if tcp_port_mappings[..index].iter().any(|existing| { + existing.guest_port == mapping.guest_port + || existing.broker_port == mapping.broker_port + }) { + return Err(Error::new( + ErrorKind::InvalidInput, + "mapped broker and guest ports must be unique", + )); + } + } + let tcp_port_mappings = tcp_port_mappings.to_vec(); Ok(Self { - reactor: Arc::new(ReactorClient::start(max_sockets)?), + reactor: Arc::new(ReactorClient::start( + max_sockets, + max_sockets_per_session, + broker_ipv4_address, + tcp_port_mappings.clone(), + )?), + tcp_port_mappings, }) } } +fn create_port_mapping_reservation( + broker_ipv4_address: Ipv4Addr, + mapping: TcpPortMapping, + reuse_address: bool, + reuse_port: bool, +) -> core::result::Result { + let socket = socket_with( + LinuxAddressFamily::INET, + LinuxSocketType::STREAM, + LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, + Some(ipproto::TCP), + )?; + if reuse_address { + sockopt::set_socket_reuseaddr(&socket, true)?; + } + if reuse_port { + sockopt::set_socket_reuseport(&socket, true)?; + } + bind( + &socket, + &SocketAddrV4::new(broker_ipv4_address, mapping.broker_port), + )?; + Ok(socket) +} + +fn create_replacement_port_mapping_reservation( + broker_ipv4_address: Ipv4Addr, + mapping: TcpPortMapping, +) -> core::result::Result { + let socket = create_port_mapping_reservation(broker_ipv4_address, mapping, true, false)?; + // Reuse is needed only while replacing a listener that may still have + // accepted children. Clear it afterward so the bound reservation is + // exclusive between mapped listeners. + sockopt::set_socket_reuseaddr(&socket, false)?; + Ok(socket) +} + impl SocketProvider for LinuxSocketProvider { + fn tcp_port_mappings(&self) -> &[TcpPortMapping] { + &self.tcp_port_mappings + } + fn create( &self, - _session_id: SessionId, + session_id: SessionId, request: CreateSocketRequest, readiness: ReadinessRegistration, ) -> BrokerResult> { @@ -86,6 +183,7 @@ impl SocketProvider for LinuxSocketProvider { }); self.reactor.request(|response| ReactorCommand::Create { id, + session_id, request, readiness, snapshot, @@ -95,7 +193,9 @@ impl SocketProvider for LinuxSocketProvider { Ok(socket) } - fn close_session(&self, _session_id: SessionId) {} + fn close_session(&self, session_id: SessionId) { + self.reactor.close_session(session_id); + } } /// Broker-core-facing handle for a reactor-owned socket. @@ -118,10 +218,15 @@ impl PlatformSocket for LinuxSocket { }) } - fn listen(&self, backlog: u32) -> BrokerResult> { + fn listen( + &self, + backlog: u32, + mapping: Option, + ) -> BrokerResult> { self.reactor.request(|response| ReactorCommand::Listen { id: self.id, backlog, + mapping, response, }) } @@ -150,7 +255,6 @@ impl PlatformSocket for LinuxSocket { SocketOutcome::Completed(accepted) => { Ok(SocketOutcome::Completed(AcceptedPlatformSocket { socket, - local_address: accepted.local_address, remote_address: accepted.remote_address, })) } @@ -275,6 +379,12 @@ impl PlatformSocket for LinuxSocket { }) } + fn retire(&self) { + if self.active.swap(false, Ordering::AcqRel) { + self.reactor.close_socket(self.id); + } + } + fn readiness(&self) -> ReadinessFlags { self.snapshot .lock() @@ -285,9 +395,7 @@ impl PlatformSocket for LinuxSocket { impl Drop for LinuxSocket { fn drop(&mut self) { - if self.active.swap(false, Ordering::AcqRel) { - self.reactor.close_socket(self.id); - } + self.retire(); } } @@ -300,9 +408,32 @@ struct ReactorClient { } impl ReactorClient { - fn start(max_sockets: usize) -> IoResult { + fn start( + max_sockets: usize, + max_sockets_per_session: usize, + broker_ipv4_address: Ipv4Addr, + tcp_port_mappings: Vec, + ) -> IoResult { let epoll_fd = epoll::create(epoll::CreateFlags::CLOEXEC)?; let wake = Arc::new(eventfd(0, EventfdFlags::CLOEXEC | EventfdFlags::NONBLOCK)?); + let port_mappings = tcp_port_mappings + .into_iter() + .map(|mapping| PortMappingState { + mapping, + reservation: None, + reservation_registered: false, + owned_by: None, + }) + .collect::>(); + let mut tcp = BrokerTcpState::default(); + tcp.stale_mapped_connections + .try_reserve(MAX_TRACKED_GUEST_CONNECTIONS) + .map_err(|_| { + Error::new( + ErrorKind::OutOfMemory, + "stale guest connection table allocation failed", + ) + })?; epoll::add( &epoll_fd, wake.as_ref(), @@ -328,9 +459,15 @@ impl ReactorClient { let mut reactor = Reactor { epoll: epoll_fd, wake: reactor_wake, + broker_ipv4_address, commands: receiver, sockets, + tcp, + sessions: HashMap::new(), + port_mappings, max_sockets, + max_sockets_per_session, + retained_connectors: 0, peek_cache: None, events, }; @@ -382,7 +519,9 @@ impl ReactorClient { Err(TrySendError::Full(_)) => return Err(BrokerError::ResourceExhausted), Err(TrySendError::Disconnected(_)) => return Err(BrokerError::Internal), } - self.signal().map_err(|_| BrokerError::Internal)?; + // Once queued, wait for acknowledgement even if the wake write fails: + // the reactor may still execute the command after another event. + let _ = self.signal(); receive.recv().map_err(|_| BrokerError::Internal)? } @@ -410,8 +549,9 @@ impl ReactorClient { )); } } - self.signal() - .map_err(|_| PlatformConnectError::PeerIndeterminate(BrokerError::Internal))?; + // A queued connect has indeterminate peer state until the reactor + // acknowledges it, regardless of whether this wake write succeeds. + let _ = self.signal(); receive .recv() .map_err(|_| PlatformConnectError::PeerIndeterminate(BrokerError::Internal))? @@ -426,8 +566,68 @@ impl ReactorClient { { return; } - // Do not release the core's socket quota until the reactor has dropped - // the descriptor, even if the redundant wake write fails. + // Wait until the reactor has dropped the descriptor or transferred it + // into retained-connector accounting, even if the wake write fails. + let _ = self.signal(); + let _ = receive.recv(); + } + + #[cfg(test)] + fn host_address(&self, kind: SocketKind, guest_port: u16) -> Option { + let (response, receive) = sync_channel(1); + self.commands + .send(ReactorCommand::HostAddress { + kind, + guest_port, + response, + }) + .unwrap(); + self.signal().unwrap(); + receive.recv().unwrap() + } + + #[cfg(test)] + fn pending_guest_connection_count(&self) -> usize { + let (response, receive) = sync_channel(1); + self.commands + .send(ReactorCommand::PendingGuestConnectionCount { response }) + .unwrap(); + self.signal().unwrap(); + receive.recv().unwrap() + } + + #[cfg(test)] + fn stale_guest_connection_count(&self) -> usize { + let (response, receive) = sync_channel(1); + self.commands + .send(ReactorCommand::StaleGuestConnectionCount { response }) + .unwrap(); + self.signal().unwrap(); + receive.recv().unwrap() + } + + #[cfg(test)] + fn retained_connector_count(&self) -> usize { + let (response, receive) = sync_channel(1); + self.commands + .send(ReactorCommand::RetainedConnectorCount { response }) + .unwrap(); + self.signal().unwrap(); + receive.recv().unwrap() + } + + fn close_session(&self, session_id: SessionId) { + let (response, receive) = sync_channel(1); + if self + .commands + .send(ReactorCommand::CloseSession { + session_id, + response, + }) + .is_err() + { + return; + } let _ = self.signal(); let _ = receive.recv(); } @@ -479,6 +679,7 @@ impl Drop for ReactorClient { enum ReactorCommand { Create { id: u64, + session_id: SessionId, request: CreateSocketRequest, readiness: ReadinessRegistration, snapshot: Arc>, @@ -498,6 +699,7 @@ enum ReactorCommand { Listen { id: u64, backlog: u32, + mapping: Option, response: SyncSender>>, }, Accept { @@ -556,6 +758,28 @@ enum ReactorCommand { id: u64, response: SyncSender<()>, }, + CloseSession { + session_id: SessionId, + response: SyncSender<()>, + }, + #[cfg(test)] + HostAddress { + kind: SocketKind, + guest_port: u16, + response: SyncSender>, + }, + #[cfg(test)] + PendingGuestConnectionCount { + response: SyncSender, + }, + #[cfg(test)] + StaleGuestConnectionCount { + response: SyncSender, + }, + #[cfg(test)] + RetainedConnectorCount { + response: SyncSender, + }, Stop { response: SyncSender<()>, }, @@ -578,17 +802,40 @@ enum ReactorReceiveFromOutcome { } struct AcceptedEndpoints { - local_address: SocketAddrV4, remote_address: SocketAddrV4, } +enum AcceptedTcpPeer { + Guest(SocketAddrV4), + Native(SocketAddrV4), + Stale, +} + +#[derive(Clone, Copy)] +enum PendingGuestConnectionDisposition { + Retain, + Discard(Option), +} + +struct RetiredTcpConnector { + connection: (SocketAddrV4, SocketAddrV4), + mapping_index: Option, + unplaced_connector: Option, +} + /// State owned and accessed exclusively by the socket reactor thread. struct Reactor { epoll: OwnedFd, wake: Arc, + broker_ipv4_address: Ipv4Addr, commands: Receiver, sockets: HashMap, + tcp: BrokerTcpState, + sessions: HashMap, + port_mappings: Vec, max_sockets: usize, + max_sockets_per_session: usize, + retained_connectors: usize, peek_cache: Option, events: Vec, } @@ -600,8 +847,10 @@ struct PeekCache { } /// Reactor-owned descriptor and its broker-facing readiness state. +#[allow(clippy::struct_excessive_bools)] // These are independent cached kernel attributes. struct SocketEntry { socket: OwnedFd, + session_id: SessionId, kind: SocketKind, readiness: ReadinessRegistration, snapshot: Arc>, @@ -609,6 +858,53 @@ struct SocketEntry { write_shutdown: bool, peek_waitall_threshold: Option, listening: bool, + abortive_close: bool, + guest_local_address: Option, + port_mapping_index: Option, + mapping_fallback_socket: Option, + tcp_no_delay: bool, + tcp_keep_alive: bool, +} + +/// The broker-wide guest TCP namespace and pending guest-to-guest connections. +#[derive(Default)] +struct BrokerTcpState { + bindings: HashMap, + pending_guest_connections: HashMap<(SocketAddrV4, SocketAddrV4), PendingGuestTcpConnection>, + stale_mapped_connections: HashMap<(usize, SocketAddrV4, SocketAddrV4), StaleTcpConnection>, +} + +/// Per-session socket ownership, quota, and teardown state. +/// +/// Sessions do not define network namespaces; all guest TCP endpoints and +/// pending guest-to-guest connections live in the reactor-wide `BrokerTcpState`. +#[derive(Default)] +struct SessionSocketState { + live_sockets: usize, + pending_guest_connections: usize, + retained_connectors: usize, + closing: bool, +} + +struct PendingGuestTcpConnection { + session_id: SessionId, + guest_address: SocketAddrV4, + listener_id: u64, + discard_on_accept: bool, + discard_deadline: Option, + // Delay an established abort until accept can identify and discard its peer. + retained_connector: Option, +} + +#[derive(Clone, Copy)] +struct GuestPortBinding { + socket_id: u64, + guest_address: SocketAddrV4, + host_address: Option, + host_peer_address: Option, + host_peer_mapping_index: Option, + host_mapped: bool, + listening: bool, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -617,6 +913,146 @@ enum SocketKind { Udp, } +impl BrokerTcpState { + fn insert_binding(&mut self, port: u16, binding: GuestPortBinding) -> BrokerResult<()> { + if binding.guest_address.port() != port { + return Err(BrokerError::Internal); + } + if self.bindings.insert(port, binding).is_some() { + return Err(BrokerError::Internal); + } + Ok(()) + } + + fn reserve_binding(&mut self) -> BrokerResult<()> { + self.bindings + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory) + } + + fn remove_binding(&mut self, port: u16, socket_id: u64) { + if self + .bindings + .get(&port) + .is_some_and(|binding| binding.socket_id == socket_id) + { + self.bindings.remove(&port); + } + } + + fn guest_binding(&self, address: SocketAddrV4) -> Option { + if !address.ip().is_loopback() { + return None; + } + let binding = self.bindings.get(&address.port())?; + if binding.guest_address.ip().is_unspecified() || binding.guest_address.ip() == address.ip() + { + Some(*binding) + } else { + None + } + } + + fn set_host_address( + &mut self, + port: u16, + socket_id: u64, + host_address: SocketAddrV4, + host_mapped: bool, + ) -> BrokerResult<()> { + let binding = self.bindings.get_mut(&port).ok_or(BrokerError::Internal)?; + if binding.socket_id != socket_id { + return Err(BrokerError::Internal); + } + binding.host_address = Some(host_address); + binding.host_mapped = host_mapped; + Ok(()) + } + + fn set_host_peer_address( + &mut self, + port: u16, + socket_id: u64, + host_peer_address: SocketAddrV4, + mapping_index: Option, + ) -> BrokerResult<(SocketAddrV4, SocketAddrV4)> { + let (host_address, guest_address) = { + let binding = self.bindings.get_mut(&port).ok_or(BrokerError::Internal)?; + if binding.socket_id != socket_id { + return Err(BrokerError::Internal); + } + binding.host_peer_address = Some(host_peer_address); + binding.host_peer_mapping_index = mapping_index; + ( + binding.host_address.ok_or(BrokerError::Internal)?, + binding.guest_address, + ) + }; + Ok((host_address, guest_address)) + } + + fn clear_host_address(&mut self, port: u16, socket_id: u64) -> BrokerResult<()> { + let binding = self.bindings.get_mut(&port).ok_or(BrokerError::Internal)?; + if binding.socket_id != socket_id { + return Err(BrokerError::Internal); + } + binding.host_address = None; + binding.host_peer_address = None; + binding.host_peer_mapping_index = None; + binding.host_mapped = false; + binding.listening = false; + Ok(()) + } + + fn mark_listening(&mut self, port: u16, socket_id: u64) -> BrokerResult<()> { + let binding = self.bindings.get_mut(&port).ok_or(BrokerError::Internal)?; + if binding.socket_id != socket_id || binding.host_address.is_none() { + return Err(BrokerError::Internal); + } + binding.listening = true; + Ok(()) + } + + fn take_connector_connection( + &mut self, + port: u16, + socket_id: u64, + ) -> Option<((SocketAddrV4, SocketAddrV4), Option)> { + self.bindings.get_mut(&port).and_then(|binding| { + (binding.socket_id == socket_id) + .then(|| { + binding + .host_address + .zip(binding.host_peer_address.take()) + .map(|connection| (connection, binding.host_peer_mapping_index.take())) + }) + .flatten() + }) + } + + fn take_pending_guest_connection( + &mut self, + remote_address: SocketAddrV4, + local_address: SocketAddrV4, + ) -> Option { + self.pending_guest_connections + .remove(&(remote_address, local_address)) + .or_else(|| { + self.pending_guest_connections.remove(&( + remote_address, + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, local_address.port()), + )) + }) + } +} + +fn retain_session_state(state: &SessionSocketState) -> bool { + !state.closing + || state.live_sockets != 0 + || state.pending_guest_connections != 0 + || state.retained_connectors != 0 +} + /// Cached connection and readiness state shared with the broker-facing handle. /// /// The reactor updates this snapshot whenever kernel state changes, allowing @@ -660,955 +1096,2835 @@ impl fmt::Display for ReactorFailure { } impl Reactor { - fn run(&mut self) -> core::result::Result<(), ReactorFailure> { - loop { - let mut events = core::mem::take(&mut self.events); - events.clear(); - match epoll::wait(&self.epoll, spare_capacity(&mut events), None) { - Ok(_) => {} - Err(Errno::INTR) => { - self.events = events; - continue; - } - Err(error) => return Err(ReactorFailure::Io(error)), - } + fn reserve_pending_guest_connection(&mut self, session_id: SessionId) -> BrokerResult<()> { + let session = self + .sessions + .get(&session_id) + .ok_or(BrokerError::Internal)?; + let session_stale = + count_session_stale_connections(self.tcp.stale_mapped_connections.values(), session_id); + if session + .pending_guest_connections + .checked_add(session_stale) + .is_none_or(|count| count >= self.max_sockets_per_session) + || self + .tcp + .pending_guest_connections + .len() + .checked_add(self.tcp.stale_mapped_connections.len()) + .is_none_or(|count| count >= MAX_TRACKED_GUEST_CONNECTIONS) + { + return Err(BrokerError::ResourceExhausted); + } + self.tcp + .pending_guest_connections + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory) + } - // Apply readiness observed by this wait before commands. A command - // that then reaches EAGAIN records the newer authoritative state. - let mut wake = false; - for event in events.drain(..) { - let id = event.data.u64(); - if id == WAKE_TOKEN { - wake = true; - } else if let Some(socket) = self.sockets.get_mut(&id) { - handle_socket_event(socket, event.flags).map_err(ReactorFailure::Broker)?; + fn decrement_pending_guest_connection_count(&mut self, session_id: SessionId) { + let session = self + .sessions + .get_mut(&session_id) + .expect("pending guest connection session state missing"); + session.pending_guest_connections = session + .pending_guest_connections + .checked_sub(1) + .expect("session pending guest connection count underflow"); + } + + fn release_pending_guest_connection_connector( + &mut self, + connection: &PendingGuestTcpConnection, + ) { + if connection.retained_connector.is_none() { + return; + } + self.retained_connectors = self + .retained_connectors + .checked_sub(1) + .expect("reactor retained connector count underflow"); + let session = self + .sessions + .get_mut(&connection.session_id) + .expect("pending guest connection session state missing"); + session.retained_connectors = session + .retained_connectors + .checked_sub(1) + .expect("session retained connector count underflow"); + } + + fn finish_removed_pending_guest_connection(&mut self, connection: &PendingGuestTcpConnection) { + self.decrement_pending_guest_connection_count(connection.session_id); + self.release_pending_guest_connection_connector(connection); + } + + fn insert_pending_guest_connection( + &mut self, + session_id: SessionId, + connection: (SocketAddrV4, SocketAddrV4), + guest_address: SocketAddrV4, + listener_id: u64, + ) -> BrokerResult<()> { + if let Some(previous) = self.tcp.pending_guest_connections.remove(&connection) { + self.finish_removed_pending_guest_connection(&previous); + } + let session = self + .sessions + .get_mut(&session_id) + .ok_or(BrokerError::Internal)?; + session.pending_guest_connections = session + .pending_guest_connections + .checked_add(1) + .ok_or(BrokerError::ResourceExhausted)?; + self.tcp.pending_guest_connections.insert( + connection, + PendingGuestTcpConnection { + session_id, + guest_address, + listener_id, + discard_on_accept: false, + discard_deadline: None, + retained_connector: None, + }, + ); + Ok(()) + } + + fn retire_pending_guest_connection_for_connector( + &mut self, + session_id: SessionId, + guest_port: u16, + socket_id: u64, + discard_on_accept: bool, + discard_deadline: Option, + mut retained_connector: Option, + ) -> Option { + let (connection, mapping_index) = + self.tcp.take_connector_connection(guest_port, socket_id)?; + let mut retained_in_pending_connection = false; + if let Some(pending) = self.tcp.pending_guest_connections.get_mut(&connection) + && pending.session_id == session_id + { + pending.discard_on_accept = discard_on_accept; + pending.discard_deadline = discard_deadline; + retained_in_pending_connection = retained_connector.is_some(); + pending.retained_connector = retained_connector.take(); + } + if retained_in_pending_connection { + self.retained_connectors = self + .retained_connectors + .checked_add(1) + .expect("reactor retained connector count overflow"); + let session = self + .sessions + .get_mut(&session_id) + .expect("pending guest connection session state missing"); + session.retained_connectors = session + .retained_connectors + .checked_add(1) + .expect("session retained connector count overflow"); + } + Some(RetiredTcpConnector { + connection, + mapping_index, + unplaced_connector: retained_connector, + }) + } + + fn owned_tcp_port_mapping(&self, id: u64) -> Option { + let socket = self.sockets.get(&id)?; + let mapping_index = socket.port_mapping_index?; + self.port_mappings + .get(mapping_index) + .is_some_and(|state| state.owned_by == Some(id)) + .then_some(mapping_index) + } + + /// Transfers one reserved broker endpoint to the socket using its mapping. + /// + /// The reservation descriptor keeps the mapped endpoint exclusively owned + /// between mapped listeners, so ownership replaces the socket's descriptor + /// rather than binding the endpoint again. + fn realize_tcp_port_mapping( + &mut self, + id: u64, + mapping_index: usize, + ) -> BrokerResult> { + let guest_port = self + .sockets + .get(&id) + .and_then(|socket| socket.guest_local_address) + .ok_or(BrokerError::Internal)? + .port(); + let mapping = self + .port_mappings + .get(mapping_index) + .ok_or(BrokerError::Internal)? + .mapping; + if mapping.guest_port != guest_port { + return Err(BrokerError::Internal); + } + if self + .port_mappings + .get(mapping_index) + .is_some_and(|state| state.owned_by.is_none() && state.reservation.is_none()) + { + match create_replacement_port_mapping_reservation(self.broker_ipv4_address, mapping) { + Ok(reservation) => { + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + state.reservation = Some(reservation); + state.reservation_registered = false; } - } - self.events = events; - if wake { - self.drain_wake()?; - if self.process_commands() { - return Ok(()); + Err(error) => { + return Ok(SocketOutcome::Failed(socket_operation_error_from_errno( + error, + )?)); } } } - } - - fn drain_wake(&self) -> core::result::Result<(), ReactorFailure> { - let mut value = [0_u8; size_of::()]; - loop { - match read(self.wake.as_ref(), &mut value) { - Ok(length) if length == value.len() => return Ok(()), - Ok(_) => return Err(ReactorFailure::Io(Errno::IO)), - Err(Errno::INTR) => {} - Err(Errno::AGAIN) => return Ok(()), - Err(error) => return Err(ReactorFailure::Io(error)), + let (reservation, reservation_registered) = { + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + if state.owned_by.is_some() { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } + let Some(reservation) = state.reservation.take() else { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + }; + let reservation_registered = state.reservation_registered; + state.reservation_registered = false; + (reservation, reservation_registered) + }; + let preparation = sockopt::set_socket_linger(&reservation, None) + .map_err(broker_error_from_errno) + .and_then(|()| self.drain_tcp_listener(&reservation, mapping_index)); + let stale_state_drained = match preparation { + Ok(drained) => drained, + Err(error) => { + self.port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal) + .map(|state| { + state.reservation = Some(reservation); + state.reservation_registered = reservation_registered; + })?; + return Err(error); + } + }; + if !stale_state_drained { + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + state.reservation = Some(reservation); + state.reservation_registered = reservation_registered; + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } + let host_address = match local_socket_address(&reservation) { + Ok(address) => address, + Err(error) => { + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + state.reservation = Some(reservation); + state.reservation_registered = reservation_registered; + return Err(error); } + }; + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + if socket.kind != SocketKind::Tcp || socket.mapping_fallback_socket.is_some() { + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + state.reservation = Some(reservation); + state.reservation_registered = reservation_registered; + return Err(BrokerError::Internal); + } + if let Err(error) = + apply_tcp_options(&reservation, socket.tcp_no_delay, socket.tcp_keep_alive) + { + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + state.reservation = Some(reservation); + state.reservation_registered = reservation_registered; + return Err(error); + } + let registration = if reservation_registered { + epoll::modify( + &self.epoll, + &reservation, + epoll::EventData::new_u64(id), + idle_epoll_events(), + ) + } else { + epoll::add( + &self.epoll, + &reservation, + epoll::EventData::new_u64(id), + idle_epoll_events(), + ) + }; + if let Err(error) = registration { + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + state.reservation = Some(reservation); + state.reservation_registered = reservation_registered; + return Err(broker_error_from_errno(error)); } + let fallback = core::mem::replace(&mut socket.socket, reservation); + socket.mapping_fallback_socket = Some(fallback); + socket.port_mapping_index = Some(mapping_index); + self.port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)? + .owned_by = Some(id); + Ok(SocketOutcome::Completed(host_address)) } - fn process_commands(&mut self) -> bool { - for _ in 0..MAX_QUEUED_SOCKET_COMMANDS { - let command = match self.commands.try_recv() { - Ok(command) => command, - Err(TryRecvError::Empty) => return false, - Err(TryRecvError::Disconnected) => return true, - }; - match command { - ReactorCommand::Create { - id, - request, - readiness, - snapshot, - active, - response, - } => { - let outcome = self.create_socket(id, request, readiness, snapshot); - let created = outcome.is_ok(); - if created { - active.store(true, Ordering::Release); - } - if response.send(outcome).is_err() && created { - self.sockets.remove(&id); - } - } - ReactorCommand::Connect { - id, - address, - response, - } => { - let outcome = match self.sockets.get_mut(&id) { - Some(socket) => connect_socket(&self.epoll, id, socket, address), - None => Err(PlatformConnectError::PeerIndeterminate( - BrokerError::Internal, - )), - }; - let _ = response.send(outcome); - } - ReactorCommand::Bind { - id, - address, - response, - } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| bind_socket(socket, address)); - let _ = response.send(outcome); - } - ReactorCommand::Listen { - id, - backlog, - response, - } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| listen_socket(&self.epoll, id, socket, backlog)); - let _ = response.send(outcome); - } - ReactorCommand::Accept { - listener_id, - accepted_id, - readiness, - snapshot, - active, - response, - } => { - let outcome = self.accept_socket(listener_id, accepted_id, readiness, snapshot); - let accepted = matches!( - &outcome, - Ok(SocketOutcome::Completed(AcceptedEndpoints { .. })) - ); - if accepted { - active.store(true, Ordering::Release); - } - if response.send(outcome).is_err() && accepted { - self.sockets.remove(&accepted_id); - } - } - ReactorCommand::Send { id, data, response } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| send_socket(socket, &data)); - let _ = response.send(outcome); - } - ReactorCommand::SendTo { - id, - data, - destination, - response, - } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| send_to_socket(socket, &data, destination)); - let _ = response.send(outcome); - } - ReactorCommand::Receive { - id, - length, - flags, - peek_offset, - peek_length, - response, - } => { - let outcome = match self.sockets.get_mut(&id) { - Some(socket) => receive_socket( - socket, - &mut self.peek_cache, - id, - length, - flags, - peek_offset, - peek_length, - ), - None => Err(BrokerError::Internal), - }; - let _ = response.send(outcome); - } - ReactorCommand::ReceiveFrom { - id, - length, - flags, - response, - } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| receive_from_socket(socket, length, flags)); - let _ = response.send(outcome); - } - ReactorCommand::Shutdown { id, mode, response } => { - if self - .peek_cache - .as_ref() - .is_some_and(|cache| cache.socket_id == id) - { - self.peek_cache = None; - } - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| shutdown_socket(socket, mode)); - let _ = response.send(outcome); - } - ReactorCommand::SetTcpOption { - id, - value, - response, - } => { - let outcome = self - .sockets - .get(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| set_tcp_option(socket, value)); - let _ = response.send(outcome); - } - ReactorCommand::GetTcpOption { id, name, response } => { - let outcome = self - .sockets - .get(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| get_tcp_option(socket, name)); - let _ = response.send(outcome); - } - ReactorCommand::Status { id, response } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(status_socket); - let _ = response.send(outcome); - } - ReactorCommand::Close { id, response } => { - if self - .peek_cache - .as_ref() - .is_some_and(|cache| cache.socket_id == id) - { - self.peek_cache = None; + /// Binds an unconfigured identity mapping at listen time. + fn realize_identity_tcp_mapping( + &mut self, + id: u64, + mapping: TcpPortMapping, + ) -> BrokerResult> { + let socket = self.sockets.get(&id).ok_or(BrokerError::Internal)?; + let guest_port = socket + .guest_local_address + .ok_or(BrokerError::Internal)? + .port(); + if socket.kind != SocketKind::Tcp + || mapping.guest_port != guest_port + || mapping.broker_port != guest_port + || socket.mapping_fallback_socket.is_some() + || socket.port_mapping_index.is_some() + { + return Err(BrokerError::Internal); + } + let replacement = match create_port_mapping_reservation( + self.broker_ipv4_address, + mapping, + false, + false, + ) { + Ok(socket) => socket, + Err(error) => { + return Ok(SocketOutcome::Failed(socket_operation_error_from_errno( + error, + )?)); + } + }; + apply_tcp_options(&replacement, socket.tcp_no_delay, socket.tcp_keep_alive)?; + epoll::add( + &self.epoll, + &replacement, + epoll::EventData::new_u64(id), + idle_epoll_events(), + ) + .map_err(broker_error_from_errno)?; + let host_address = local_socket_address(&replacement)?; + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + let fallback = core::mem::replace(&mut socket.socket, replacement); + socket.mapping_fallback_socket = Some(fallback); + Ok(SocketOutcome::Completed(host_address)) + } + + fn release_failed_identity_tcp_mapping(&mut self, id: u64) -> BrokerResult<()> { + let guest_port = self + .sockets + .get(&id) + .and_then(|socket| { + (!socket.listening + && socket.port_mapping_index.is_none() + && socket.mapping_fallback_socket.is_some()) + .then_some(())?; + socket.guest_local_address.map(|address| address.port()) + }) + .ok_or(BrokerError::Internal)?; + let host_address_is_set = self + .tcp + .bindings + .get(&guest_port) + .filter(|binding| binding.socket_id == id) + .map(|binding| binding.host_mapped); + let replacement = { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + let fallback = socket + .mapping_fallback_socket + .take() + .ok_or(BrokerError::Internal)?; + core::mem::replace(&mut socket.socket, fallback) + }; + drop(replacement); + let Some(host_address_is_set) = host_address_is_set else { + return Err(BrokerError::Internal); + }; + if host_address_is_set { + self.tcp.clear_host_address(guest_port, id)?; + } + Ok(()) + } + + fn drain_tcp_listener(&mut self, socket: &OwnedFd, mapping_index: usize) -> BrokerResult { + for _ in 0..MAX_STALE_PORT_MAPPING_CONNECTIONS { + match acceptfrom_with( + socket, + LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, + ) { + Ok((accepted, remote_address)) => { + if let Some(remote_address) = remote_address { + let remote_address = SocketAddrV4::try_from(remote_address) + .map_err(|_| BrokerError::Internal)?; + let local_address = local_socket_address(&accepted)?; + self.take_stale_tcp_connection( + mapping_index, + remote_address, + local_address, + ); + self.remove_pending_guest_connection_except( + None, + remote_address, + local_address, + ); } - self.sockets.remove(&id); - let _ = response.send(()); + drop(accepted); } - ReactorCommand::Stop { response } => { - self.sockets.clear(); - let _ = response.send(()); - return true; + Err(Errno::INTR) => {} + Err(Errno::AGAIN | Errno::INVAL) => return Ok(true), + Err(error) => { + // Linux reports pending per-connection network errors from + // accept. Skip stale connection failures while preserving + // broker resource failures. + let _ = socket_operation_error_from_errno(error)?; } } } - false + Ok(false) } - fn create_socket( - &mut self, - id: u64, - request: CreateSocketRequest, - readiness: ReadinessRegistration, - snapshot: Arc>, - ) -> BrokerResult<()> { - if self.sockets.len() >= self.max_sockets { - return Err(BrokerError::ResourceExhausted); + fn stop_listening_socket(&mut self, id: u64) -> BrokerResult> { + let owned_mapping = { + let socket = self.sockets.get(&id).ok_or(BrokerError::Internal)?; + if socket.kind != SocketKind::Tcp || !socket.listening { + return self + .sockets + .get_mut(&id) + .ok_or(BrokerError::Internal) + .and_then(|socket| shutdown_socket(socket, ShutdownMode::StopListening)); + } + self.owned_tcp_port_mapping(id) + }; + let Some(mapping_index) = owned_mapping else { + let socket = self.sockets.get(&id).ok_or(BrokerError::Internal)?; + if socket.port_mapping_index.is_some() { + return Err(BrokerError::Internal); + } + return self + .sockets + .get_mut(&id) + .ok_or(BrokerError::Internal) + .and_then(|socket| shutdown_socket(socket, ShutdownMode::StopListening)); + }; + if !self + .port_mappings + .get(mapping_index) + .is_some_and(|state| state.owned_by == Some(id) && state.reservation.is_none()) + { + return Err(BrokerError::Internal); } - let kind = socket_kind(request).ok_or(BrokerError::Internal)?; - if self.sockets.contains_key(&id) { + let (guest_port, no_delay, keep_alive) = self + .sockets + .get(&id) + .filter(|socket| socket.mapping_fallback_socket.is_none()) + .and_then(|socket| { + socket.guest_local_address.map(|guest_address| { + ( + guest_address.port(), + socket.tcp_no_delay, + socket.tcp_keep_alive, + ) + }) + }) + .ok_or(BrokerError::Internal)?; + if !self + .tcp + .bindings + .get(&guest_port) + .is_some_and(|binding| binding.socket_id == id && binding.host_mapped) + { return Err(BrokerError::Internal); } - let (linux_type, protocol, epoll_events, initial_readiness) = match kind { - SocketKind::Tcp => ( - LinuxSocketType::STREAM, - ipproto::TCP, - idle_epoll_events(), - ReadinessFlags::default(), - ), - SocketKind::Udp => ( - LinuxSocketType::DGRAM, - ipproto::UDP, - active_epoll_events(), - ReadinessFlags::WRITE, - ), - }; - let socket = socket_with( + let replacement = socket_with( LinuxAddressFamily::INET, - linux_type, + LinuxSocketType::STREAM, LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, - Some(protocol), + Some(ipproto::TCP), ) .map_err(broker_error_from_errno)?; + apply_tcp_options(&replacement, no_delay, keep_alive)?; epoll::add( &self.epoll, - &socket, + &replacement, epoll::EventData::new_u64(id), - epoll_events, + idle_epoll_events(), ) .map_err(broker_error_from_errno)?; - if initial_readiness != ReadinessFlags::default() { - snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness = initial_readiness; - readiness.publish(initial_readiness)?; + let old_socket = &self.sockets.get(&id).ok_or(BrokerError::Internal)?.socket; + if let Err(error) = quiesce_tcp_listener(old_socket) { + let _ = delete_epoll_registration(&self.epoll, &replacement); + return Err(error); } - self.sockets.insert( - id, - SocketEntry { - socket, - kind, - readiness, - snapshot, - read_shutdown: false, - write_shutdown: false, - peek_waitall_threshold: None, - listening: false, - }, - ); - Ok(()) + if !delete_epoll_registration(&self.epoll, old_socket) { + let _ = delete_epoll_registration(&self.epoll, &replacement); + return Err(BrokerError::Internal); + } + let reservation = { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + let reservation = core::mem::replace(&mut socket.socket, replacement); + socket.listening = false; + socket.read_shutdown = true; + socket.peek_waitall_threshold = None; + socket.port_mapping_index = None; + reservation + }; + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + state.owned_by = None; + state.reservation = Some(reservation); + state.reservation_registered = false; + Ok(SocketOutcome::Completed(())) } - fn accept_socket( + fn release_failed_tcp_port_mapping( &mut self, - listener_id: u64, - accepted_id: u64, - readiness: ReadinessRegistration, - snapshot: Arc>, - ) -> BrokerResult> { - if self.sockets.len() >= self.max_sockets { - return Err(BrokerError::ResourceExhausted); - } - if self.sockets.contains_key(&accepted_id) { + id: u64, + mapping_index: usize, + ) -> BrokerResult<()> { + if !self + .port_mappings + .get(mapping_index) + .is_some_and(|state| state.owned_by == Some(id) && state.reservation.is_none()) + { return Err(BrokerError::Internal); } - let listener = self + let guest_port = self .sockets - .get_mut(&listener_id) + .get(&id) + .and_then(|socket| { + (!socket.listening + && socket.port_mapping_index == Some(mapping_index) + && socket.mapping_fallback_socket.is_some()) + .then_some(())?; + socket + .guest_local_address + .map(|guest_address| guest_address.port()) + }) .ok_or(BrokerError::Internal)?; - if listener.kind != SocketKind::Tcp || !listener.listening { - return Ok(SocketOutcome::Failed(SocketError::NotConnected)); - } - let (socket, remote_address) = loop { - match acceptfrom_with( - &listener.socket, - LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, - ) { - Ok((socket, address)) => break (socket, address), - Err(Errno::INTR) => {} - Err(Errno::AGAIN) => { - clear_readiness(listener, ReadinessFlags::READ)?; - return Err(BrokerError::WouldBlock); - } - Err(error) => { - return Ok(SocketOutcome::Failed(socket_operation_error_from_errno( - error, - )?)); - } - } + let host_address_is_set = self + .tcp + .bindings + .get(&guest_port) + .filter(|binding| binding.socket_id == id) + .map(|binding| binding.host_mapped); + let reservation = { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + let fallback = socket + .mapping_fallback_socket + .take() + .ok_or(BrokerError::Internal)?; + socket.port_mapping_index = None; + core::mem::replace(&mut socket.socket, fallback) }; - let no_wait = Timespec { - tv_sec: 0, - tv_nsec: 0, + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + state.owned_by = None; + state.reservation = Some(reservation); + state.reservation_registered = true; + let Some(host_address_is_set) = host_address_is_set else { + return Err(BrokerError::Internal); }; - loop { - let mut poll_fd = [PollFd::new(&listener.socket, PollFlags::IN)]; - match poll(&mut poll_fd, Some(&no_wait)) { - Ok(_) if poll_fd[0].revents().contains(PollFlags::IN) => break, - Ok(_) => { - clear_readiness(listener, ReadinessFlags::READ)?; - break; - } - Err(Errno::INTR) => {} - Err(error) => return Err(broker_error_from_errno(error)), + if host_address_is_set { + self.tcp.clear_host_address(guest_port, id)?; + } + Ok(()) + } + + /// Returns whether `address` names a host endpoint backing a guest socket. + /// + /// Private backend endpoints are an implementation detail of the guest + /// namespace and must not be reachable as guest destinations. + fn is_private_host_endpoint( + &self, + kind: SocketKind, + address: SocketAddrV4, + ) -> BrokerResult { + if kind != SocketKind::Tcp { + return Ok(false); + } + for binding in self.tcp.bindings.values() { + let Some(host_address) = binding.host_address else { + continue; + }; + if binding.host_mapped || host_address.port() != address.port() { + continue; + } + if host_address.ip() == address.ip() + || (host_address.ip().is_unspecified() + && host_ipv4_address_is_local(*address.ip())?) + { + return Ok(true); } } - let remote_address = SocketAddrV4::try_from(remote_address.ok_or(BrokerError::Internal)?) - .map_err(|_| BrokerError::Internal)?; - let local_address = local_socket_address(&socket)?; - epoll::add( - &self.epoll, - &socket, - epoll::EventData::new_u64(accepted_id), - active_epoll_events(), - ) - .map_err(broker_error_from_errno)?; + Ok(false) + } + + /// Resolves a guest destination to the host endpoint that should receive it. + fn resolve_guest_destination( + &self, + kind: SocketKind, + mut address: SocketAddrV4, + ) -> BrokerResult)>> { + if address.ip().is_unspecified() { + address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, address.port()); + } + if kind == SocketKind::Tcp + && let Some(binding) = self.tcp.guest_binding(address) { - let mut snapshot = snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - snapshot.status = SocketConnectionStatus::Connected; - snapshot.local_address = Some(local_address); - snapshot.readiness = ReadinessFlags::WRITE; + if !binding.listening { + return Ok(SocketOutcome::Failed(SocketError::ConnectionRefused)); + } + return match binding.host_address { + Some(host_address) if host_address.ip().is_unspecified() => { + Ok(SocketOutcome::Completed(( + SocketAddrV4::new(*address.ip(), host_address.port()), + Some(binding.socket_id), + ))) + } + Some(host_address) => Ok(SocketOutcome::Completed(( + host_address, + Some(binding.socket_id), + ))), + None => Ok(SocketOutcome::Failed(SocketError::ConnectionRefused)), + }; + } + if self.is_private_host_endpoint(kind, address)? { + Ok(SocketOutcome::Failed(SocketError::ConnectionRefused)) + } else { + Ok(SocketOutcome::Completed((address, None))) } - readiness.publish(ReadinessFlags::WRITE)?; - self.sockets.insert( - accepted_id, - SocketEntry { - socket, - kind: SocketKind::Tcp, - readiness, - snapshot, - read_shutdown: false, - write_shutdown: false, - peek_waitall_threshold: None, - listening: false, - }, - ); - Ok(SocketOutcome::Completed(AcceptedEndpoints { - local_address, - remote_address, - })) } - fn fail_all_sockets(&mut self) { - for socket in self.sockets.values() { - let mut snapshot = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - snapshot.status = SocketConnectionStatus::Failed(SocketError::Other); - snapshot.readiness = ReadinessFlags::ERROR; - drop(snapshot); - // The readiness path may itself be why the reactor is failing. The - // cached terminal snapshot remains authoritative if publication is - // no longer available. - let _ = socket.readiness.publish(ReadinessFlags::ERROR); + /// Binds a socket to its guest-local address. + /// + /// A stream socket reserves a guest port without consuming a host port; the + /// host endpoint is chosen when the socket listens or connects. A datagram + /// socket is still bound directly to the requested host endpoint. + fn bind_socket( + &mut self, + id: u64, + requested_address: SocketAddrV4, + ) -> BrokerResult> { + let (kind, already_bound) = { + let socket = self.sockets.get(&id).ok_or(BrokerError::Internal)?; + (socket.kind, socket.guest_local_address.is_some()) + }; + if kind != SocketKind::Tcp { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + return match bind_host_socket(socket, requested_address)? { + SocketOutcome::Completed(local_address) => { + socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .local_address = Some(local_address); + Ok(SocketOutcome::Completed(local_address)) + } + SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), + }; } - self.sockets.clear(); + if already_bound { + return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); + } + let guest_port = requested_address.port(); + if guest_port == 0 { + return Err(BrokerError::Internal); + } + if self.tcp.bindings.contains_key(&guest_port) { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } + self.tcp.reserve_binding()?; + let guest_address = SocketAddrV4::new(*requested_address.ip(), guest_port); + self.tcp.insert_binding( + guest_port, + GuestPortBinding { + socket_id: id, + guest_address, + host_address: None, + host_peer_address: None, + host_peer_mapping_index: None, + host_mapped: false, + listening: false, + }, + )?; + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + socket.guest_local_address = Some(guest_address); + socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .local_address = Some(guest_address); + Ok(SocketOutcome::Completed(guest_address)) } -} -fn connect_socket( - epoll_fd: &OwnedFd, - id: u64, - socket: &mut SocketEntry, - address: SocketAddrV4, -) -> core::result::Result { - if socket.kind == SocketKind::Udp { - return connect_datagram_socket(socket, address); - } - if let Err(error) = epoll::modify( - epoll_fd, - &socket.socket, - epoll::EventData::new_u64(id), - active_epoll_events(), - ) { - return Err(PlatformConnectError::PeerUnchanged( - broker_error_from_errno(error), - )); - } - let status = loop { - match connect(&socket.socket, &address) { - Ok(()) | Err(Errno::ISCONN) => break SocketConnectionStatus::Connected, - Err(Errno::INTR) => {} - Err(Errno::INPROGRESS | Errno::ALREADY) => { - break SocketConnectionStatus::Connecting; + fn listen_socket( + &mut self, + id: u64, + backlog: u32, + mapping: Option, + ) -> BrokerResult> { + let (kind, guest_address, current_mapping_index) = self + .sockets + .get(&id) + .map(|socket| { + ( + socket.kind, + socket.guest_local_address, + socket.port_mapping_index, + ) + }) + .ok_or(BrokerError::Internal)?; + if kind != SocketKind::Tcp { + return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); + } + let guest_address = guest_address.ok_or(BrokerError::Internal)?; + let needs_host_bind = + local_socket_address(&self.sockets.get(&id).ok_or(BrokerError::Internal)?.socket)? + .port() + == 0; + if mapping.is_some_and(|mapping| mapping.guest_port != guest_address.port()) { + return Err(BrokerError::Internal); + } + let requested_mapping_index = mapping.and_then(|mapping| { + self.port_mappings + .iter() + .position(|state| state.mapping == mapping) + }); + let identity_mapping = mapping.filter(|_| requested_mapping_index.is_none()); + if identity_mapping.is_some_and(|mapping| mapping.broker_port != mapping.guest_port) { + return Err(BrokerError::Internal); + } + let currently_host_mapped = self + .tcp + .bindings + .get(&guest_address.port()) + .filter(|binding| binding.socket_id == id) + .map(|binding| binding.host_mapped) + .ok_or(BrokerError::Internal)?; + if let Some(current_mapping_index) = current_mapping_index { + if requested_mapping_index != Some(current_mapping_index) { + return Err(BrokerError::Internal); } - Err(error) => { - let error = match socket_operation_error_from_errno(error) { - Ok(error) => error, - Err(error) => { - update_snapshot( - socket, - Some(SocketConnectionStatus::Failed(SocketError::Other)), - ReadinessFlags::ERROR, - ) - .map_err(PlatformConnectError::PeerIndeterminate)?; - return Err(PlatformConnectError::PeerIndeterminate(error)); - } - }; - break SocketConnectionStatus::Failed(error); + if self + .port_mappings + .get(current_mapping_index) + .is_none_or(|state| state.owned_by != Some(id)) + { + return Err(BrokerError::Internal); + } + } else if !needs_host_bind { + if mapping.is_some() != currently_host_mapped { + return Err(BrokerError::Internal); + } + if let Some(mapping) = mapping + && local_socket_address( + &self.sockets.get(&id).ok_or(BrokerError::Internal)?.socket, + )? + .port() + != mapping.broker_port + { + return Err(BrokerError::Internal); } } - }; - let readiness = match status { - SocketConnectionStatus::Connected | SocketConnectionStatus::Connecting => { - let local_address = local_socket_address(&socket.socket) - .map_err(PlatformConnectError::PeerIndeterminate)?; - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(local_address); - if status == SocketConnectionStatus::Connected { - ReadinessFlags::WRITE + let new_mapping = current_mapping_index + .is_none() + .then_some(requested_mapping_index) + .flatten(); + let new_identity_mapping = needs_host_bind.then_some(identity_mapping).flatten(); + if needs_host_bind { + let (host_address, host_mapped) = if let Some(mapping_index) = requested_mapping_index { + match self.realize_tcp_port_mapping(id, mapping_index)? { + SocketOutcome::Completed(address) => (address, true), + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + } + } else if let Some(mapping) = identity_mapping { + match self.realize_identity_tcp_mapping(id, mapping)? { + SocketOutcome::Completed(address) => (address, true), + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + } } else { - ReadinessFlags::default() + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + let address = + match bind_host_socket(socket, SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0))? { + SocketOutcome::Completed(address) => address, + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + }; + (address, false) + }; + if let Err(error) = + self.tcp + .set_host_address(guest_address.port(), id, host_address, host_mapped) + { + if let Some(mapping_index) = new_mapping { + self.release_failed_tcp_port_mapping(id, mapping_index)?; + } else if new_identity_mapping.is_some() { + self.release_failed_identity_tcp_mapping(id)?; + } + return Err(error); } } - SocketConnectionStatus::Failed(_) => ReadinessFlags::ERROR, - SocketConnectionStatus::Unconnected => ReadinessFlags::default(), - _ => { - return Err(PlatformConnectError::PeerIndeterminate( - BrokerError::Internal, - )); + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + let host_mapped = mapping.is_some(); + let outcome = listen_tcp_socket(&self.epoll, id, socket, backlog, host_mapped); + match outcome { + Ok(SocketOutcome::Completed(())) => { + self.tcp.mark_listening(guest_address.port(), id)?; + if let Some(mapping_index) = new_mapping { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + if socket.port_mapping_index != Some(mapping_index) { + return Err(BrokerError::Internal); + } + drop( + socket + .mapping_fallback_socket + .take() + .ok_or(BrokerError::Internal)?, + ); + } else if new_identity_mapping.is_some() { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + if socket.port_mapping_index.is_some() { + return Err(BrokerError::Internal); + } + drop( + socket + .mapping_fallback_socket + .take() + .ok_or(BrokerError::Internal)?, + ); + } + Ok(SocketOutcome::Completed(guest_address)) + } + Ok(SocketOutcome::Failed(error)) => { + if let Some(mapping_index) = new_mapping { + self.release_failed_tcp_port_mapping(id, mapping_index)?; + } else if new_identity_mapping.is_some() { + self.release_failed_identity_tcp_mapping(id)?; + } + Ok(SocketOutcome::Failed(error)) + } + Err(error) => { + if let Some(mapping_index) = new_mapping { + self.release_failed_tcp_port_mapping(id, mapping_index)?; + } else if new_identity_mapping.is_some() { + self.release_failed_identity_tcp_mapping(id)?; + } + Err(error) + } } - }; - update_snapshot(socket, Some(status), readiness) - .map_err(PlatformConnectError::PeerIndeterminate)?; - Ok(status) -} + } -fn connect_datagram_socket( - socket: &mut SocketEntry, - address: SocketAddrV4, -) -> core::result::Result { - loop { - match connect(&socket.socket, &address) { - Ok(()) | Err(Errno::ISCONN) => { - let local_address = local_socket_address(&socket.socket) - .map_err(PlatformConnectError::PeerIndeterminate)?; - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(local_address); - let readiness = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness; - let readiness = if socket.write_shutdown { - ReadinessFlags(readiness.0 & !ReadinessFlags::WRITE.0) - } else { - readiness | ReadinessFlags::WRITE - }; - update_snapshot(socket, Some(SocketConnectionStatus::Connected), readiness) - .map_err(PlatformConnectError::PeerIndeterminate)?; - return Ok(SocketConnectionStatus::Connected); - } - Err(Errno::INTR) => {} - Err(error) => { - update_local_address(socket).map_err(PlatformConnectError::PeerIndeterminate)?; - let error = socket_operation_error_from_errno(error) - .map_err(PlatformConnectError::PeerIndeterminate)?; - return Ok(SocketConnectionStatus::Failed(error)); - } + fn connect_socket( + &mut self, + id: u64, + address: SocketAddrV4, + ) -> core::result::Result { + let kind = self.sockets.get(&id).map(|socket| socket.kind).ok_or( + PlatformConnectError::PeerIndeterminate(BrokerError::Internal), + )?; + if kind != SocketKind::Tcp { + let socket = + self.sockets + .get_mut(&id) + .ok_or(PlatformConnectError::PeerIndeterminate( + BrokerError::Internal, + ))?; + return connect_datagram_socket(socket, address); + } + let owns_mapping = self.owned_tcp_port_mapping(id).is_some(); + if owns_mapping { + let status = SocketConnectionStatus::Failed(SocketError::InvalidArgument); + update_snapshot( + self.sockets + .get(&id) + .ok_or(PlatformConnectError::PeerUnchanged(BrokerError::Internal))?, + Some(status), + ReadinessFlags::ERROR, + ) + .map_err(PlatformConnectError::PeerIndeterminate)?; + return Ok(status); } + self.connect_guest_tcp_socket(id, address) } -} -fn bind_socket( - socket: &mut SocketEntry, - address: SocketAddrV4, -) -> BrokerResult> { - loop { - match bind(&socket.socket, &address) { - Ok(()) => { - let local_address = local_socket_address(&socket.socket)?; - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(local_address); - return Ok(SocketOutcome::Completed(local_address)); + fn connect_guest_tcp_socket( + &mut self, + id: u64, + guest_address: SocketAddrV4, + ) -> core::result::Result { + let session_id = { + let socket = self + .sockets + .get(&id) + .ok_or(PlatformConnectError::PeerUnchanged(BrokerError::Internal))?; + if socket.guest_local_address.is_none() { + return Err(PlatformConnectError::PeerUnchanged(BrokerError::Internal)); } - Err(Errno::INTR) => {} - Err(error) => { - return Ok(SocketOutcome::Failed(socket_operation_error_from_errno( - error, - )?)); + socket.session_id + }; + let (network_address, guest_listener_id) = match self + .resolve_guest_destination(SocketKind::Tcp, guest_address) + .map_err(PlatformConnectError::PeerUnchanged)? + { + SocketOutcome::Completed(destination) => destination, + SocketOutcome::Failed(error) => { + let status = SocketConnectionStatus::Failed(error); + update_snapshot( + self.sockets + .get(&id) + .ok_or(PlatformConnectError::PeerUnchanged(BrokerError::Internal))?, + Some(status), + ReadinessFlags::ERROR, + ) + .map_err(PlatformConnectError::PeerIndeterminate)?; + return Ok(status); + } + }; + if guest_listener_id.is_none() { + let source_ip = self + .sockets + .get(&id) + .and_then(|socket| socket.guest_local_address) + .map(|address| { + if address.ip().is_unspecified() { + self.broker_ipv4_address + } else { + *address.ip() + } + }) + .ok_or(PlatformConnectError::PeerUnchanged(BrokerError::Internal))?; + let socket = self + .sockets + .get_mut(&id) + .ok_or(PlatformConnectError::PeerUnchanged(BrokerError::Internal))?; + if local_socket_address(&socket.socket) + .map_err(PlatformConnectError::PeerUnchanged)? + .port() + == 0 + { + match bind_host_socket(socket, SocketAddrV4::new(source_ip, 0)) + .map_err(PlatformConnectError::PeerUnchanged)? + { + SocketOutcome::Completed(_) => {} + SocketOutcome::Failed(error) => { + let status = SocketConnectionStatus::Failed(error); + update_snapshot(socket, Some(status), ReadinessFlags::ERROR) + .map_err(PlatformConnectError::PeerIndeterminate)?; + return Ok(status); + } + } } } - } -} - -fn listen_socket( - epoll_fd: &OwnedFd, - id: u64, - socket: &mut SocketEntry, - backlog: u32, -) -> BrokerResult> { - if socket.kind != SocketKind::Tcp { - return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); - } - let backlog = i32::try_from(backlog).map_err(|_| BrokerError::UnsupportedOperation)?; - let was_listening = socket.listening; - if !was_listening { - epoll::modify( - epoll_fd, - &socket.socket, - epoll::EventData::new_u64(id), - active_epoll_events(), - ) - .map_err(broker_error_from_errno)?; - } - loop { - match listen(&socket.socket, backlog) { - Ok(()) => break, - Err(Errno::INTR) => {} - Err(error) => { - if !was_listening { - epoll::modify( - epoll_fd, - &socket.socket, - epoll::EventData::new_u64(id), - idle_epoll_events(), + let guest_mapping_index = guest_listener_id + .and_then(|listener_id| self.sockets.get(&listener_id)) + .and_then(|listener| listener.port_mapping_index); + if guest_listener_id.is_some() { + let now = Instant::now(); + self.expire_deadlined_state(now); + self.reserve_pending_guest_connection(session_id) + .map_err(PlatformConnectError::PeerUnchanged)?; + } + let (outcome, readiness) = { + let socket = self + .sockets + .get_mut(&id) + .ok_or(PlatformConnectError::PeerUnchanged(BrokerError::Internal))?; + connect_tcp_socket(&self.epoll, id, socket, network_address)? + }; + if matches!( + outcome, + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + ) { + let (local_guest_address, host_address) = + { + let socket = + self.sockets + .get(&id) + .ok_or(PlatformConnectError::PeerIndeterminate( + BrokerError::Internal, + ))?; + ( + socket.guest_local_address.ok_or( + PlatformConnectError::PeerIndeterminate(BrokerError::Internal), + )?, + local_socket_address(&socket.socket) + .map_err(PlatformConnectError::PeerIndeterminate)?, ) - .map_err(broker_error_from_errno)?; + }; + if guest_listener_id.is_some() { + if let Some(mapping_index) = guest_mapping_index { + self.port_mappings.get(mapping_index).ok_or( + PlatformConnectError::PeerIndeterminate(BrokerError::Internal), + )?; + self.take_stale_tcp_connection(mapping_index, host_address, network_address); } - return Ok(SocketOutcome::Failed(socket_operation_error_from_errno( - error, - )?)); + self.remove_pending_guest_connection_except( + Some(session_id), + host_address, + network_address, + ); + } + self.tcp + .set_host_address(local_guest_address.port(), id, host_address, false) + .map_err(PlatformConnectError::PeerIndeterminate)?; + if let Some(listener_id) = guest_listener_id { + let (host_address, guest_address) = self + .tcp + .set_host_peer_address( + local_guest_address.port(), + id, + network_address, + guest_mapping_index, + ) + .map_err(PlatformConnectError::PeerIndeterminate)?; + self.insert_pending_guest_connection( + session_id, + (host_address, network_address), + guest_address, + listener_id, + ) + .map_err(PlatformConnectError::PeerIndeterminate)?; } } + update_snapshot( + self.sockets + .get(&id) + .ok_or(PlatformConnectError::PeerIndeterminate( + BrokerError::Internal, + ))?, + Some(outcome), + readiness, + ) + .map_err(PlatformConnectError::PeerIndeterminate)?; + Ok(outcome) } - socket.listening = true; - // An unbound TCP socket is implicitly bound by listen(2), so query the assigned - // address only after listen succeeds. - let local_address = local_socket_address(&socket.socket)?; - let mut snapshot = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - snapshot.local_address = Some(local_address); - if !was_listening { - snapshot.readiness = ReadinessFlags::default(); - } - Ok(SocketOutcome::Completed(local_address)) -} -fn send_socket(socket: &mut SocketEntry, data: &[u8]) -> BrokerResult> { - if socket.kind != SocketKind::Tcp { - return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); - } - if socket.write_shutdown { - return Ok(SocketOutcome::Failed(SocketError::Other)); - } - loop { - match send(&socket.socket, data, LinuxSendFlags::NOSIGNAL) { - Ok(sent) => return Ok(SocketOutcome::Completed(sent)), - Err(Errno::INTR) => {} - Err(Errno::AGAIN) => { - clear_readiness(socket, ReadinessFlags::WRITE)?; - return Err(BrokerError::WouldBlock); + fn remove_pending_guest_connection_for_connector( + &mut self, + session_id: SessionId, + kind: SocketKind, + guest_port: u16, + socket_id: u64, + ) { + if kind != SocketKind::Tcp { + return; + } + if let Some((connection, mapping_index)) = + self.tcp.take_connector_connection(guest_port, socket_id) + { + if let Some(pending) = self.tcp.pending_guest_connections.remove(&connection) { + debug_assert_eq!(pending.session_id, session_id); + self.finish_removed_pending_guest_connection(&pending); } - Err(error) => { - let error = socket_operation_error_from_errno(error)?; - consume_synchronous_error(socket)?; - return Ok(SocketOutcome::Failed(error)); + if let Some(mapping_index) = mapping_index + && let Some(stale) = self.tcp.stale_mapped_connections.remove(&( + mapping_index, + connection.0, + connection.1, + )) + { + self.release_stale_tcp_connection(stale); } } } -} -fn send_to_socket( - socket: &mut SocketEntry, - data: &[u8], - destination: Option, -) -> BrokerResult> { - if socket.kind != SocketKind::Udp || data.len() > MAX_UDP_DATAGRAM_SIZE as usize { - return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); - } - if socket.write_shutdown { - return Ok(SocketOutcome::Failed(SocketError::Other)); - } - loop { - let result = match destination { - Some(address) => sendto(&socket.socket, data, LinuxSendFlags::NOSIGNAL, &address), - None => send(&socket.socket, data, LinuxSendFlags::NOSIGNAL), + /// Drops a socket, releasing its guest port and any owned host mapping. + fn remove_socket(&mut self, id: u64) { + let port_mapping = self + .owned_tcp_port_mapping(id) + .map(|mapping_index| (mapping_index, self.port_mappings[mapping_index].mapping)); + let retired_listener_address = self.sockets.get(&id).and_then(|socket| { + (socket.kind == SocketKind::Tcp && socket.listening) + .then(|| local_socket_address(&socket.socket).ok()) + .flatten() + }); + // Keep a mapped listener active with a zero backlog so the endpoint + // remains exclusive without accumulating an unbounded stale queue. + let retain_original = port_mapping.is_some_and(|_| { + self.sockets.get(&id).is_some_and(|socket| { + delete_epoll_registration(&self.epoll, &socket.socket) + && if socket.listening { + quiesce_tcp_listener(&socket.socket).is_ok() + } else { + sockopt::set_socket_reuseaddr(&socket.socket, false).is_ok() + } + }) + }); + let Some(socket) = self.sockets.remove(&id) else { + return; }; - match result { - Ok(sent) if sent == data.len() => { - update_local_address(socket)?; - return Ok(SocketOutcome::Completed(sent)); - } - Ok(_) => { - update_local_address(socket)?; - return Err(BrokerError::Internal); + let SocketEntry { + socket, + session_id, + kind, + snapshot, + abortive_close, + guest_local_address, + .. + } = socket; + let mut socket = Some(socket); + let connecting = snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .status + == SocketConnectionStatus::Connecting; + let pending_connection_disposition = + match getpeername(socket.as_ref().expect("removed socket descriptor missing")) { + Ok(Some(_)) if abortive_close => PendingGuestConnectionDisposition::Discard(None), + Ok(Some(_)) => PendingGuestConnectionDisposition::Retain, + Ok(None) | Err(_) => { + if abortive_close || connecting { + let _ = sockopt::set_socket_linger( + socket.as_ref().expect("removed socket descriptor missing"), + Some(Duration::ZERO), + ); + } + PendingGuestConnectionDisposition::Discard(Some( + Instant::now() + PENDING_CONNECT_DISCARD_LIFETIME, + )) + } + }; + if retired_listener_address.is_some() + && let Some((mapping_index, _)) = port_mapping + && !retain_original + { + self.clear_stale_tcp_connections(mapping_index); + } + if let Some(listener_address) = retired_listener_address { + if let Some((mapping_index, _)) = port_mapping + && retain_original + { + self.move_pending_guest_connections_for_listener(listener_address, mapping_index); + } else { + self.remove_pending_guest_connections_for_listener(listener_address); } - Err(Errno::INTR) => {} - Err(Errno::AGAIN) => { - update_local_address(socket)?; - clear_readiness(socket, ReadinessFlags::WRITE)?; - return Err(BrokerError::WouldBlock); + } + let mut discarded_connector = None; + if kind == SocketKind::Tcp + && let Some(address) = guest_local_address + { + match pending_connection_disposition { + PendingGuestConnectionDisposition::Retain => { + let retained_connector = + socket.take().expect("removed socket descriptor missing"); + discarded_connector = self.retire_pending_guest_connection_for_connector( + session_id, + address.port(), + id, + false, + None, + Some(retained_connector), + ); + } + PendingGuestConnectionDisposition::Discard(discard_deadline) => { + let retained_connector = discard_deadline + .is_none() + .then(|| socket.take().expect("removed socket descriptor missing")); + discarded_connector = self.retire_pending_guest_connection_for_connector( + session_id, + address.port(), + id, + true, + discard_deadline, + retained_connector, + ); + } } - Err(error) => { - update_local_address(socket)?; - let error = socket_operation_error_from_errno(error)?; - consume_synchronous_error(socket)?; - return Ok(SocketOutcome::Failed(error)); + if let Some(retired) = discarded_connector.as_mut() + && let Some(mapping_index) = retired.mapping_index + && let Some(stale) = self.tcp.stale_mapped_connections.get_mut(&( + mapping_index, + retired.connection.0, + retired.connection.1, + )) + && stale.session_id == session_id + { + stale.deadline = match (stale.deadline, pending_connection_disposition) { + (existing, PendingGuestConnectionDisposition::Discard(None)) => existing, + (_, PendingGuestConnectionDisposition::Retain) => None, + (None, PendingGuestConnectionDisposition::Discard(Some(deadline))) => { + Some(deadline) + } + ( + Some(existing), + PendingGuestConnectionDisposition::Discard(Some(deadline)), + ) => Some(existing.max(deadline)), + }; + if stale.retained_connector.is_none() + && let Some(retained_connector) = retired.unplaced_connector.take() + { + stale.deadline = None; + stale.retained_connector = Some(retained_connector); + self.retained_connectors = self + .retained_connectors + .checked_add(1) + .expect("reactor retained connector count overflow"); + let session = self + .sessions + .get_mut(&session_id) + .expect("socket session state missing"); + session.retained_connectors = session + .retained_connectors + .checked_add(1) + .expect("session retained connector count overflow"); + } } + self.tcp.remove_binding(address.port(), id); } - } -} - -fn receive_socket( - socket: &mut SocketEntry, - peek_cache: &mut Option, - socket_id: u64, - length: usize, - flags: ReceiveFlags, - peek_offset: usize, - peek_length: usize, -) -> BrokerResult { - if socket.kind != SocketKind::Tcp { - return Ok(ReactorReceiveOutcome::Failed(SocketError::InvalidArgument)); - } - let peek = flags.contains(ReceiveFlags::PEEK); - if !peek { - if peek_offset != 0 || peek_length != 0 { - return Err(BrokerError::UnsupportedOperation); + if let Some(session) = self.sessions.get_mut(&session_id) { + session.live_sockets = session + .live_sockets + .checked_sub(1) + .expect("session socket count underflow"); } - if peek_cache - .as_ref() - .is_some_and(|cache| cache.socket_id == socket_id) + drop(discarded_connector); + if self + .sessions + .get(&session_id) + .is_some_and(|session| session.closing) { - *peek_cache = None; + self.retire_session_connectors(session_id); + } + let remove_session = self + .sessions + .get(&session_id) + .is_some_and(|session| !retain_session_state(session)); + if remove_session { + self.sessions.remove(&session_id); + } + let replacement_reservation = if retain_original { + Some( + socket + .take() + .expect("mapping owner descriptor unexpectedly retained"), + ) + } else { + drop(socket.take()); + port_mapping.and_then(|(_, mapping)| { + create_replacement_port_mapping_reservation(self.broker_ipv4_address, mapping).ok() + }) + }; + if let Some((mapping_index, _)) = port_mapping + && let Some(state) = self.port_mappings.get_mut(mapping_index) + && state.owned_by == Some(id) + { + state.owned_by = None; + state.reservation = replacement_reservation; + state.reservation_registered = false; } - return receive_socket_once(socket, zeroed_vec(length)?, LinuxRecvFlags::empty()); } - let peek_end = peek_offset - .checked_add(length) - .ok_or(BrokerError::UnsupportedOperation)?; - let canonical_length = peek_length - .checked_sub(peek_offset) - .map(|remaining| remaining.min(MAX_SOCKET_TRANSFER_SIZE as usize)); - if !peek_offset.is_multiple_of(MAX_SOCKET_TRANSFER_SIZE as usize) - || canonical_length != Some(length) - || peek_length < peek_end - || peek_length > MAX_SOCKET_PEEK_SIZE as usize - { - return Err(BrokerError::UnsupportedOperation); - } - if flags.contains(ReceiveFlags::WAITALL) { - let snapshot = *socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - let terminal = socket.read_shutdown - || snapshot.readiness.contains(ReadinessFlags::HANGUP) - || snapshot.readiness.contains(ReadinessFlags::ERROR); - if snapshot.status == SocketConnectionStatus::Connected - && !terminal - && ioctl_fionread(&socket.socket).map_err(broker_error_from_errno)? - < peek_length.try_into().map_err(|_| BrokerError::Internal)? - { - socket.peek_waitall_threshold = Some( - socket - .peek_waitall_threshold - .map_or(peek_length, |threshold| threshold.min(peek_length)), - ); - return Err(BrokerError::WouldBlock); - } - } - - let refresh = peek_offset == 0 - || !peek_cache.as_ref().is_some_and(|cache| { - cache.socket_id == socket_id && cache.requested_length == peek_length - }); - if refresh { - *peek_cache = None; - let flags = if flags.contains(ReceiveFlags::WAITALL) { - LinuxRecvFlags::PEEK | LinuxRecvFlags::WAITALL - } else { - LinuxRecvFlags::PEEK - }; - match receive_socket_once(socket, zeroed_vec(peek_length)?, flags)? { - ReactorReceiveOutcome::Received(data) => { - *peek_cache = Some(PeekCache { - socket_id, - requested_length: peek_length, - data, - }); - } - outcome => return Ok(outcome), - } - } - - let cache = peek_cache.as_ref().ok_or(BrokerError::Internal)?; - if cache.data.len() <= peek_offset { - let readiness = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness; - let terminal = socket.read_shutdown - || readiness.contains(ReadinessFlags::HANGUP) - || readiness.contains(ReadinessFlags::ERROR); - *peek_cache = None; - return if terminal { - Ok(ReactorReceiveOutcome::EndOfStream) - } else { - Err(BrokerError::WouldBlock) - }; - } - let end = peek_end.min(cache.data.len()); - let mut data = Vec::new(); - data.try_reserve_exact(end - peek_offset) - .map_err(|_| BrokerError::OutOfMemory)?; - data.extend_from_slice(&cache.data[peek_offset..end]); - if end < peek_end || end == peek_length { - *peek_cache = None; - } - Ok(ReactorReceiveOutcome::Received(data)) -} - -fn receive_from_socket( - socket: &mut SocketEntry, - length: usize, - flags: ReceiveFromFlags, -) -> BrokerResult { - if socket.kind != SocketKind::Udp - || length > MAX_UDP_DATAGRAM_SIZE as usize - || flags.has_unsupported_bits() - { - return Ok(ReactorReceiveFromOutcome::Failed( - SocketError::InvalidArgument, - )); - } - if socket.read_shutdown { - return Ok(ReactorReceiveFromOutcome::Failed(SocketError::NotConnected)); - } - let mut data = zeroed_vec(length)?; - let mut linux_flags = LinuxRecvFlags::TRUNC; - if flags.contains(ReceiveFromFlags::PEEK) { - linux_flags |= LinuxRecvFlags::PEEK; - } - loop { - match recvfrom(&socket.socket, data.as_mut_slice(), linux_flags) { - Ok((received, datagram_length, address)) => { - let source_address = SocketAddrV4::try_from(address.ok_or(BrokerError::Internal)?) - .map_err(|_| BrokerError::Internal)?; - data.truncate(received); - update_local_address(socket)?; - return Ok(ReactorReceiveFromOutcome::Received { - data, - datagram_length, - source_address, - }); - } - Err(Errno::INTR) => {} - Err(Errno::AGAIN) => { - clear_readiness(socket, ReadinessFlags::READ)?; - return Err(BrokerError::WouldBlock); - } - Err(error) => { - let error = socket_operation_error_from_errno(error)?; - consume_synchronous_error(socket)?; - return Ok(ReactorReceiveFromOutcome::Failed(error)); + fn run(&mut self) -> core::result::Result<(), ReactorFailure> { + loop { + let mut events = core::mem::take(&mut self.events); + events.clear(); + let now = Instant::now(); + let timeout = self.next_cleanup_deadline().map(|deadline| { + let duration = deadline.saturating_duration_since(now); + Timespec { + tv_sec: i64::try_from(duration.as_secs()).unwrap_or(i64::MAX), + tv_nsec: i64::from(duration.subsec_nanos()), + } + }); + match epoll::wait(&self.epoll, spare_capacity(&mut events), timeout.as_ref()) { + Ok(_) => {} + Err(Errno::INTR) => { + self.events = events; + continue; + } + Err(error) => return Err(ReactorFailure::Io(error)), } - } - } -} - -fn zeroed_vec(length: usize) -> BrokerResult> { - let mut data = Vec::new(); - data.try_reserve_exact(length) - .map_err(|_| BrokerError::OutOfMemory)?; - data.resize(length, 0); - Ok(data) -} + self.expire_deadlined_state(Instant::now()); -fn receive_socket_once( - socket: &mut SocketEntry, - mut data: Vec, - flags: LinuxRecvFlags, -) -> BrokerResult { - loop { - match recv(&socket.socket, data.as_mut_slice(), flags) { - Ok((_buffer, 0)) => { - let readiness = if socket.read_shutdown { - ReadinessFlags::READ + // Apply readiness observed by this wait before commands. A command + // that then reaches EAGAIN records the newer authoritative state. + let mut wake = false; + for event in events.drain(..) { + let id = event.data.u64(); + if id == WAKE_TOKEN { + wake = true; } else { - ReadinessFlags::READ | ReadinessFlags::HANGUP - }; - add_readiness(socket, readiness)?; - return Ok(ReactorReceiveOutcome::EndOfStream); - } - Ok((_buffer, received)) => { - data.truncate(received); - let terminal_readable = socket.read_shutdown - || socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness - .contains(ReadinessFlags::HANGUP); - if !flags.contains(LinuxRecvFlags::PEEK) - && !terminal_readable - && ioctl_fionread(&socket.socket).map_err(broker_error_from_errno)? == 0 - { - clear_readiness(socket, ReadinessFlags::READ)?; + let failed_connector = if let Some(socket) = self.sockets.get_mut(&id) { + let session_id = socket.session_id; + let guest_port = socket.guest_local_address.map(|address| address.port()); + handle_socket_event(socket, event.flags) + .map_err(ReactorFailure::Broker)? + .then_some((session_id, socket.kind, guest_port)) + } else { + None + }; + if let Some((session_id, kind, Some(guest_port))) = failed_connector { + self.remove_pending_guest_connection_for_connector( + session_id, kind, guest_port, id, + ); + } } - return Ok(ReactorReceiveOutcome::Received(data)); - } - Err(Errno::INTR) => {} - Err(Errno::AGAIN) => { - clear_readiness(socket, ReadinessFlags::READ)?; - return Err(BrokerError::WouldBlock); } - Err(error) => { - let error = socket_operation_error_from_errno(error)?; - consume_synchronous_error(socket)?; - return Ok(ReactorReceiveOutcome::Failed(error)); + self.events = events; + if wake { + self.drain_wake()?; + if self.process_commands() { + return Ok(()); + } } } } -} -fn set_tcp_option(socket: &SocketEntry, value: TcpOptionValue) -> BrokerResult<()> { - if socket.kind != SocketKind::Tcp { - return Err(BrokerError::UnsupportedOperation); - } - match value { - TcpOptionValue::NoDelay(value) => sockopt::set_tcp_nodelay(&socket.socket, value), - TcpOptionValue::KeepAlive(value) => sockopt::set_socket_keepalive(&socket.socket, value), - _ => return Err(BrokerError::UnsupportedOperation), + fn next_cleanup_deadline(&self) -> Option { + self.tcp + .pending_guest_connections + .values() + .filter_map(|connection| connection.discard_deadline) + .chain( + self.tcp + .stale_mapped_connections + .values() + .filter_map(|stale| stale.deadline), + ) + .min() } - .map_err(broker_error_from_errno) -} -fn get_tcp_option(socket: &SocketEntry, name: TcpOptionName) -> BrokerResult { - if socket.kind != SocketKind::Tcp { - return Err(BrokerError::UnsupportedOperation); - } - match name { - TcpOptionName::NoDelay => sockopt::tcp_nodelay(&socket.socket) - .map(TcpOptionValue::NoDelay) - .map_err(broker_error_from_errno), - TcpOptionName::KeepAlive => sockopt::socket_keepalive(&socket.socket) - .map(TcpOptionValue::KeepAlive) - .map_err(broker_error_from_errno), - _ => Err(BrokerError::UnsupportedOperation), + fn expire_deadlined_state(&mut self, now: Instant) { + let Reactor { + tcp, + sessions, + retained_connectors, + .. + } = self; + tcp.pending_guest_connections.retain(|_, connection| { + let retain = !connection.discard_on_accept + || connection + .discard_deadline + .is_none_or(|deadline| deadline > now); + if !retain { + let session = sessions + .get_mut(&connection.session_id) + .expect("pending guest connection session state missing"); + session.pending_guest_connections = session + .pending_guest_connections + .checked_sub(1) + .expect("session pending guest connection count underflow"); + if connection.retained_connector.is_some() { + session.retained_connectors = session + .retained_connectors + .checked_sub(1) + .expect("session retained connector count underflow"); + *retained_connectors = retained_connectors + .checked_sub(1) + .expect("reactor retained connector count underflow"); + } + } + retain + }); + sessions.retain(|_, session| retain_session_state(session)); + tcp.stale_mapped_connections.retain(|_, stale| { + stale.retained_connector.is_some() + || stale.deadline.is_none_or(|deadline| deadline > now) + }); } -} -fn shutdown_socket( - socket: &mut SocketEntry, - mode: ShutdownMode, -) -> BrokerResult> { - if socket.kind == SocketKind::Udp - && matches!(mode, ShutdownMode::Abort | ShutdownMode::StopListening) - { - return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); - } - if mode == ShutdownMode::Abort { - sockopt::set_socket_linger(&socket.socket, Some(Duration::ZERO)) - .map_err(broker_error_from_errno)?; - return Ok(SocketOutcome::Completed(())); - } - let stop_listening = mode == ShutdownMode::StopListening; - if stop_listening { - if !socket.listening { - return Ok(SocketOutcome::Failed(SocketError::NotConnected)); + fn drain_wake(&self) -> core::result::Result<(), ReactorFailure> { + let mut value = [0_u8; size_of::()]; + loop { + match read(self.wake.as_ref(), &mut value) { + Ok(length) if length == value.len() => return Ok(()), + Ok(_) => return Err(ReactorFailure::Io(Errno::IO)), + Err(Errno::INTR) => {} + Err(Errno::AGAIN) => return Ok(()), + Err(error) => return Err(ReactorFailure::Io(error)), + } } - } else if socket.listening { - return Ok(SocketOutcome::Failed(SocketError::NotConnected)); } - let (mode, add, clear, shuts_down_read, shuts_down_write) = match mode { - ShutdownMode::Read => ( - LinuxShutdown::Read, - ReadinessFlags::READ, - ReadinessFlags::default(), - true, + + fn process_commands(&mut self) -> bool { + for _ in 0..MAX_QUEUED_SOCKET_COMMANDS { + let command = match self.commands.try_recv() { + Ok(command) => command, + Err(TryRecvError::Empty) => return false, + Err(TryRecvError::Disconnected) => return true, + }; + match command { + ReactorCommand::Create { + id, + session_id, + request, + readiness, + snapshot, + active, + response, + } => { + let outcome = self.create_socket(id, session_id, request, readiness, snapshot); + let created = outcome.is_ok(); + if created { + active.store(true, Ordering::Release); + } + if response.send(outcome).is_err() && created { + self.remove_socket(id); + } + } + ReactorCommand::Connect { + id, + address, + response, + } => { + let outcome = self.connect_socket(id, address); + let _ = response.send(outcome); + } + ReactorCommand::Bind { + id, + address, + response, + } => { + let outcome = self.bind_socket(id, address); + let _ = response.send(outcome); + } + ReactorCommand::Listen { + id, + backlog, + mapping, + response, + } => { + let outcome = self.listen_socket(id, backlog, mapping); + if response.send(outcome).is_err() { + // An unacknowledged listen may have committed a + // mapping. Retire the socket before processing any + // later command that could observe released core + // authority. + self.remove_socket(id); + } + } + ReactorCommand::Accept { + listener_id, + accepted_id, + readiness, + snapshot, + active, + response, + } => { + let outcome = self.accept_socket(listener_id, accepted_id, readiness, snapshot); + let accepted = matches!( + &outcome, + Ok(SocketOutcome::Completed(AcceptedEndpoints { .. })) + ); + if accepted { + active.store(true, Ordering::Release); + } + if response.send(outcome).is_err() && accepted { + self.remove_socket(accepted_id); + } + } + ReactorCommand::Send { id, data, response } => { + let outcome = self + .sockets + .get_mut(&id) + .ok_or(BrokerError::Internal) + .and_then(|socket| send_socket(socket, &data)); + let _ = response.send(outcome); + } + ReactorCommand::SendTo { + id, + data, + destination, + response, + } => { + let outcome = self + .sockets + .get_mut(&id) + .ok_or(BrokerError::Internal) + .and_then(|socket| send_to_socket(socket, &data, destination)); + let _ = response.send(outcome); + } + ReactorCommand::Receive { + id, + length, + flags, + peek_offset, + peek_length, + response, + } => { + let outcome = match self.sockets.get_mut(&id) { + Some(socket) => receive_socket( + socket, + &mut self.peek_cache, + id, + length, + flags, + peek_offset, + peek_length, + ), + None => Err(BrokerError::Internal), + }; + let _ = response.send(outcome); + } + ReactorCommand::ReceiveFrom { + id, + length, + flags, + response, + } => { + let outcome = self + .sockets + .get_mut(&id) + .ok_or(BrokerError::Internal) + .and_then(|socket| receive_from_socket(socket, length, flags)); + let _ = response.send(outcome); + } + ReactorCommand::Shutdown { id, mode, response } => { + if self + .peek_cache + .as_ref() + .is_some_and(|cache| cache.socket_id == id) + { + self.peek_cache = None; + } + let outcome = (|| { + let retired_listener = if mode == ShutdownMode::StopListening { + let socket = self.sockets.get(&id).ok_or(BrokerError::Internal)?; + (socket.kind == SocketKind::Tcp && socket.listening) + .then(|| { + local_socket_address(&socket.socket).and_then(|address| { + socket + .guest_local_address + .map(|guest_address| { + ( + address, + guest_address.port(), + socket.port_mapping_index, + ) + }) + .ok_or(BrokerError::Internal) + }) + }) + .transpose()? + } else { + None + }; + let outcome = if mode == ShutdownMode::StopListening { + self.stop_listening_socket(id)? + } else { + self.sockets + .get_mut(&id) + .ok_or(BrokerError::Internal) + .and_then(|socket| shutdown_socket(socket, mode))? + }; + if matches!(outcome, SocketOutcome::Completed(())) + && let Some((listener_address, guest_port, mapping_index)) = + retired_listener + { + if let Some(mapping_index) = mapping_index { + self.move_pending_guest_connections_for_listener( + listener_address, + mapping_index, + ); + } else { + self.remove_pending_guest_connections_for_listener( + listener_address, + ); + } + self.tcp.clear_host_address(guest_port, id)?; + update_snapshot( + self.sockets.get(&id).ok_or(BrokerError::Internal)?, + Some(SocketConnectionStatus::Failed(SocketError::NotConnected)), + ReadinessFlags::WRITE | ReadinessFlags::HANGUP, + )?; + } + Ok(outcome) + })(); + let _ = response.send(outcome); + } + ReactorCommand::SetTcpOption { + id, + value, + response, + } => { + let outcome = self + .sockets + .get_mut(&id) + .ok_or(BrokerError::Internal) + .and_then(|socket| set_tcp_option(socket, value)); + let _ = response.send(outcome); + } + ReactorCommand::GetTcpOption { id, name, response } => { + let outcome = self + .sockets + .get(&id) + .ok_or(BrokerError::Internal) + .and_then(|socket| get_tcp_option(socket, name)); + let _ = response.send(outcome); + } + ReactorCommand::Status { id, response } => { + let outcome = self + .sockets + .get_mut(&id) + .ok_or(BrokerError::Internal) + .and_then(status_socket); + let _ = response.send(outcome); + } + ReactorCommand::Close { id, response } => { + if self + .peek_cache + .as_ref() + .is_some_and(|cache| cache.socket_id == id) + { + self.peek_cache = None; + } + self.remove_socket(id); + let _ = response.send(()); + } + ReactorCommand::CloseSession { + session_id, + response, + } => { + if let Some(session) = self.sessions.get_mut(&session_id) { + session.closing = true; + } + self.retire_session_connectors(session_id); + if self + .sessions + .get(&session_id) + .is_some_and(|session| !retain_session_state(session)) + { + self.sessions.remove(&session_id); + } + let _ = response.send(()); + } + #[cfg(test)] + ReactorCommand::HostAddress { + kind, + guest_port, + response, + } => { + let host_address = (kind == SocketKind::Tcp) + .then(|| self.tcp.bindings.get(&guest_port)) + .flatten() + .and_then(|binding| binding.host_address); + let _ = response.send(host_address); + } + #[cfg(test)] + ReactorCommand::PendingGuestConnectionCount { response } => { + let _ = response.send(self.tcp.pending_guest_connections.len()); + } + #[cfg(test)] + ReactorCommand::StaleGuestConnectionCount { response } => { + let _ = response.send(self.tcp.stale_mapped_connections.len()); + } + #[cfg(test)] + ReactorCommand::RetainedConnectorCount { response } => { + let _ = response.send(self.retained_connectors); + } + ReactorCommand::Stop { response } => { + self.sockets.clear(); + self.tcp.bindings.clear(); + self.tcp.pending_guest_connections.clear(); + self.sessions.clear(); + let _ = response.send(()); + return true; + } + } + } + false + } + + fn create_socket( + &mut self, + id: u64, + session_id: SessionId, + request: CreateSocketRequest, + readiness: ReadinessRegistration, + snapshot: Arc>, + ) -> BrokerResult<()> { + if self + .sockets + .len() + .checked_add(self.retained_connectors) + .is_none_or(|count| count >= self.max_sockets) + { + return Err(BrokerError::ResourceExhausted); + } + let kind = socket_kind(request).ok_or(BrokerError::Internal)?; + if self.sockets.contains_key(&id) { + return Err(BrokerError::Internal); + } + let session = self.sessions.entry(session_id).or_default(); + if session.closing { + return Err(BrokerError::UnknownObject); + } + if session + .live_sockets + .checked_add(session.retained_connectors) + .is_none_or(|count| count >= self.max_sockets_per_session) + { + return Err(BrokerError::ResourceExhausted); + } + let (linux_type, protocol, epoll_events, initial_readiness) = match kind { + SocketKind::Tcp => ( + LinuxSocketType::STREAM, + ipproto::TCP, + idle_epoll_events(), + ReadinessFlags::default(), + ), + SocketKind::Udp => ( + LinuxSocketType::DGRAM, + ipproto::UDP, + active_epoll_events(), + ReadinessFlags::WRITE, + ), + }; + let socket = socket_with( + LinuxAddressFamily::INET, + linux_type, + LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, + Some(protocol), + ) + .map_err(broker_error_from_errno)?; + epoll::add( + &self.epoll, + &socket, + epoll::EventData::new_u64(id), + epoll_events, + ) + .map_err(broker_error_from_errno)?; + if initial_readiness != ReadinessFlags::default() { + snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness = initial_readiness; + readiness.publish(initial_readiness)?; + } + self.sockets.insert( + id, + SocketEntry { + socket, + session_id, + kind, + readiness, + snapshot, + read_shutdown: false, + write_shutdown: false, + peek_waitall_threshold: None, + listening: false, + abortive_close: false, + guest_local_address: None, + port_mapping_index: None, + mapping_fallback_socket: None, + tcp_no_delay: false, + tcp_keep_alive: false, + }, + ); + session.live_sockets = session + .live_sockets + .checked_add(1) + .ok_or(BrokerError::ResourceExhausted)?; + Ok(()) + } + + fn take_pending_guest_connection_for_accept( + &mut self, + listener_id: u64, + mapping_index: Option, + remote_address: SocketAddrV4, + local_address: SocketAddrV4, + ) -> AcceptedTcpPeer { + let now = Instant::now(); + self.expire_deadlined_state(now); + if let Some(mapping_index) = mapping_index + && self.take_stale_tcp_connection(mapping_index, remote_address, local_address) + { + return AcceptedTcpPeer::Stale; + } + let pending_connection = self + .tcp + .take_pending_guest_connection(remote_address, local_address); + let peer = if let Some(connection) = pending_connection { + self.finish_removed_pending_guest_connection(&connection); + drop(connection.retained_connector); + if connection.discard_on_accept { + AcceptedTcpPeer::Stale + } else if connection.listener_id == listener_id { + AcceptedTcpPeer::Guest(connection.guest_address) + } else { + AcceptedTcpPeer::Stale + } + } else { + AcceptedTcpPeer::Native(remote_address) + }; + self.sessions + .retain(|_, session| retain_session_state(session)); + peer + } + + fn take_stale_tcp_connection( + &mut self, + mapping_index: usize, + remote_address: SocketAddrV4, + local_address: SocketAddrV4, + ) -> bool { + let stale = self + .tcp + .stale_mapped_connections + .remove(&(mapping_index, remote_address, local_address)) + .or_else(|| { + self.tcp.stale_mapped_connections.remove(&( + mapping_index, + remote_address, + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, local_address.port()), + )) + }); + let Some(stale) = stale else { + return false; + }; + self.release_stale_tcp_connection(stale); + true + } + + fn release_stale_tcp_connection(&mut self, stale: StaleTcpConnection) { + if stale.retained_connector.is_some() { + self.retained_connectors = self + .retained_connectors + .checked_sub(1) + .expect("reactor retained connector count underflow"); + let remove_session = self + .sessions + .get_mut(&stale.session_id) + .is_some_and(|session| { + session.retained_connectors = session + .retained_connectors + .checked_sub(1) + .expect("session retained connector count underflow"); + !retain_session_state(session) + }); + if remove_session { + self.sessions.remove(&stale.session_id); + } + } + drop(stale); + } + + fn clear_stale_tcp_connections(&mut self, mapping_index: usize) { + let Reactor { + tcp, + sessions, + retained_connectors, + .. + } = self; + tcp.stale_mapped_connections.retain(|(index, _, _), stale| { + let retain = *index != mapping_index; + if !retain && stale.retained_connector.is_some() { + *retained_connectors = retained_connectors + .checked_sub(1) + .expect("reactor retained connector count underflow"); + if let Some(session) = sessions.get_mut(&stale.session_id) { + session.retained_connectors = session + .retained_connectors + .checked_sub(1) + .expect("session retained connector count underflow"); + } + } + retain + }); + sessions.retain(|_, session| retain_session_state(session)); + } + + fn retire_session_connectors(&mut self, session_id: SessionId) { + let mut pending_connections_released = 0; + for connection in self + .tcp + .pending_guest_connections + .values_mut() + .filter(|connection| connection.session_id == session_id) + { + if let Some(connector) = connection.retained_connector.take() { + let _ = sockopt::set_socket_linger(&connector, None); + drop(connector); + connection.discard_on_accept = true; + // An established child can remain queued after its connector + // closes. Keep its discard marker until accept or listener + // teardown so it can never fall through as a native peer. + connection.discard_deadline = None; + pending_connections_released += 1; + } + } + let mut stale_released = 0; + for stale in self + .tcp + .stale_mapped_connections + .values_mut() + .filter(|stale| stale.session_id == session_id) + { + if let Some(connector) = stale.retained_connector.take() { + let _ = sockopt::set_socket_linger(&connector, None); + drop(connector); + stale.deadline = None; + stale_released += 1; + } + } + self.retained_connectors = self + .retained_connectors + .checked_sub(pending_connections_released + stale_released) + .expect("reactor retained connector count underflow"); + if let Some(session) = self.sessions.get_mut(&session_id) { + session.retained_connectors = session + .retained_connectors + .checked_sub(pending_connections_released + stale_released) + .expect("session retained connector count underflow"); + } + } + + fn remove_pending_guest_connection_except( + &mut self, + excluded_session_id: Option, + remote_address: SocketAddrV4, + local_address: SocketAddrV4, + ) -> bool { + let exact = (remote_address, local_address); + let unspecified = ( + remote_address, + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, local_address.port()), + ); + let connection = if self.tcp.pending_guest_connections.contains_key(&exact) { + Some(exact) + } else if self + .tcp + .pending_guest_connections + .contains_key(&unspecified) + { + Some(unspecified) + } else { + None + }; + let removed = connection + .filter(|connection| { + self.tcp + .pending_guest_connections + .get(connection) + .is_some_and(|pending| Some(pending.session_id) != excluded_session_id) + }) + .and_then(|connection| self.tcp.pending_guest_connections.remove(&connection)); + if let Some(pending) = &removed { + self.finish_removed_pending_guest_connection(pending); + } + self.sessions + .retain(|_, session| retain_session_state(session)); + removed.is_some() + } + + fn remove_pending_guest_connections_for_listener(&mut self, listener_address: SocketAddrV4) { + let Reactor { + tcp, + sessions, + retained_connectors, + .. + } = self; + tcp.pending_guest_connections + .retain(|(_, destination), connection| { + let retain = destination.port() != listener_address.port() + || (!listener_address.ip().is_unspecified() + && destination.ip() != listener_address.ip()); + if !retain { + let session = sessions + .get_mut(&connection.session_id) + .expect("pending guest connection session state missing"); + session.pending_guest_connections = session + .pending_guest_connections + .checked_sub(1) + .expect("session pending guest connection count underflow"); + if connection.retained_connector.is_some() { + session.retained_connectors = session + .retained_connectors + .checked_sub(1) + .expect("session retained connector count underflow"); + *retained_connectors = retained_connectors + .checked_sub(1) + .expect("reactor retained connector count underflow"); + } + } + retain + }); + sessions.retain(|_, session| retain_session_state(session)); + } + + fn move_pending_guest_connections_for_listener( + &mut self, + listener_address: SocketAddrV4, + mapping_index: usize, + ) { + let retirement_deadline = Instant::now() + PENDING_CONNECT_DISCARD_LIFETIME; + let Reactor { + tcp, + sessions, + retained_connectors, + .. + } = self; + let BrokerTcpState { + pending_guest_connections, + stale_mapped_connections, + .. + } = tcp; + pending_guest_connections.retain(|connection_tuple, connection| { + let destination = connection_tuple.1; + let matches = destination.port() == listener_address.port() + && (listener_address.ip().is_unspecified() + || destination.ip() == listener_address.ip()); + if !matches { + return true; + } + let session = sessions + .get_mut(&connection.session_id) + .expect("pending guest connection session state missing"); + session.pending_guest_connections = session + .pending_guest_connections + .checked_sub(1) + .expect("session pending guest connection count underflow"); + let mut deadline = connection + .discard_deadline + .map(|deadline| deadline.max(retirement_deadline)); + let mut retained_connector = connection.retained_connector.take(); + if retained_connector.is_some() { + deadline = None; + } + match stale_mapped_connections.entry(( + mapping_index, + connection_tuple.0, + connection_tuple.1, + )) { + std::collections::hash_map::Entry::Vacant(entry) => { + entry.insert(StaleTcpConnection { + session_id: connection.session_id, + deadline, + retained_connector, + }); + } + std::collections::hash_map::Entry::Occupied(mut entry) => { + let existing = entry.get_mut(); + existing.deadline = match (existing.deadline, deadline) { + (None, _) | (_, None) => None, + (Some(existing), Some(deadline)) => Some(existing.max(deadline)), + }; + if existing.session_id == connection.session_id + && existing.retained_connector.is_none() + { + existing.retained_connector = retained_connector.take(); + } + if retained_connector.is_some() { + session.retained_connectors = session + .retained_connectors + .checked_sub(1) + .expect("session retained connector count underflow"); + *retained_connectors = retained_connectors + .checked_sub(1) + .expect("reactor retained connector count underflow"); + } + } + } + false + }); + sessions.retain(|_, session| retain_session_state(session)); + } + + fn accept_socket( + &mut self, + listener_id: u64, + accepted_id: u64, + readiness: ReadinessRegistration, + snapshot: Arc>, + ) -> BrokerResult> { + if self.sockets.contains_key(&accepted_id) { + return Err(BrokerError::Internal); + } + let ( + listener_session_id, + listener_tcp_no_delay, + listener_tcp_keep_alive, + local_address, + mapping_index, + ) = { + let listener = self + .sockets + .get(&listener_id) + .ok_or(BrokerError::Internal)?; + if listener.kind != SocketKind::Tcp || !listener.listening { + return Ok(SocketOutcome::Failed(SocketError::NotConnected)); + } + ( + listener.session_id, + listener.tcp_no_delay, + listener.tcp_keep_alive, + // Accepted connections inherit the trusted guest-local address, + // not the private host endpoint behind the listener. + listener.guest_local_address.ok_or(BrokerError::Internal)?, + listener.port_mapping_index, + ) + }; + let (socket, remote_address) = loop { + let (socket, remote_address) = loop { + let listener = self + .sockets + .get_mut(&listener_id) + .ok_or(BrokerError::Internal)?; + match acceptfrom_with( + &listener.socket, + LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, + ) { + Ok((socket, address)) => break (socket, address), + Err(Errno::INTR) => {} + Err(Errno::AGAIN) => { + clear_readiness(listener, ReadinessFlags::READ)?; + return Err(BrokerError::WouldBlock); + } + Err(error) => { + let _ = socket_operation_error_from_errno(error)?; + } + } + }; + let remote_address = + SocketAddrV4::try_from(remote_address.ok_or(BrokerError::Internal)?) + .map_err(|_| BrokerError::Internal)?; + let host_local_address = local_socket_address(&socket)?; + match self.take_pending_guest_connection_for_accept( + listener_id, + mapping_index, + remote_address, + host_local_address, + ) { + AcceptedTcpPeer::Guest(guest_address) => break (socket, guest_address), + AcceptedTcpPeer::Native(host_address) => break (socket, host_address), + AcceptedTcpPeer::Stale => drop(socket), + } + }; + if self + .sockets + .len() + .checked_add(self.retained_connectors) + .is_none_or(|count| count >= self.max_sockets) + || self + .sessions + .get(&listener_session_id) + .ok_or(BrokerError::Internal)? + .live_sockets + .checked_add( + self.sessions + .get(&listener_session_id) + .ok_or(BrokerError::Internal)? + .retained_connectors, + ) + .is_none_or(|count| count >= self.max_sockets_per_session) + { + return Err(BrokerError::ResourceExhausted); + } + let listener = self + .sockets + .get_mut(&listener_id) + .ok_or(BrokerError::Internal)?; + let no_wait = Timespec { + tv_sec: 0, + tv_nsec: 0, + }; + loop { + let mut poll_fd = [PollFd::new(&listener.socket, PollFlags::IN)]; + match poll(&mut poll_fd, Some(&no_wait)) { + Ok(_) if poll_fd[0].revents().contains(PollFlags::IN) => break, + Ok(_) => { + clear_readiness(listener, ReadinessFlags::READ)?; + break; + } + Err(Errno::INTR) => {} + Err(error) => return Err(broker_error_from_errno(error)), + } + } + epoll::add( + &self.epoll, + &socket, + epoll::EventData::new_u64(accepted_id), + active_epoll_events(), + ) + .map_err(broker_error_from_errno)?; + { + let mut snapshot = snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + snapshot.status = SocketConnectionStatus::Connected; + snapshot.local_address = Some(local_address); + snapshot.readiness = ReadinessFlags::WRITE; + } + readiness.publish(ReadinessFlags::WRITE)?; + self.sockets.insert( + accepted_id, + SocketEntry { + socket, + session_id: listener_session_id, + kind: SocketKind::Tcp, + readiness, + snapshot, + read_shutdown: false, + write_shutdown: false, + peek_waitall_threshold: None, + listening: false, + abortive_close: false, + guest_local_address: Some(local_address), + port_mapping_index: None, + mapping_fallback_socket: None, + tcp_no_delay: listener_tcp_no_delay, + tcp_keep_alive: listener_tcp_keep_alive, + }, + ); + let session = self + .sessions + .get_mut(&listener_session_id) + .ok_or(BrokerError::Internal)?; + session.live_sockets = session + .live_sockets + .checked_add(1) + .ok_or(BrokerError::ResourceExhausted)?; + Ok(SocketOutcome::Completed(AcceptedEndpoints { + remote_address, + })) + } + + fn fail_all_sockets(&mut self) { + for socket in self.sockets.values() { + let mut snapshot = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + snapshot.status = SocketConnectionStatus::Failed(SocketError::Other); + snapshot.readiness = ReadinessFlags::ERROR; + drop(snapshot); + // The readiness path may itself be why the reactor is failing. The + // cached terminal snapshot remains authoritative if mapping is + // no longer available. + let _ = socket.readiness.publish(ReadinessFlags::ERROR); + } + self.sockets.clear(); + self.sessions.clear(); + } +} + +fn connect_tcp_socket( + epoll_fd: &OwnedFd, + id: u64, + socket: &mut SocketEntry, + address: SocketAddrV4, +) -> core::result::Result<(SocketConnectionStatus, ReadinessFlags), PlatformConnectError> { + if let Err(error) = epoll::modify( + epoll_fd, + &socket.socket, + epoll::EventData::new_u64(id), + active_epoll_events(), + ) { + return Err(PlatformConnectError::PeerUnchanged( + broker_error_from_errno(error), + )); + } + let status = loop { + match connect(&socket.socket, &address) { + Ok(()) | Err(Errno::ISCONN) => break SocketConnectionStatus::Connected, + Err(Errno::INTR) => {} + Err(Errno::INPROGRESS | Errno::ALREADY) => { + break SocketConnectionStatus::Connecting; + } + Err(error) => { + let error = match socket_operation_error_from_errno(error) { + Ok(error) => error, + Err(error) => { + update_snapshot( + socket, + Some(SocketConnectionStatus::Failed(SocketError::Other)), + ReadinessFlags::ERROR, + ) + .map_err(PlatformConnectError::PeerIndeterminate)?; + return Err(PlatformConnectError::PeerIndeterminate(error)); + } + }; + break SocketConnectionStatus::Failed(error); + } + } + }; + let readiness = match status { + SocketConnectionStatus::Connected | SocketConnectionStatus::Connecting => { + // The guest-local address is assigned by bind; the private host + // endpoint behind it is never reported to the guest. + if socket.guest_local_address.is_none() { + return Err(PlatformConnectError::PeerIndeterminate( + BrokerError::Internal, + )); + } + if status == SocketConnectionStatus::Connected { + ReadinessFlags::WRITE + } else { + ReadinessFlags::default() + } + } + SocketConnectionStatus::Failed(_) => ReadinessFlags::ERROR, + SocketConnectionStatus::Unconnected => ReadinessFlags::default(), + _ => { + return Err(PlatformConnectError::PeerIndeterminate( + BrokerError::Internal, + )); + } + }; + Ok((status, readiness)) +} + +fn connect_datagram_socket( + socket: &mut SocketEntry, + address: SocketAddrV4, +) -> core::result::Result { + loop { + match connect(&socket.socket, &address) { + Ok(()) | Err(Errno::ISCONN) => { + let local_address = local_socket_address(&socket.socket) + .map_err(PlatformConnectError::PeerIndeterminate)?; + socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .local_address = Some(local_address); + let readiness = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + let readiness = if socket.write_shutdown { + ReadinessFlags(readiness.0 & !ReadinessFlags::WRITE.0) + } else { + readiness | ReadinessFlags::WRITE + }; + update_snapshot(socket, Some(SocketConnectionStatus::Connected), readiness) + .map_err(PlatformConnectError::PeerIndeterminate)?; + return Ok(SocketConnectionStatus::Connected); + } + Err(Errno::INTR) => {} + Err(error) => { + update_local_address(socket).map_err(PlatformConnectError::PeerIndeterminate)?; + let error = socket_operation_error_from_errno(error) + .map_err(PlatformConnectError::PeerIndeterminate)?; + return Ok(SocketConnectionStatus::Failed(error)); + } + } + } +} + +fn bind_host_socket( + socket: &mut SocketEntry, + address: SocketAddrV4, +) -> BrokerResult> { + loop { + match bind(&socket.socket, &address) { + Ok(()) => { + let local_address = local_socket_address(&socket.socket)?; + return Ok(SocketOutcome::Completed(local_address)); + } + Err(Errno::INTR) => {} + Err(error) => { + return Ok(SocketOutcome::Failed(socket_operation_error_from_errno( + error, + )?)); + } + } + } +} + +fn delete_epoll_registration(epoll_fd: &OwnedFd, socket: &OwnedFd) -> bool { + loop { + match epoll::delete(epoll_fd, socket) { + Ok(()) | Err(Errno::NOENT) => return true, + Err(Errno::INTR) => {} + Err(_) => return false, + } + } +} + +fn quiesce_tcp_listener(socket: &OwnedFd) -> BrokerResult<()> { + loop { + match listen(socket, 0) { + Ok(()) => return Ok(()), + Err(Errno::INTR) => {} + Err(error) => return Err(broker_error_from_errno(error)), + } + } +} + +/// Returns whether a host IPv4 address names one of this machine's interfaces. +fn host_ipv4_address_is_local(address: Ipv4Addr) -> BrokerResult { + let socket = socket_with( + LinuxAddressFamily::INET, + LinuxSocketType::DGRAM, + LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, + Some(ipproto::UDP), + ) + .map_err(broker_error_from_errno)?; + match bind(&socket, &SocketAddrV4::new(address, 0)) { + Ok(()) => Ok(true), + Err(Errno::ADDRNOTAVAIL) => Ok(false), + Err(error) => Err(broker_error_from_errno(error)), + } +} + +fn listen_tcp_socket( + epoll_fd: &OwnedFd, + id: u64, + socket: &mut SocketEntry, + backlog: u32, + host_mapped: bool, +) -> BrokerResult> { + if socket.kind != SocketKind::Tcp { + return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); + } + let backlog = i32::try_from(backlog).map_err(|_| BrokerError::UnsupportedOperation)?; + let was_listening = socket.listening; + if !was_listening { + epoll::modify( + epoll_fd, + &socket.socket, + epoll::EventData::new_u64(id), + active_epoll_events(), + ) + .map_err(broker_error_from_errno)?; + } + if host_mapped && let Err(error) = sockopt::set_socket_reuseaddr(&socket.socket, true) { + if !was_listening { + epoll::modify( + epoll_fd, + &socket.socket, + epoll::EventData::new_u64(id), + idle_epoll_events(), + ) + .map_err(broker_error_from_errno)?; + } + return Err(broker_error_from_errno(error)); + } + loop { + match listen(&socket.socket, backlog) { + Ok(()) => break, + Err(Errno::INTR) => {} + Err(error) => { + let reuse_error = host_mapped + .then(|| sockopt::set_socket_reuseaddr(&socket.socket, false)) + .transpose() + .err(); + let epoll_error = if was_listening { + None + } else { + epoll::modify( + epoll_fd, + &socket.socket, + epoll::EventData::new_u64(id), + idle_epoll_events(), + ) + .err() + }; + if let Some(error) = reuse_error.or(epoll_error) { + return Err(broker_error_from_errno(error)); + } + return Ok(SocketOutcome::Failed(socket_operation_error_from_errno( + error, + )?)); + } + } + } + socket.listening = true; + let mut snapshot = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + if !was_listening { + snapshot.readiness = ReadinessFlags::default(); + } + Ok(SocketOutcome::Completed(())) +} + +fn send_socket(socket: &mut SocketEntry, data: &[u8]) -> BrokerResult> { + if socket.kind != SocketKind::Tcp { + return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); + } + if socket.write_shutdown { + return Ok(SocketOutcome::Failed(SocketError::Other)); + } + loop { + match send(&socket.socket, data, LinuxSendFlags::NOSIGNAL) { + Ok(sent) => return Ok(SocketOutcome::Completed(sent)), + Err(Errno::INTR) => {} + Err(Errno::AGAIN) => { + clear_readiness(socket, ReadinessFlags::WRITE)?; + return Err(BrokerError::WouldBlock); + } + Err(error) => { + let error = socket_operation_error_from_errno(error)?; + consume_synchronous_error(socket)?; + return Ok(SocketOutcome::Failed(error)); + } + } + } +} + +fn send_to_socket( + socket: &mut SocketEntry, + data: &[u8], + destination: Option, +) -> BrokerResult> { + if socket.kind != SocketKind::Udp || data.len() > MAX_UDP_DATAGRAM_SIZE as usize { + return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); + } + if socket.write_shutdown { + return Ok(SocketOutcome::Failed(SocketError::Other)); + } + loop { + let result = match destination { + Some(address) => sendto(&socket.socket, data, LinuxSendFlags::NOSIGNAL, &address), + None => send(&socket.socket, data, LinuxSendFlags::NOSIGNAL), + }; + match result { + Ok(sent) if sent == data.len() => { + update_local_address(socket)?; + return Ok(SocketOutcome::Completed(sent)); + } + Ok(_) => { + update_local_address(socket)?; + return Err(BrokerError::Internal); + } + Err(Errno::INTR) => {} + Err(Errno::AGAIN) => { + update_local_address(socket)?; + clear_readiness(socket, ReadinessFlags::WRITE)?; + return Err(BrokerError::WouldBlock); + } + Err(error) => { + update_local_address(socket)?; + let error = socket_operation_error_from_errno(error)?; + consume_synchronous_error(socket)?; + return Ok(SocketOutcome::Failed(error)); + } + } + } +} + +fn receive_socket( + socket: &mut SocketEntry, + peek_cache: &mut Option, + socket_id: u64, + length: usize, + flags: ReceiveFlags, + peek_offset: usize, + peek_length: usize, +) -> BrokerResult { + if socket.kind != SocketKind::Tcp { + return Ok(ReactorReceiveOutcome::Failed(SocketError::InvalidArgument)); + } + let peek = flags.contains(ReceiveFlags::PEEK); + if !peek { + if peek_offset != 0 || peek_length != 0 { + return Err(BrokerError::UnsupportedOperation); + } + if peek_cache + .as_ref() + .is_some_and(|cache| cache.socket_id == socket_id) + { + *peek_cache = None; + } + return receive_socket_once(socket, zeroed_vec(length)?, LinuxRecvFlags::empty()); + } + + let peek_end = peek_offset + .checked_add(length) + .ok_or(BrokerError::UnsupportedOperation)?; + let canonical_length = peek_length + .checked_sub(peek_offset) + .map(|remaining| remaining.min(MAX_SOCKET_TRANSFER_SIZE as usize)); + if !peek_offset.is_multiple_of(MAX_SOCKET_TRANSFER_SIZE as usize) + || canonical_length != Some(length) + || peek_length < peek_end + || peek_length > MAX_SOCKET_PEEK_SIZE as usize + { + return Err(BrokerError::UnsupportedOperation); + } + if flags.contains(ReceiveFlags::WAITALL) { + let snapshot = *socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + let terminal = socket.read_shutdown + || snapshot.readiness.contains(ReadinessFlags::HANGUP) + || snapshot.readiness.contains(ReadinessFlags::ERROR); + if snapshot.status == SocketConnectionStatus::Connected + && !terminal + && ioctl_fionread(&socket.socket).map_err(broker_error_from_errno)? + < peek_length.try_into().map_err(|_| BrokerError::Internal)? + { + socket.peek_waitall_threshold = Some( + socket + .peek_waitall_threshold + .map_or(peek_length, |threshold| threshold.min(peek_length)), + ); + return Err(BrokerError::WouldBlock); + } + } + + let refresh = peek_offset == 0 + || !peek_cache.as_ref().is_some_and(|cache| { + cache.socket_id == socket_id && cache.requested_length == peek_length + }); + if refresh { + *peek_cache = None; + let flags = if flags.contains(ReceiveFlags::WAITALL) { + LinuxRecvFlags::PEEK | LinuxRecvFlags::WAITALL + } else { + LinuxRecvFlags::PEEK + }; + match receive_socket_once(socket, zeroed_vec(peek_length)?, flags)? { + ReactorReceiveOutcome::Received(data) => { + *peek_cache = Some(PeekCache { + socket_id, + requested_length: peek_length, + data, + }); + } + outcome => return Ok(outcome), + } + } + + let cache = peek_cache.as_ref().ok_or(BrokerError::Internal)?; + if cache.data.len() <= peek_offset { + let readiness = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + let terminal = socket.read_shutdown + || readiness.contains(ReadinessFlags::HANGUP) + || readiness.contains(ReadinessFlags::ERROR); + *peek_cache = None; + return if terminal { + Ok(ReactorReceiveOutcome::EndOfStream) + } else { + Err(BrokerError::WouldBlock) + }; + } + let end = peek_end.min(cache.data.len()); + let mut data = Vec::new(); + data.try_reserve_exact(end - peek_offset) + .map_err(|_| BrokerError::OutOfMemory)?; + data.extend_from_slice(&cache.data[peek_offset..end]); + if end < peek_end || end == peek_length { + *peek_cache = None; + } + Ok(ReactorReceiveOutcome::Received(data)) +} + +fn receive_from_socket( + socket: &mut SocketEntry, + length: usize, + flags: ReceiveFromFlags, +) -> BrokerResult { + if socket.kind != SocketKind::Udp + || length > MAX_UDP_DATAGRAM_SIZE as usize + || flags.has_unsupported_bits() + { + return Ok(ReactorReceiveFromOutcome::Failed( + SocketError::InvalidArgument, + )); + } + if socket.read_shutdown { + return Ok(ReactorReceiveFromOutcome::Failed(SocketError::NotConnected)); + } + let mut data = zeroed_vec(length)?; + let mut linux_flags = LinuxRecvFlags::TRUNC; + if flags.contains(ReceiveFromFlags::PEEK) { + linux_flags |= LinuxRecvFlags::PEEK; + } + loop { + match recvfrom(&socket.socket, data.as_mut_slice(), linux_flags) { + Ok((received, datagram_length, address)) => { + let source_address = SocketAddrV4::try_from(address.ok_or(BrokerError::Internal)?) + .map_err(|_| BrokerError::Internal)?; + data.truncate(received); + update_local_address(socket)?; + return Ok(ReactorReceiveFromOutcome::Received { + data, + datagram_length, + source_address, + }); + } + Err(Errno::INTR) => {} + Err(Errno::AGAIN) => { + clear_readiness(socket, ReadinessFlags::READ)?; + return Err(BrokerError::WouldBlock); + } + Err(error) => { + let error = socket_operation_error_from_errno(error)?; + consume_synchronous_error(socket)?; + return Ok(ReactorReceiveFromOutcome::Failed(error)); + } + } + } +} + +fn zeroed_vec(length: usize) -> BrokerResult> { + let mut data = Vec::new(); + data.try_reserve_exact(length) + .map_err(|_| BrokerError::OutOfMemory)?; + data.resize(length, 0); + Ok(data) +} + +fn receive_socket_once( + socket: &mut SocketEntry, + mut data: Vec, + flags: LinuxRecvFlags, +) -> BrokerResult { + loop { + match recv(&socket.socket, data.as_mut_slice(), flags) { + Ok((_buffer, 0)) => { + let readiness = if socket.read_shutdown { + ReadinessFlags::READ + } else { + ReadinessFlags::READ | ReadinessFlags::HANGUP + }; + add_readiness(socket, readiness)?; + return Ok(ReactorReceiveOutcome::EndOfStream); + } + Ok((_buffer, received)) => { + data.truncate(received); + let terminal_readable = socket.read_shutdown + || socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness + .contains(ReadinessFlags::HANGUP); + if !flags.contains(LinuxRecvFlags::PEEK) + && !terminal_readable + && ioctl_fionread(&socket.socket).map_err(broker_error_from_errno)? == 0 + { + clear_readiness(socket, ReadinessFlags::READ)?; + } + return Ok(ReactorReceiveOutcome::Received(data)); + } + Err(Errno::INTR) => {} + Err(Errno::AGAIN) => { + clear_readiness(socket, ReadinessFlags::READ)?; + return Err(BrokerError::WouldBlock); + } + Err(error) => { + let error = socket_operation_error_from_errno(error)?; + consume_synchronous_error(socket)?; + return Ok(ReactorReceiveOutcome::Failed(error)); + } + } + } +} + +fn set_tcp_option(socket: &mut SocketEntry, value: TcpOptionValue) -> BrokerResult<()> { + if socket.kind != SocketKind::Tcp { + return Err(BrokerError::UnsupportedOperation); + } + // Cached values are reapplied when a mapped endpoint replaces this + // socket's descriptor. + match value { + TcpOptionValue::NoDelay(value) => { + sockopt::set_tcp_nodelay(&socket.socket, value).map_err(broker_error_from_errno)?; + socket.tcp_no_delay = value; + } + TcpOptionValue::KeepAlive(value) => { + sockopt::set_socket_keepalive(&socket.socket, value) + .map_err(broker_error_from_errno)?; + socket.tcp_keep_alive = value; + } + _ => return Err(BrokerError::UnsupportedOperation), + } + Ok(()) +} + +fn apply_tcp_options(socket: &OwnedFd, no_delay: bool, keep_alive: bool) -> BrokerResult<()> { + sockopt::set_tcp_nodelay(socket, no_delay).map_err(broker_error_from_errno)?; + sockopt::set_socket_keepalive(socket, keep_alive).map_err(broker_error_from_errno) +} + +fn get_tcp_option(socket: &SocketEntry, name: TcpOptionName) -> BrokerResult { + if socket.kind != SocketKind::Tcp { + return Err(BrokerError::UnsupportedOperation); + } + match name { + TcpOptionName::NoDelay => sockopt::tcp_nodelay(&socket.socket) + .map(TcpOptionValue::NoDelay) + .map_err(broker_error_from_errno), + TcpOptionName::KeepAlive => sockopt::socket_keepalive(&socket.socket) + .map(TcpOptionValue::KeepAlive) + .map_err(broker_error_from_errno), + _ => Err(BrokerError::UnsupportedOperation), + } +} + +fn count_session_stale_connections<'a>( + stale_connections: impl Iterator, + session_id: SessionId, +) -> usize { + stale_connections + .filter(|stale| stale.session_id == session_id) + .count() +} + +fn shutdown_socket( + socket: &mut SocketEntry, + mode: ShutdownMode, +) -> BrokerResult> { + if socket.kind == SocketKind::Udp + && matches!(mode, ShutdownMode::Abort | ShutdownMode::StopListening) + { + return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); + } + if mode == ShutdownMode::Abort { + sockopt::set_socket_linger(&socket.socket, Some(Duration::ZERO)) + .map_err(broker_error_from_errno)?; + socket.abortive_close = true; + return Ok(SocketOutcome::Completed(())); + } + let stop_listening = mode == ShutdownMode::StopListening; + if stop_listening { + if !socket.listening { + return Ok(SocketOutcome::Failed(SocketError::NotConnected)); + } + } else if socket.listening { + return Ok(SocketOutcome::Failed(SocketError::NotConnected)); + } + let (mode, add, clear, shuts_down_read, shuts_down_write) = match mode { + ShutdownMode::Read => ( + LinuxShutdown::Read, + ReadinessFlags::READ, + ReadinessFlags::default(), + true, false, ), ShutdownMode::Write => ( @@ -1634,548 +3950,1934 @@ fn shutdown_socket( ), _ => return Err(BrokerError::UnsupportedOperation), }; - loop { - match shutdown(&socket.socket, mode) { - Ok(()) => {} - Err(Errno::INTR) => continue, - // Linux applies directional shutdown to unconnected UDP sockets even - // though it reports ENOTCONN. - Err(Errno::NOTCONN) if socket.kind == SocketKind::Udp => {} - Err(Errno::NOTCONN) => { - return Ok(SocketOutcome::Failed(SocketError::NotConnected)); - } - Err(error) => { - let error = socket_operation_error_from_errno(error)?; - return Ok(SocketOutcome::Failed(error)); - } + loop { + match shutdown(&socket.socket, mode) { + Ok(()) => {} + Err(Errno::INTR) => continue, + // Linux applies directional shutdown to unconnected UDP sockets even + // though it reports ENOTCONN. + Err(Errno::NOTCONN) if socket.kind == SocketKind::Udp => {} + Err(Errno::NOTCONN) => { + return Ok(SocketOutcome::Failed(SocketError::NotConnected)); + } + Err(error) => { + let error = socket_operation_error_from_errno(error)?; + return Ok(SocketOutcome::Failed(error)); + } + } + if stop_listening { + socket.listening = false; + socket.read_shutdown = true; + socket.peek_waitall_threshold = None; + return Ok(SocketOutcome::Completed(())); + } + socket.read_shutdown |= shuts_down_read; + socket.write_shutdown |= shuts_down_write; + let republish_readiness = shuts_down_read && socket.peek_waitall_threshold.take().is_some(); + if clear.0 != 0 { + clear_readiness(socket, clear)?; + } + if add.0 != 0 { + add_readiness(socket, add)?; + } + if republish_readiness { + let readiness = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + socket.readiness.republish(readiness)?; + } + return Ok(SocketOutcome::Completed(())); + } +} + +fn handle_socket_event(socket: &mut SocketEntry, events: epoll::EventFlags) -> BrokerResult { + if socket.listening { + update_snapshot(socket, None, readiness_from_epoll(socket, events))?; + return Ok(false); + } + if socket.kind == SocketKind::Udp { + update_snapshot(socket, None, readiness_from_epoll(socket, events))?; + return Ok(false); + } + let republish_readiness = if events.contains(epoll::EventFlags::IN) + && let Some(threshold) = socket.peek_waitall_threshold + { + let threshold_reached = ioctl_fionread(&socket.socket) + .ok() + .and_then(|available| usize::try_from(available).ok()) + .is_none_or(|available| available >= threshold); + if threshold_reached { + socket.peek_waitall_threshold = None; + } + threshold_reached + } else { + false + }; + let status = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .status; + let failed_connector = match status { + SocketConnectionStatus::Unconnected => false, + SocketConnectionStatus::Connecting => { + matches!( + complete_connect(socket, events)?, + SocketConnectionStatus::Failed(_) + ) + } + SocketConnectionStatus::Connected => { + update_snapshot(socket, None, readiness_from_epoll(socket, events))?; + false + } + SocketConnectionStatus::Failed(SocketError::NotConnected) if socket.read_shutdown => { + update_snapshot(socket, None, ReadinessFlags::WRITE | ReadinessFlags::HANGUP)?; + false + } + SocketConnectionStatus::Failed(_) => { + update_snapshot(socket, None, ReadinessFlags::ERROR)?; + false + } + _ => return Err(BrokerError::Internal), + }; + if republish_readiness { + let readiness = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + socket.readiness.republish(readiness)?; + } + Ok(failed_connector) +} + +fn complete_connect( + socket: &mut SocketEntry, + events: epoll::EventFlags, +) -> BrokerResult { + let status = match sockopt::socket_error(&socket.socket) { + Ok(Ok(())) => match getpeername(&socket.socket) { + Ok(Some(_)) => SocketConnectionStatus::Connected, + Ok(None) | Err(Errno::NOTCONN) => SocketConnectionStatus::Connecting, + Err(error) => SocketConnectionStatus::Failed(socket_error_from_errno(error)), + }, + Ok(Err(error)) | Err(error) => { + SocketConnectionStatus::Failed(socket_error_from_errno(error)) + } + }; + let readiness = match status { + SocketConnectionStatus::Connected => { + readiness_from_epoll(socket, events) | ReadinessFlags::WRITE + } + SocketConnectionStatus::Connecting => ReadinessFlags::default(), + SocketConnectionStatus::Failed(_) => ReadinessFlags::ERROR, + _ => return Err(BrokerError::Internal), + }; + update_snapshot(socket, Some(status), readiness)?; + Ok(status) +} + +fn take_socket_error(socket: &SocketEntry) -> BrokerResult> { + match sockopt::socket_error(&socket.socket) { + Ok(Ok(())) => Ok(None), + Ok(Err(error)) | Err(error) => socket_operation_error_from_errno(error).map(Some), + } +} + +fn local_socket_address(socket: &OwnedFd) -> BrokerResult { + match getsockname(socket) { + Ok(address) => SocketAddrV4::try_from(address).map_err(|_| BrokerError::Internal), + Err(_) => Err(BrokerError::Internal), + } +} + +fn update_local_address(socket: &SocketEntry) -> BrokerResult<()> { + let needs_address = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .local_address + .is_none(); + if needs_address { + let address = local_socket_address(&socket.socket)?; + if address.port() != 0 { + socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .local_address = Some(address); + } + } + Ok(()) +} + +fn status_socket(socket: &mut SocketEntry) -> BrokerResult { + let query_socket_error = { + let snapshot = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + (snapshot.status == SocketConnectionStatus::Connected || socket.kind == SocketKind::Udp) + && snapshot.readiness.contains(ReadinessFlags::ERROR) + }; + let socket_error = if query_socket_error { + take_socket_error(socket)? + } else { + None + }; + let (response, readiness, republish_error) = { + let mut snapshot = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + let cached_error = snapshot.pending_error.take(); + let (pending_error, next_pending_error) = shift_pending_error(cached_error, socket_error); + snapshot.pending_error = next_pending_error; + if pending_error.is_some() && next_pending_error.is_none() { + snapshot.readiness = ReadinessFlags(snapshot.readiness.0 & !ReadinessFlags::ERROR.0); + } + ( + SocketStatusResponse { + status: snapshot.status, + local_address: snapshot.local_address, + pending_error, + }, + snapshot.readiness, + next_pending_error.is_some(), + ) + }; + if republish_error { + socket.readiness.republish(readiness)?; + } else if response.pending_error.is_some() { + socket.readiness.publish(readiness)?; + } + Ok(response) +} + +fn shift_pending_error( + cached_error: Option, + socket_error: Option, +) -> (Option, Option) { + match cached_error { + Some(error) => (Some(error), socket_error), + None => (socket_error, None), + } +} + +fn readiness_from_epoll(socket: &SocketEntry, events: epoll::EventFlags) -> ReadinessFlags { + let mut readiness = ReadinessFlags::default(); + if events.contains(epoll::EventFlags::IN) { + readiness = readiness | ReadinessFlags::READ; + } + if events.contains(epoll::EventFlags::OUT) && !socket.write_shutdown { + readiness = readiness | ReadinessFlags::WRITE; + } + if socket.kind == SocketKind::Tcp + && !socket.read_shutdown + && events.intersects(epoll::EventFlags::RDHUP | epoll::EventFlags::HUP) + { + readiness = readiness | ReadinessFlags::READ | ReadinessFlags::HANGUP; + } + + if events.contains(epoll::EventFlags::ERR) { + readiness = readiness | ReadinessFlags::ERROR; + } + let previous = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + ReadinessFlags(readiness.0 | previous.0) +} + +const fn socket_kind(request: CreateSocketRequest) -> Option { + match ( + request.address_family, + request.socket_type, + request.protocol, + ) { + (AddressFamily::Ipv4, SocketType::Stream, IpProtocol::Tcp) => Some(SocketKind::Tcp), + (AddressFamily::Ipv4, SocketType::Datagram, IpProtocol::Udp) => Some(SocketKind::Udp), + _ => None, + } +} + +fn idle_epoll_events() -> epoll::EventFlags { + epoll::EventFlags::RDHUP | epoll::EventFlags::ET +} + +fn active_epoll_events() -> epoll::EventFlags { + // Cached readiness turns these edge-triggered kernel events into the + // level-triggered snapshots consumed by the broker protocol. + epoll::EventFlags::IN + | epoll::EventFlags::OUT + | epoll::EventFlags::RDHUP + | epoll::EventFlags::ET +} + +fn update_snapshot( + socket: &SocketEntry, + status: Option, + readiness: ReadinessFlags, +) -> BrokerResult<()> { + let readiness_changed = { + let mut snapshot = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + if let Some(status) = status { + snapshot.status = status; + } + let changed = snapshot.readiness != readiness; + snapshot.readiness = readiness; + changed + }; + if readiness_changed { + socket.readiness.publish(readiness)?; + } + Ok(()) +} + +fn add_readiness(socket: &SocketEntry, readiness: ReadinessFlags) -> BrokerResult<()> { + let current = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + update_snapshot(socket, None, ReadinessFlags(current.0 | readiness.0)) +} + +fn clear_readiness(socket: &SocketEntry, readiness: ReadinessFlags) -> BrokerResult<()> { + let current = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + update_snapshot(socket, None, ReadinessFlags(current.0 & !readiness.0)) +} + +fn consume_synchronous_error(socket: &SocketEntry) -> BrokerResult<()> { + let query_socket_error = { + let snapshot = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + if !can_consume_synchronous_error(socket.kind, snapshot.status) { + return Ok(()); } - if stop_listening { - socket.listening = false; - socket.read_shutdown = true; - socket.peek_waitall_threshold = None; - update_snapshot( - socket, - Some(SocketConnectionStatus::Failed(SocketError::NotConnected)), - ReadinessFlags::WRITE | ReadinessFlags::HANGUP, - )?; - return Ok(SocketOutcome::Completed(())); + snapshot.pending_error.is_none() + }; + let socket_error = if query_socket_error { + take_socket_error(socket)? + } else { + None + }; + let (readiness, changed) = { + let mut snapshot = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned"); + if snapshot.pending_error.is_none() { + snapshot.pending_error = socket_error; } - socket.read_shutdown |= shuts_down_read; - socket.write_shutdown |= shuts_down_write; - let republish_readiness = shuts_down_read && socket.peek_waitall_threshold.take().is_some(); - if clear.0 != 0 { - clear_readiness(socket, clear)?; + let readiness = if snapshot.pending_error.is_some() { + snapshot.readiness | ReadinessFlags::ERROR + } else { + ReadinessFlags(snapshot.readiness.0 & !ReadinessFlags::ERROR.0) + }; + let changed = readiness != snapshot.readiness; + snapshot.readiness = readiness; + (readiness, changed) + }; + if changed { + socket.readiness.publish(readiness)?; + } + Ok(()) +} + +const fn can_consume_synchronous_error(kind: SocketKind, status: SocketConnectionStatus) -> bool { + matches!(kind, SocketKind::Udp) || matches!(status, SocketConnectionStatus::Connected) +} + +const fn socket_error_from_errno(error: Errno) -> SocketError { + match error { + Errno::CONNREFUSED => SocketError::ConnectionRefused, + Errno::CONNRESET | Errno::PIPE => SocketError::ConnectionReset, + Errno::CONNABORTED => SocketError::ConnectionAborted, + Errno::NETUNREACH => SocketError::NetworkUnreachable, + Errno::HOSTUNREACH => SocketError::HostUnreachable, + Errno::TIMEDOUT => SocketError::TimedOut, + Errno::ADDRINUSE => SocketError::AddressInUse, + Errno::ADDRNOTAVAIL => SocketError::AddressNotAvailable, + Errno::NOTCONN => SocketError::NotConnected, + Errno::INVAL => SocketError::InvalidArgument, + _ => SocketError::Other, + } +} + +const fn socket_operation_error_from_errno(error: Errno) -> BrokerResult { + match broker_resource_error_from_errno(error) { + Some(error) => Err(error), + None => Ok(socket_error_from_errno(error)), + } +} + +const fn broker_error_from_errno(error: Errno) -> BrokerError { + match broker_resource_error_from_errno(error) { + Some(error) => error, + None => BrokerError::Internal, + } +} + +const fn broker_resource_error_from_errno(error: Errno) -> Option { + match error { + Errno::NOMEM => Some(BrokerError::OutOfMemory), + Errno::MFILE | Errno::NFILE | Errno::NOBUFS | Errno::NOSPC => { + Some(BrokerError::ResourceExhausted) } - if add.0 != 0 { - add_readiness(socket, add)?; + _ => None, + } +} + +#[cfg(test)] +mod tests { + use std::io::{Read as _, Write as _}; + use std::net::Ipv4Addr; + use std::net::{Shutdown, TcpListener, TcpStream, UdpSocket}; + use std::sync::mpsc::{Receiver, Sender, channel}; + use std::time::{Duration, Instant}; + + use super::*; + use litebox_broker_core::readiness::ReadinessSink; + use litebox_broker_core::{ + BrokerCore, BrokerCoreLimits, BrokerSession, CallerCredential, DestinationPortRange, + DestinationRule, Ipv4Cidr, ObjectRights, PolicyEngine, SocketPolicy, + }; + use litebox_broker_protocol::ObjectHandle; + use litebox_broker_protocol::socket::{Ipv4Address, Port, ReceiveSocketResponse}; + + const TEST_TIMEOUT: Duration = Duration::from_secs(5); + const FIRST_GUEST_EPHEMERAL_PORT: u16 = 49152; + + fn test_reactor( + max_sockets_per_session: usize, + tcp_port_mappings: &[TcpPortMapping], + ) -> Reactor { + let epoll = epoll::create(epoll::CreateFlags::CLOEXEC).unwrap(); + let wake = Arc::new(eventfd(0, EventfdFlags::CLOEXEC | EventfdFlags::NONBLOCK).unwrap()); + let (_, commands) = sync_channel(1); + Reactor { + epoll, + wake, + broker_ipv4_address: Ipv4Addr::LOCALHOST, + commands, + sockets: HashMap::new(), + tcp: BrokerTcpState::default(), + sessions: HashMap::new(), + port_mappings: tcp_port_mappings + .iter() + .copied() + .map(|mapping| PortMappingState { + mapping, + reservation: None, + reservation_registered: false, + owned_by: None, + }) + .collect(), + max_sockets: MAX_TRACKED_GUEST_CONNECTIONS, + max_sockets_per_session, + retained_connectors: 0, + peek_cache: None, + events: Vec::new(), } - if republish_readiness { - let readiness = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness; - socket.readiness.republish(readiness)?; + } + + #[allow(clippy::too_many_arguments)] + fn insert_test_pending_guest_connection( + reactor: &mut Reactor, + session_id: SessionId, + guest_port: u16, + socket_id: u64, + host_address: SocketAddrV4, + listener_address: SocketAddrV4, + listener_id: u64, + mapping_index: Option, + ) { + reactor.sessions.entry(session_id).or_default(); + reactor + .reserve_pending_guest_connection(session_id) + .unwrap(); + reactor + .tcp + .insert_binding( + guest_port, + GuestPortBinding { + socket_id, + guest_address: SocketAddrV4::new(Ipv4Addr::LOCALHOST, guest_port), + host_address: Some(host_address), + host_peer_address: None, + host_peer_mapping_index: None, + host_mapped: false, + listening: false, + }, + ) + .unwrap(); + let (host_address, guest_address) = reactor + .tcp + .set_host_peer_address(guest_port, socket_id, listener_address, mapping_index) + .unwrap(); + reactor + .insert_pending_guest_connection( + session_id, + (host_address, listener_address), + guest_address, + listener_id, + ) + .unwrap(); + } + + #[derive(Clone, Copy, Debug, PartialEq, Eq)] + struct ReceivedPlatformDatagram { + received: usize, + datagram_length: usize, + source_address: SocketAddrV4, + } + + fn send_bytes( + session: &BrokerSession, + handle: ObjectHandle, + data: &[u8], + flags: SendFlags, + ) -> BrokerResult> { + litebox_broker_core::socket::send(session, handle, data.to_vec(), flags) + } + + fn send_datagram( + session: &BrokerSession, + handle: ObjectHandle, + data: &[u8], + flags: SendFlags, + destination: Option, + ) -> BrokerResult> { + litebox_broker_core::socket::send_to(session, handle, data.to_vec(), flags, destination) + } + + fn receive_into( + session: &BrokerSession, + handle: ObjectHandle, + data: &mut [u8], + flags: ReceiveFlags, + peek_offset: u32, + peek_length: u32, + ) -> BrokerResult> { + match litebox_broker_core::socket::receive( + session, + handle, + data.len(), + flags, + peek_offset, + peek_length, + )? { + SocketOutcome::Completed(PlatformStreamReceive::Received(received)) => { + data[..received.len()].copy_from_slice(&received); + Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received( + received + .len() + .try_into() + .map_err(|_| BrokerError::Internal)?, + ))) + } + SocketOutcome::Completed(PlatformStreamReceive::EndOfStream) => { + Ok(SocketOutcome::Completed(ReceiveSocketResponse::EndOfStream)) + } + SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), + } + } + + fn receive_datagram_into( + session: &BrokerSession, + handle: ObjectHandle, + data: &mut [u8], + flags: ReceiveFromFlags, + ) -> BrokerResult> { + match litebox_broker_core::socket::receive_from(session, handle, data.len(), flags)? { + SocketOutcome::Completed(received) => { + data[..received.data.len()].copy_from_slice(&received.data); + Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { + received: received.data.len(), + datagram_length: received.datagram_length, + source_address: received.source_address, + })) + } + SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), } - return Ok(SocketOutcome::Completed(())); } -} -fn handle_socket_event(socket: &mut SocketEntry, events: epoll::EventFlags) -> BrokerResult<()> { - if socket.listening { - return update_snapshot(socket, None, readiness_from_epoll(socket, events)); + #[test] + fn cached_socket_error_precedes_a_new_kernel_error() { + assert_eq!( + shift_pending_error( + Some(SocketError::ConnectionRefused), + Some(SocketError::NetworkUnreachable), + ), + ( + Some(SocketError::ConnectionRefused), + Some(SocketError::NetworkUnreachable), + ) + ); + assert_eq!( + shift_pending_error(None, Some(SocketError::NetworkUnreachable)), + (Some(SocketError::NetworkUnreachable), None) + ); } - if socket.kind == SocketKind::Udp { - return update_snapshot(socket, None, readiness_from_epoll(socket, events)); + + #[test] + fn synchronous_errors_do_not_consume_tcp_connect_status() { + assert!(!can_consume_synchronous_error( + SocketKind::Tcp, + SocketConnectionStatus::Connecting, + )); + assert!(can_consume_synchronous_error( + SocketKind::Tcp, + SocketConnectionStatus::Connected, + )); + assert!(can_consume_synchronous_error( + SocketKind::Udp, + SocketConnectionStatus::Unconnected, + )); } - let republish_readiness = if events.contains(epoll::EventFlags::IN) - && let Some(threshold) = socket.peek_waitall_threshold - { - let threshold_reached = ioctl_fionread(&socket.socket) - .ok() - .and_then(|available| usize::try_from(available).ok()) - .is_none_or(|available| available >= threshold); - if threshold_reached { - socket.peek_waitall_threshold = None; - } - threshold_reached - } else { - false - }; - let status = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .status; - let result = match status { - SocketConnectionStatus::Unconnected => Ok(()), - SocketConnectionStatus::Connecting => complete_connect(socket, events), - SocketConnectionStatus::Connected => { - update_snapshot(socket, None, readiness_from_epoll(socket, events)) - } - SocketConnectionStatus::Failed(SocketError::NotConnected) if socket.read_shutdown => { - update_snapshot(socket, None, ReadinessFlags::WRITE | ReadinessFlags::HANGUP) - } - SocketConnectionStatus::Failed(_) => update_snapshot(socket, None, ReadinessFlags::ERROR), - _ => Err(BrokerError::Internal), - }; - result?; - if republish_readiness { - let readiness = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness; - socket.readiness.republish(readiness)?; + + struct TestReadinessSink { + published: Sender<(ObjectHandle, ReadinessFlags)>, + retired: Sender, } - Ok(()) -} -fn complete_connect(socket: &mut SocketEntry, events: epoll::EventFlags) -> BrokerResult<()> { - let status = match sockopt::socket_error(&socket.socket) { - Ok(Ok(())) => match getpeername(&socket.socket) { - Ok(Some(_)) => { - let local_address = local_socket_address(&socket.socket)?; - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(local_address); - SocketConnectionStatus::Connected - } - Ok(None) | Err(Errno::NOTCONN) => SocketConnectionStatus::Connecting, - Err(error) => SocketConnectionStatus::Failed(socket_error_from_errno(error)), - }, - Ok(Err(error)) | Err(error) => { - SocketConnectionStatus::Failed(socket_error_from_errno(error)) + impl ReadinessSink for TestReadinessSink { + fn max_tracked_objects(&self) -> usize { + 8 } - }; - let readiness = match status { - SocketConnectionStatus::Connected => { - readiness_from_epoll(socket, events) | ReadinessFlags::WRITE + + fn publish(&self, handle: ObjectHandle, readiness: ReadinessFlags) -> BrokerResult<()> { + self.published + .send((handle, readiness)) + .map_err(|_| BrokerError::Internal) } - SocketConnectionStatus::Connecting => ReadinessFlags::default(), - SocketConnectionStatus::Failed(_) => ReadinessFlags::ERROR, - _ => return Err(BrokerError::Internal), - }; - update_snapshot(socket, Some(status), readiness) -} -fn take_socket_error(socket: &SocketEntry) -> BrokerResult> { - match sockopt::socket_error(&socket.socket) { - Ok(Ok(())) => Ok(None), - Ok(Err(error)) | Err(error) => socket_operation_error_from_errno(error).map(Some), + fn republish(&self, handle: ObjectHandle, readiness: ReadinessFlags) -> BrokerResult<()> { + self.publish(handle, readiness) + } + + fn retire(&self, handle: ObjectHandle) { + let _ = self.retired.send(handle); + } } -} -fn local_socket_address(socket: &OwnedFd) -> BrokerResult { - match getsockname(socket) { - Ok(address) => SocketAddrV4::try_from(address).map_err(|_| BrokerError::Internal), - Err(_) => Err(BrokerError::Internal), + #[test] + fn port_mapping_reservation_is_close_on_exec() { + let host_address = unused_tcp_address(); + let retained = create_port_mapping_reservation( + *host_address.ip(), + TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }, + false, + false, + ) + .unwrap(); + + assert!( + rustix::io::fcntl_getfd(&retained) + .unwrap() + .contains(rustix::io::FdFlags::CLOEXEC) + ); } -} -fn update_local_address(socket: &SocketEntry) -> BrokerResult<()> { - let needs_address = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address - .is_none(); - if needs_address { - let address = local_socket_address(&socket.socket)?; - if address.port() != 0 { - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(address); - } + #[test] + fn unavailable_mapped_port_rejects_listen() { + let occupied = TcpListener::bind("127.0.0.1:0").unwrap(); + let host_address = socket_address_v4(occupied.local_addr().unwrap()); + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings( + 1, + 1, + *host_address.ip(), + &[TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(2, 0, 1, 1), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let listener = create_socket(&session, readiness); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80); + assert_eq!( + litebox_broker_core::socket::bind(&session, listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, listener, 1), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + ); } - Ok(()) -} -fn status_socket(socket: &mut SocketEntry) -> BrokerResult { - let query_socket_error = { - let snapshot = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - (snapshot.status == SocketConnectionStatus::Connected || socket.kind == SocketKind::Udp) - && snapshot.readiness.contains(ReadinessFlags::ERROR) - }; - let socket_error = if query_socket_error { - take_socket_error(socket)? - } else { - None - }; - let (response, readiness, republish_error) = { - let mut snapshot = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - let cached_error = snapshot.pending_error.take(); - let (pending_error, next_pending_error) = shift_pending_error(cached_error, socket_error); - snapshot.pending_error = next_pending_error; - if pending_error.is_some() && next_pending_error.is_none() { - snapshot.readiness = ReadinessFlags(snapshot.readiness.0 & !ReadinessFlags::ERROR.0); - } - ( - SocketStatusResponse { - status: snapshot.status, - local_address: snapshot.local_address, - pending_error, - }, - snapshot.readiness, - next_pending_error.is_some(), + #[test] + fn private_backend_endpoints_are_not_guest_destinations() { + let provider = Arc::new(LinuxSocketProvider::new(2, 2).unwrap()); + let all_destinations = DestinationRule::new( + CallerCredential::Unauthenticated, + Ipv4Cidr::new(Ipv4Address([0, 0, 0, 0]), 0).unwrap(), + DestinationPortRange::new(Port(1), Port(u16::MAX)).unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()).with_socket_policy( + SocketPolicy::from_tcp_udp_destination_rules( + &[all_destinations], + &[all_destinations], + ) + .unwrap(), + ), + BrokerCoreLimits::new_with_all_limits(4, 0, 2, 2), + provider.clone(), ) - }; - if republish_error { - socket.readiness.republish(readiness)?; - } else if response.pending_error.is_some() { - socket.readiness.publish(readiness)?; + .unwrap(); + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + + let listener = create_socket(&first_session, readiness.clone()); + let guest_listener_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + assert_eq!( + litebox_broker_core::socket::bind(&first_session, listener, guest_listener_address,), + Ok(SocketOutcome::Completed(guest_listener_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&first_session, listener, 1), + Ok(SocketOutcome::Completed(guest_listener_address)) + ); + let private_tcp_address = provider + .reactor + .host_address(SocketKind::Tcp, guest_listener_address.port()) + .unwrap(); + let tcp_client = create_socket(&second_session, readiness); + let private_tcp_alias = + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, private_tcp_address.port()); + + assert_eq!( + litebox_broker_core::socket::connect(&second_session, tcp_client, private_tcp_alias,), + Ok(SocketOutcome::Completed(SocketConnectionStatus::Failed( + SocketError::ConnectionRefused, + ))) + ); } - Ok(response) -} -fn shift_pending_error( - cached_error: Option, - socket_error: Option, -) -> (Option, Option) { - match cached_error { - Some(error) => (Some(error), socket_error), - None => (socket_error, None), + #[test] + fn pending_guest_tcp_connection_uses_the_complete_host_tuple() { + let mut reactor = test_reactor(MAX_TRACKED_GUEST_CONNECTIONS, &[]); + let session_id = SessionId(7); + let shared_host_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000); + for (guest_port, host_peer_port) in [(1000, 5000), (1001, 5001)] { + insert_test_pending_guest_connection( + &mut reactor, + session_id, + guest_port, + u64::from(guest_port), + shared_host_address, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, host_peer_port), + u64::from(guest_port), + None, + ); + } + + assert_eq!( + reactor + .tcp + .take_pending_guest_connection( + shared_host_address, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5001), + ) + .unwrap() + .guest_address, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1001) + ); + assert_eq!(reactor.tcp.pending_guest_connections.len(), 1); + assert!( + reactor + .tcp + .take_pending_guest_connection( + shared_host_address, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5001), + ) + .is_none() + ); + assert_eq!( + reactor + .tcp + .take_pending_guest_connection( + shared_host_address, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000), + ) + .unwrap() + .guest_address, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1000) + ); + assert!(reactor.tcp.pending_guest_connections.is_empty()); } -} -fn readiness_from_epoll(socket: &SocketEntry, events: epoll::EventFlags) -> ReadinessFlags { - let mut readiness = ReadinessFlags::default(); - if events.contains(epoll::EventFlags::IN) { - readiness = readiness | ReadinessFlags::READ; + #[test] + fn connector_binding_is_not_routable_as_a_guest_listener() { + let mut reactor = test_reactor(1, &[]); + let guest_port = 1000; + let socket_id = 1; + let host_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000); + reactor + .tcp + .insert_binding( + guest_port, + GuestPortBinding { + socket_id, + guest_address: SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, guest_port), + host_address: Some(host_address), + host_peer_address: None, + host_peer_mapping_index: None, + host_mapped: false, + listening: false, + }, + ) + .unwrap(); + + assert_eq!( + reactor.resolve_guest_destination( + SocketKind::Tcp, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, guest_port), + ), + Ok(SocketOutcome::Failed(SocketError::ConnectionRefused)) + ); + + reactor.tcp.mark_listening(guest_port, socket_id).unwrap(); + assert_eq!( + reactor.resolve_guest_destination( + SocketKind::Tcp, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, guest_port), + ), + Ok(SocketOutcome::Completed((host_address, Some(socket_id)))) + ); } - if events.contains(epoll::EventFlags::OUT) && !socket.write_shutdown { - readiness = readiness | ReadinessFlags::WRITE; + + #[test] + fn pending_guest_connection_capacity_counts_all_stale_mappings() { + let mut reactor = test_reactor(1, &[]); + let session_id = SessionId(7); + let foreign_session_id = SessionId(8); + reactor + .sessions + .insert(session_id, SessionSocketState::default()); + + reactor + .reserve_pending_guest_connection(session_id) + .unwrap(); + reactor.tcp.stale_mapped_connections.insert( + ( + 0, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000), + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000), + ), + StaleTcpConnection { + session_id, + deadline: None, + retained_connector: None, + }, + ); + assert_eq!( + reactor.reserve_pending_guest_connection(session_id), + Err(BrokerError::ResourceExhausted) + ); + reactor.tcp.stale_mapped_connections.clear(); + reactor + .tcp + .stale_mapped_connections + .try_reserve(MAX_TRACKED_GUEST_CONNECTIONS) + .unwrap(); + for mapping_index in 0..MAX_TRACKED_GUEST_CONNECTIONS { + reactor.tcp.stale_mapped_connections.insert( + ( + mapping_index, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000), + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000), + ), + StaleTcpConnection { + session_id: foreign_session_id, + deadline: None, + retained_connector: None, + }, + ); + } + assert_eq!( + reactor.reserve_pending_guest_connection(session_id), + Err(BrokerError::ResourceExhausted) + ); } - if socket.kind == SocketKind::Tcp - && !socket.read_shutdown - && events.intersects(epoll::EventFlags::RDHUP | epoll::EventFlags::HUP) - { - readiness = readiness | ReadinessFlags::READ | ReadinessFlags::HANGUP; + + #[test] + fn stale_guest_connection_ownership_is_aggregated_across_mappings() { + let session_id = SessionId(7); + let foreign_session_id = SessionId(8); + let stale = [ + StaleTcpConnection { + session_id, + deadline: None, + retained_connector: None, + }, + StaleTcpConnection { + session_id, + deadline: None, + retained_connector: None, + }, + StaleTcpConnection { + session_id: foreign_session_id, + deadline: None, + retained_connector: None, + }, + ]; + + assert_eq!(count_session_stale_connections(stale.iter(), session_id), 2); } - if events.contains(epoll::EventFlags::ERR) { - readiness = readiness | ReadinessFlags::ERROR; + #[test] + fn retiring_listener_removes_its_pending_guest_connections() { + let mut reactor = test_reactor(MAX_TRACKED_GUEST_CONNECTIONS, &[]); + let session_id = SessionId(7); + let first_listener = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000); + let second_listener = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5001); + for (guest_port, listener) in [(1000, first_listener), (1001, second_listener)] { + insert_test_pending_guest_connection( + &mut reactor, + session_id, + guest_port, + u64::from(guest_port), + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000 + guest_port), + listener, + u64::from(guest_port), + None, + ); + } + + reactor.remove_pending_guest_connections_for_listener(first_listener); + + assert_eq!(reactor.tcp.pending_guest_connections.len(), 1); + assert_eq!( + reactor + .sessions + .get(&session_id) + .unwrap() + .pending_guest_connections, + 1 + ); + assert_eq!( + reactor + .tcp + .take_pending_guest_connection( + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 41001), + second_listener, + ) + .unwrap() + .guest_address, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1001) + ); + assert!(reactor.tcp.pending_guest_connections.is_empty()); } - let previous = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness; - ReadinessFlags(readiness.0 | previous.0) -} -const fn socket_kind(request: CreateSocketRequest) -> Option { - match ( - request.address_family, - request.socket_type, - request.protocol, - ) { - (AddressFamily::Ipv4, SocketType::Stream, IpProtocol::Tcp) => Some(SocketKind::Tcp), - (AddressFamily::Ipv4, SocketType::Datagram, IpProtocol::Udp) => Some(SocketKind::Udp), - _ => None, + #[test] + fn moving_live_pending_guest_connection_keeps_nonexpiring_session_identity() { + let mapping = TcpPortMapping { + guest_port: 5000, + broker_port: 5000, + }; + let mut reactor = test_reactor(MAX_TRACKED_GUEST_CONNECTIONS, &[mapping]); + let session_id = SessionId(7); + let guest_port = 1000; + let socket_id = 1; + let host_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000); + let listener_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000); + insert_test_pending_guest_connection( + &mut reactor, + session_id, + guest_port, + socket_id, + host_address, + listener_address, + 2, + Some(0), + ); + reactor.retire_pending_guest_connection_for_connector( + session_id, + guest_port, + socket_id, + false, + None, + Some(eventfd(0, EventfdFlags::CLOEXEC).unwrap()), + ); + + reactor.move_pending_guest_connections_for_listener(listener_address, 0); + + let stale = reactor + .tcp + .stale_mapped_connections + .get(&(0, host_address, listener_address)) + .unwrap(); + assert_eq!(stale.session_id, session_id); + assert_eq!(stale.deadline, None); + assert!(stale.retained_connector.is_some()); + let session = reactor.sessions.get(&session_id).unwrap(); + assert_eq!(session.pending_guest_connections, 0); + assert_eq!(session.retained_connectors, 1); } -} -fn idle_epoll_events() -> epoll::EventFlags { - epoll::EventFlags::RDHUP | epoll::EventFlags::ET -} + #[test] + fn retiring_failed_connector_removes_its_pending_guest_connection() { + let mut reactor = test_reactor(MAX_TRACKED_GUEST_CONNECTIONS, &[]); + let session_id = SessionId(7); + let guest_port = 1000; + let socket_id = 1; + let host_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000); + let listener_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000); + insert_test_pending_guest_connection( + &mut reactor, + session_id, + guest_port, + socket_id, + host_address, + listener_address, + 2, + None, + ); -fn active_epoll_events() -> epoll::EventFlags { - // Cached readiness turns these edge-triggered kernel events into the - // level-triggered snapshots consumed by the broker protocol. - epoll::EventFlags::IN - | epoll::EventFlags::OUT - | epoll::EventFlags::RDHUP - | epoll::EventFlags::ET -} + reactor.remove_pending_guest_connection_for_connector( + session_id, + SocketKind::Tcp, + guest_port, + socket_id, + ); -fn update_snapshot( - socket: &SocketEntry, - status: Option, - readiness: ReadinessFlags, -) -> BrokerResult<()> { - let readiness_changed = { - let mut snapshot = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - if let Some(status) = status { - snapshot.status = status; - } - let changed = snapshot.readiness != readiness; - snapshot.readiness = readiness; - changed - }; - if readiness_changed { - socket.readiness.publish(readiness)?; + assert!(reactor.tcp.pending_guest_connections.is_empty()); + assert_eq!( + reactor + .tcp + .bindings + .get(&guest_port) + .unwrap() + .host_peer_address, + None + ); + assert_eq!( + reactor + .sessions + .get(&session_id) + .unwrap() + .pending_guest_connections, + 0 + ); } - Ok(()) -} -fn add_readiness(socket: &SocketEntry, readiness: ReadinessFlags) -> BrokerResult<()> { - let current = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness; - update_snapshot(socket, None, ReadinessFlags(current.0 | readiness.0)) -} + #[test] + fn aborting_connector_marks_its_pending_guest_connection_for_discard() { + let mut reactor = test_reactor(MAX_TRACKED_GUEST_CONNECTIONS, &[]); + let session_id = SessionId(7); + let guest_port = 1000; + let socket_id = 1; + let host_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000); + let listener_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000); + insert_test_pending_guest_connection( + &mut reactor, + session_id, + guest_port, + socket_id, + host_address, + listener_address, + 2, + None, + ); -fn clear_readiness(socket: &SocketEntry, readiness: ReadinessFlags) -> BrokerResult<()> { - let current = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness; - update_snapshot(socket, None, ReadinessFlags(current.0 & !readiness.0)) -} + let _ = reactor.retire_pending_guest_connection_for_connector( + session_id, + guest_port, + socket_id, + true, + None, + Some(eventfd(0, EventfdFlags::CLOEXEC).unwrap()), + ); -fn consume_synchronous_error(socket: &SocketEntry) -> BrokerResult<()> { - let query_socket_error = { - let snapshot = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - if !can_consume_synchronous_error(socket.kind, snapshot.status) { - return Ok(()); - } - snapshot.pending_error.is_none() - }; - let socket_error = if query_socket_error { - take_socket_error(socket)? - } else { - None - }; - let (readiness, changed) = { - let mut snapshot = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned"); - if snapshot.pending_error.is_none() { - snapshot.pending_error = socket_error; - } - let readiness = if snapshot.pending_error.is_some() { - snapshot.readiness | ReadinessFlags::ERROR - } else { - ReadinessFlags(snapshot.readiness.0 & !ReadinessFlags::ERROR.0) - }; - let changed = readiness != snapshot.readiness; - snapshot.readiness = readiness; - (readiness, changed) - }; - if changed { - socket.readiness.publish(readiness)?; + let pending = reactor + .tcp + .take_pending_guest_connection(host_address, listener_address) + .unwrap(); + assert!(pending.discard_on_accept); + assert!(pending.retained_connector.is_some()); } - Ok(()) -} -const fn can_consume_synchronous_error(kind: SocketKind, status: SocketConnectionStatus) -> bool { - matches!(kind, SocketKind::Udp) || matches!(status, SocketConnectionStatus::Connected) -} + #[test] + fn closing_session_keeps_established_discard_marker_until_accept() { + let mut reactor = test_reactor(MAX_TRACKED_GUEST_CONNECTIONS, &[]); + let session_id = SessionId(7); + let guest_port = 1000; + let socket_id = 1; + let host_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000); + let listener_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5000); + insert_test_pending_guest_connection( + &mut reactor, + session_id, + guest_port, + socket_id, + host_address, + listener_address, + 2, + None, + ); + reactor.retire_pending_guest_connection_for_connector( + session_id, + guest_port, + socket_id, + false, + None, + Some(eventfd(0, EventfdFlags::CLOEXEC).unwrap()), + ); -const fn socket_error_from_errno(error: Errno) -> SocketError { - match error { - Errno::CONNREFUSED => SocketError::ConnectionRefused, - Errno::CONNRESET | Errno::PIPE => SocketError::ConnectionReset, - Errno::CONNABORTED => SocketError::ConnectionAborted, - Errno::NETUNREACH => SocketError::NetworkUnreachable, - Errno::HOSTUNREACH => SocketError::HostUnreachable, - Errno::TIMEDOUT => SocketError::TimedOut, - Errno::ADDRINUSE => SocketError::AddressInUse, - Errno::ADDRNOTAVAIL => SocketError::AddressNotAvailable, - Errno::NOTCONN => SocketError::NotConnected, - Errno::INVAL => SocketError::InvalidArgument, - _ => SocketError::Other, + reactor.retire_session_connectors(session_id); + reactor.expire_deadlined_state( + Instant::now() + PENDING_CONNECT_DISCARD_LIFETIME + Duration::from_secs(1), + ); + + let pending = reactor + .tcp + .pending_guest_connections + .get(&(host_address, listener_address)) + .unwrap(); + assert!(pending.discard_on_accept); + assert_eq!(pending.discard_deadline, None); + assert!(pending.retained_connector.is_none()); + assert_eq!(reactor.retained_connectors, 0); } -} -const fn socket_operation_error_from_errno(error: Errno) -> BrokerResult { - match broker_resource_error_from_errno(error) { - Some(error) => Err(error), - None => Ok(socket_error_from_errno(error)), + #[test] + fn guest_tcp_ports_are_broker_wide_and_do_not_bind_private_host_ports() { + let occupied_host_listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let occupied_host_address = socket_address_v4(occupied_host_listener.local_addr().unwrap()); + let provider = Arc::new(LinuxSocketProvider::new(4, 4).unwrap()); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(8, 0, 4, 4), + provider, + ) + .unwrap(); + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let first = create_socket(&first_session, readiness.clone()); + let second = create_socket(&second_session, readiness.clone()); + let occupied_host_port = create_socket(&first_session, readiness); + let guest_port_80 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + + assert_eq!( + litebox_broker_core::socket::bind(&first_session, first, guest_port_80), + Ok(SocketOutcome::Completed(guest_port_80)) + ); + assert_eq!( + litebox_broker_core::socket::bind(&second_session, second, guest_port_80), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + ); + assert_eq!( + litebox_broker_core::socket::bind( + &first_session, + occupied_host_port, + occupied_host_address, + ), + Ok(SocketOutcome::Completed(occupied_host_address)) + ); } -} -const fn broker_error_from_errno(error: Errno) -> BrokerError { - match broker_resource_error_from_errno(error) { - Some(error) => error, - None => BrokerError::Internal, + #[test] + fn guest_tcp_loopback_routes_across_sessions() { + let provider = Arc::new(LinuxSocketProvider::new(3, 3).unwrap()); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 3, 3), + provider.clone(), + ) + .unwrap(); + let listener_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let client_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let listener = create_socket(&listener_session, readiness.clone()); + let guest_listener_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + assert_eq!( + litebox_broker_core::socket::bind(&listener_session, listener, guest_listener_address,), + Ok(SocketOutcome::Completed(guest_listener_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&listener_session, listener, 2), + Ok(SocketOutcome::Completed(guest_listener_address)) + ); + + let client = create_socket(&client_session, readiness.clone()); + let connect = + litebox_broker_core::socket::connect(&client_session, client, guest_listener_address) + .unwrap(); + assert!(matches!( + connect, + SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + ) + )); + wait_until_connected(&client_session, client, &publications); + let client_address = litebox_broker_core::socket::status(&client_session, client) + .unwrap() + .local_address + .expect("connected client must have a guest-local address"); + client_session.close_object_reference(client).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), client); + assert_eq!(provider.reactor.retained_connector_count(), 1); + + let replacement = create_socket(&client_session, readiness.clone()); + let replacement_address = + SocketAddrV4::new(Ipv4Addr::LOCALHOST, FIRST_GUEST_EPHEMERAL_PORT + 1); + assert_eq!( + litebox_broker_core::socket::bind(&client_session, replacement, replacement_address,), + Ok(SocketOutcome::Completed(replacement_address)) + ); + let connect = litebox_broker_core::socket::connect( + &client_session, + replacement, + guest_listener_address, + ) + .unwrap(); + assert!(matches!( + connect, + SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + ) + )); + wait_until_connected(&client_session, replacement, &publications); + if !listener_session + .check_readiness(listener) + .unwrap() + .contains(ReadinessFlags::READ) + { + wait_for_readiness(&publications, listener, ReadinessFlags::READ); + } + let accepted = match litebox_broker_core::socket::accept( + &listener_session, + listener, + readiness.clone(), + ) + .unwrap() + { + SocketOutcome::Completed(accepted) => accepted, + SocketOutcome::Failed(error) => panic!("accept failed: {error:?}"), + }; + assert_eq!(accepted.local_address, guest_listener_address); + assert_eq!(accepted.remote_address, client_address); + assert_eq!(provider.reactor.retained_connector_count(), 0); + listener_session + .close_object_reference(accepted.handle) + .unwrap(); + assert_eq!( + retirements.recv_timeout(TEST_TIMEOUT).unwrap(), + accepted.handle + ); + + let accepted = + match litebox_broker_core::socket::accept(&listener_session, listener, readiness) + .unwrap() + { + SocketOutcome::Completed(accepted) => accepted, + SocketOutcome::Failed(error) => panic!("replacement accept failed: {error:?}"), + }; + assert_eq!(accepted.local_address, guest_listener_address); + assert_eq!(accepted.remote_address, replacement_address); + assert_eq!( + litebox_broker_core::socket::status(&client_session, replacement) + .unwrap() + .local_address, + Some(replacement_address) + ); + assert_eq!( + send_bytes(&client_session, replacement, b"x", SendFlags::NONE), + Ok(SocketOutcome::Completed(1)) + ); + wait_for_readiness(&publications, accepted.handle, ReadinessFlags::READ); + let mut byte = [0]; + assert_eq!( + receive_into( + &listener_session, + accepted.handle, + &mut byte, + ReceiveFlags::NONE, + 0, + 0, + ), + Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received(1))) + ); + assert_eq!(byte, *b"x"); } -} -const fn broker_resource_error_from_errno(error: Errno) -> Option { - match error { - Errno::NOMEM => Some(BrokerError::OutOfMemory), - Errno::MFILE | Errno::NFILE | Errno::NOBUFS | Errno::NOSPC => { - Some(BrokerError::ResourceExhausted) + #[test] + fn closing_connector_session_preserves_cross_session_discard_marker() { + let provider = Arc::new(LinuxSocketProvider::new(2, 2).unwrap()); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 2, 2), + provider.clone(), + ) + .unwrap(); + let listener_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let connector_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 8080); + let listener = create_socket(&listener_session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&listener_session, listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&listener_session, listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + let connector = create_socket(&connector_session, readiness.clone()); + assert!(matches!( + litebox_broker_core::socket::connect(&connector_session, connector, guest_address), + Ok(SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + )) + )); + wait_until_connected(&connector_session, connector, &publications); + connector_session.close_object_reference(connector).unwrap(); + assert_eq!(provider.reactor.pending_guest_connection_count(), 1); + assert_eq!(provider.reactor.retained_connector_count(), 1); + + drop(connector_session); + assert_eq!(provider.reactor.pending_guest_connection_count(), 1); + assert_eq!(provider.reactor.retained_connector_count(), 0); + if !listener_session + .check_readiness(listener) + .unwrap() + .contains(ReadinessFlags::READ) + { + wait_for_readiness(&publications, listener, ReadinessFlags::READ); } - _ => None, + assert!(matches!( + litebox_broker_core::socket::accept(&listener_session, listener, readiness), + Err(BrokerError::WouldBlock) + )); + assert_eq!(provider.reactor.pending_guest_connection_count(), 0); + listener_session.close_object_reference(listener).unwrap(); } -} -#[cfg(test)] -mod tests { - use std::io::{Read as _, Write as _}; - use std::net::Ipv4Addr; - use std::net::{Shutdown, TcpListener, TcpStream, UdpSocket}; - use std::sync::mpsc::{Receiver, Sender, channel}; - use std::time::{Duration, Instant}; + #[test] + fn mapped_tcp_listener_can_be_replaced_while_an_accepted_socket_remains() { + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings( + 2, + 2, + *host_address.ip(), + &[TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 2, 2), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80); + let listener = create_socket(&session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&session, listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); - use super::*; - use litebox_broker_core::readiness::ReadinessSink; - use litebox_broker_core::{ - BrokerCore, BrokerCoreLimits, BrokerSession, CallerCredential, DestinationPortRange, - DestinationRule, Ipv4Cidr, ObjectRights, PolicyEngine, SocketPolicy, - }; - use litebox_broker_protocol::ObjectHandle; - use litebox_broker_protocol::socket::{Ipv4Address, Port, ReceiveSocketResponse}; + let first_client = TcpStream::connect(host_address).unwrap(); + wait_for_readiness(&publications, listener, ReadinessFlags::READ); + let accepted = + match litebox_broker_core::socket::accept(&session, listener, readiness.clone()) + .unwrap() + { + SocketOutcome::Completed(accepted) => accepted, + SocketOutcome::Failed(error) => panic!("accept failed: {error:?}"), + }; + assert_eq!( + litebox_broker_core::socket::shutdown(&session, listener, ShutdownMode::StopListening,), + Ok(SocketOutcome::Completed(())) + ); + assert_eq!( + create_port_mapping_reservation( + *host_address.ip(), + TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }, + true, + true, + ) + .unwrap_err(), + Errno::ADDRINUSE + ); + session.close_object_reference(listener).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), listener); - const TEST_TIMEOUT: Duration = Duration::from_secs(5); + let replacement = create_socket(&session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&session, replacement, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, replacement, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + let second_client = TcpStream::connect(host_address).unwrap(); - #[derive(Clone, Copy, Debug, PartialEq, Eq)] - struct ReceivedPlatformDatagram { - received: usize, - datagram_length: usize, - source_address: SocketAddrV4, + drop((first_client, second_client)); + session.close_object_reference(accepted.handle).unwrap(); } - fn send_bytes( - session: &BrokerSession, - handle: ObjectHandle, - data: &[u8], - flags: SendFlags, - ) -> BrokerResult> { - litebox_broker_core::socket::send(session, handle, data.to_vec(), flags) - } + #[test] + fn unacknowledged_mapped_listen_retires_platform_ownership() { + let host_address = unused_tcp_address(); + let mapping = TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }; + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings(2, 1, *host_address.ip(), &[mapping]) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 2, 1), + provider.clone(), + ) + .unwrap(); + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, mapping.guest_port); + let first = create_socket(&first_session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&first_session, first, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + let first_id = provider + .reactor + .next_socket_id + .load(Ordering::Relaxed) + .checked_sub(1) + .unwrap(); + let (response, receive) = sync_channel(1); + drop(receive); + provider + .reactor + .commands + .send(ReactorCommand::Listen { + id: first_id, + backlog: 1, + mapping: Some(mapping), + response, + }) + .unwrap(); + provider.reactor.signal().unwrap(); + assert_eq!( + provider + .reactor + .host_address(SocketKind::Tcp, mapping.guest_port), + None + ); + first_session.close_object_reference(first).unwrap(); - fn send_datagram( - session: &BrokerSession, - handle: ObjectHandle, - data: &[u8], - flags: SendFlags, - destination: Option, - ) -> BrokerResult> { - litebox_broker_core::socket::send_to(session, handle, data.to_vec(), flags, destination) + let second = create_socket(&second_session, readiness); + assert_eq!( + litebox_broker_core::socket::bind(&second_session, second, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&second_session, second, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); } - fn receive_into( - session: &BrokerSession, - handle: ObjectHandle, - data: &mut [u8], - flags: ReceiveFlags, - peek_offset: u32, - peek_length: u32, - ) -> BrokerResult> { - match litebox_broker_core::socket::receive( - session, - handle, - data.len(), - flags, - peek_offset, - peek_length, - )? { - SocketOutcome::Completed(PlatformStreamReceive::Received(received)) => { - data[..received.len()].copy_from_slice(&received); - Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received( - received - .len() - .try_into() - .map_err(|_| BrokerError::Internal)?, - ))) - } - SocketOutcome::Completed(PlatformStreamReceive::EndOfStream) => { - Ok(SocketOutcome::Completed(ReceiveSocketResponse::EndOfStream)) - } - SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), - } + #[test] + fn stopped_mapped_listener_keeps_global_guest_port_until_close() { + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings( + 3, + 2, + *host_address.ip(), + &[TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(6, 0, 3, 2), + provider, + ) + .unwrap(); + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80); + let first_listener = create_socket(&first_session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&first_session, first_listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&first_session, first_listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::shutdown( + &first_session, + first_listener, + ShutdownMode::StopListening, + ), + Ok(SocketOutcome::Completed(())) + ); + + let second_listener = create_socket(&second_session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&second_session, second_listener, guest_address), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + ); + let first_client = create_socket(&first_session, readiness); + let guest_destination = SocketAddrV4::new(Ipv4Addr::LOCALHOST, guest_address.port()); + assert_eq!( + litebox_broker_core::socket::connect(&first_session, first_client, guest_destination,), + Ok(SocketOutcome::Completed(SocketConnectionStatus::Failed( + SocketError::ConnectionRefused, + ))) + ); + + first_session.close_object_reference(first_client).unwrap(); + first_session + .close_object_reference(first_listener) + .unwrap(); + assert_eq!( + litebox_broker_core::socket::bind(&second_session, second_listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&second_session, second_listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + second_session + .close_object_reference(second_listener) + .unwrap(); } - fn receive_datagram_into( - session: &BrokerSession, - handle: ObjectHandle, - data: &mut [u8], - flags: ReceiveFromFlags, - ) -> BrokerResult> { - match litebox_broker_core::socket::receive_from(session, handle, data.len(), flags)? { - SocketOutcome::Completed(received) => { - data[..received.data.len()].copy_from_slice(&received.data); - Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { - received: received.data.len(), - datagram_length: received.datagram_length, - source_address: received.source_address, - })) - } - SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), - } + #[test] + fn mapped_listener_handoff_drains_bounded_stale_connections() { + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings( + 3, + 2, + *host_address.ip(), + &[TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }], + ) + .unwrap(), + ); + let socket_provider: Arc = provider.clone(); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(6, 0, 3, 2), + socket_provider, + ) + .unwrap(); + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80); + let first_listener = create_socket(&first_session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&first_session, first_listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&first_session, first_listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + let first_client = create_socket(&second_session, readiness.clone()); + let guest_destination = SocketAddrV4::new(Ipv4Addr::LOCALHOST, guest_address.port()); + assert!(matches!( + litebox_broker_core::socket::connect(&second_session, first_client, guest_destination,), + Ok(SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + )) + )); + wait_until_connected(&second_session, first_client, &publications); + assert_eq!(provider.reactor.pending_guest_connection_count(), 1); + assert_eq!( + litebox_broker_core::socket::shutdown( + &first_session, + first_listener, + ShutdownMode::StopListening, + ), + Ok(SocketOutcome::Completed(())) + ); + assert_eq!(provider.reactor.pending_guest_connection_count(), 0); + assert_eq!(provider.reactor.stale_guest_connection_count(), 1); + assert_eq!( + litebox_broker_core::socket::shutdown( + &second_session, + first_client, + ShutdownMode::Abort, + ), + Ok(SocketOutcome::Completed(())) + ); + second_session.close_object_reference(first_client).unwrap(); + assert_eq!(provider.reactor.retained_connector_count(), 1); + first_session + .close_object_reference(first_listener) + .unwrap(); + drop(first_session); + assert_eq!(provider.reactor.retained_connector_count(), 1); + assert_eq!(provider.reactor.stale_guest_connection_count(), 1); + + let second_listener = create_socket(&second_session, readiness); + assert_eq!( + litebox_broker_core::socket::bind(&second_session, second_listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&second_session, second_listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!(provider.reactor.stale_guest_connection_count(), 0); + assert_eq!(provider.reactor.retained_connector_count(), 0); + + second_session + .close_object_reference(second_listener) + .unwrap(); } #[test] - fn cached_socket_error_precedes_a_new_kernel_error() { + fn stopped_private_listener_cleans_cross_session_pending_guest_connections() { + let provider = Arc::new(LinuxSocketProvider::new(3, 2).unwrap()); + let socket_provider: Arc = provider.clone(); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(6, 0, 3, 2), + socket_provider, + ) + .unwrap(); + let listener_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let client_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 8000); + let listener = create_socket(&listener_session, readiness.clone()); assert_eq!( - shift_pending_error( - Some(SocketError::ConnectionRefused), - Some(SocketError::NetworkUnreachable), + litebox_broker_core::socket::bind(&listener_session, listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&listener_session, listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + let connected_client = create_socket(&client_session, readiness.clone()); + assert!(matches!( + litebox_broker_core::socket::connect(&client_session, connected_client, guest_address,), + Ok(SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + )) + )); + wait_until_connected(&client_session, connected_client, &publications); + assert_eq!(provider.reactor.pending_guest_connection_count(), 1); + assert_eq!( + litebox_broker_core::socket::shutdown( + &listener_session, + listener, + ShutdownMode::StopListening, ), - ( - Some(SocketError::ConnectionRefused), - Some(SocketError::NetworkUnreachable), - ) + Ok(SocketOutcome::Completed(())) ); + assert_eq!(provider.reactor.pending_guest_connection_count(), 0); + + let client = create_socket(&client_session, readiness); assert_eq!( - shift_pending_error(None, Some(SocketError::NetworkUnreachable)), - (Some(SocketError::NetworkUnreachable), None) + litebox_broker_core::socket::connect(&client_session, client, guest_address), + Ok(SocketOutcome::Completed(SocketConnectionStatus::Failed( + SocketError::ConnectionRefused, + ))) ); + assert!( + client_session + .check_readiness(client) + .unwrap() + .contains(ReadinessFlags::ERROR) + ); + assert_eq!(provider.reactor.pending_guest_connection_count(), 0); + + client_session.close_object_reference(client).unwrap(); + client_session + .close_object_reference(connected_client) + .unwrap(); + listener_session.close_object_reference(listener).unwrap(); } #[test] - fn synchronous_errors_do_not_consume_tcp_connect_status() { - assert!(!can_consume_synchronous_error( - SocketKind::Tcp, - SocketConnectionStatus::Connecting, + fn aborting_guest_routed_connector_retains_a_discard_marker() { + let provider = Arc::new(LinuxSocketProvider::new(3, 2).unwrap()); + let socket_provider: Arc = provider.clone(); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(8, 0, 6, 4), + socket_provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 8000); + let listener = create_socket(&session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&session, listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + let client = create_socket(&session, readiness.clone()); + assert!(matches!( + litebox_broker_core::socket::connect(&session, client, guest_address), + Ok(SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + )) )); - assert!(can_consume_synchronous_error( - SocketKind::Tcp, - SocketConnectionStatus::Connected, + wait_until_connected(&session, client, &publications); + assert_eq!(provider.reactor.pending_guest_connection_count(), 1); + + assert_eq!( + litebox_broker_core::socket::shutdown(&session, client, ShutdownMode::Abort), + Ok(SocketOutcome::Completed(())) + ); + session.close_object_reference(client).unwrap(); + + assert_eq!(provider.reactor.pending_guest_connection_count(), 1); + assert_eq!( + litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Stream, + protocol: IpProtocol::Tcp, + }, + readiness.clone(), + ), + Err(BrokerError::ResourceExhausted) + ); + let other_session_socket = create_socket(&second_session, readiness.clone()); + assert!(matches!( + litebox_broker_core::socket::accept(&session, listener, readiness.clone()), + Err(BrokerError::WouldBlock) )); - assert!(can_consume_synchronous_error( - SocketKind::Udp, - SocketConnectionStatus::Unconnected, + assert_eq!(provider.reactor.pending_guest_connection_count(), 0); + let replacement = create_socket(&session, readiness.clone()); + session.close_object_reference(replacement).unwrap(); + second_session + .close_object_reference(other_session_socket) + .unwrap(); + let retiring_client = create_socket(&session, readiness.clone()); + assert!(matches!( + litebox_broker_core::socket::connect(&session, retiring_client, guest_address), + Ok(SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + )) )); + wait_until_connected(&session, retiring_client, &publications); + assert_eq!( + litebox_broker_core::socket::shutdown(&session, retiring_client, ShutdownMode::Abort,), + Ok(SocketOutcome::Completed(())) + ); + session.close_object_reference(retiring_client).unwrap(); + assert_eq!(provider.reactor.pending_guest_connection_count(), 1); + session.close_object_reference(listener).unwrap(); + assert_eq!(provider.reactor.pending_guest_connection_count(), 0); + let after_retirement = create_socket(&second_session, readiness); + second_session + .close_object_reference(after_retirement) + .unwrap(); } - struct TestReadinessSink { - published: Sender<(ObjectHandle, ReadinessFlags)>, - retired: Sender, - } - - impl ReadinessSink for TestReadinessSink { - fn max_tracked_objects(&self) -> usize { - 8 - } + #[test] + fn closing_mapped_listener_discards_more_than_sixty_four_queued_connections() { + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings( + 2, + 2, + *host_address.ip(), + &[TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 2, 2), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80); + let listener = create_socket(&session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&session, listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, listener, 65), + Ok(SocketOutcome::Completed(guest_address)) + ); - fn publish(&self, handle: ObjectHandle, readiness: ReadinessFlags) -> BrokerResult<()> { - self.published - .send((handle, readiness)) - .map_err(|_| BrokerError::Internal) - } + let clients = (0..65) + .map(|_| TcpStream::connect(host_address).unwrap()) + .collect::>(); + session.close_object_reference(listener).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), listener); + assert_eq!( + create_port_mapping_reservation( + *host_address.ip(), + TcpPortMapping { + broker_port: host_address.port(), + guest_port: 80, + }, + true, + true, + ) + .unwrap_err(), + Errno::ADDRINUSE + ); - fn republish(&self, handle: ObjectHandle, readiness: ReadinessFlags) -> BrokerResult<()> { - self.publish(handle, readiness) - } + let replacement = create_socket(&session, readiness); + assert_eq!( + litebox_broker_core::socket::bind(&session, replacement, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, replacement, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + TcpStream::connect(host_address).unwrap(); - fn retire(&self, handle: ObjectHandle) { - let _ = self.retired.send(handle); - } + drop(clients); } #[test] @@ -2193,8 +5895,10 @@ mod tests { let address = listener.local_addr().unwrap(); let (allow_response, response_allowed) = channel(); let (allow_end_of_stream, end_of_stream_allowed) = channel(); + let (peer_address, peer_addresses) = channel(); let server = thread::spawn(move || { - let (mut stream, _) = listener.accept().unwrap(); + let (mut stream, peer) = listener.accept().unwrap(); + peer_address.send(peer).unwrap(); stream.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); stream.set_write_timeout(Some(TEST_TIMEOUT)).unwrap(); let mut request = [0_u8; 4]; @@ -2232,7 +5936,11 @@ mod tests { assert_eq!(error.kind(), std::io::ErrorKind::ConnectionReset); }); - let provider = Arc::new(LinuxSocketProvider::new(8).unwrap()); + let broker_ipv4_address = Ipv4Addr::new(127, 0, 0, 2); + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings(8, 8, broker_ipv4_address, &[]) + .unwrap(), + ); let broker = BrokerCore::new_with_limits( PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) .with_socket_policy(SocketPolicy::Ipv4Loopback), @@ -2257,13 +5965,17 @@ mod tests { ) )); wait_until_connected(&session, handle, &publications); + assert_eq!( + socket_address_v4(peer_addresses.recv_timeout(TEST_TIMEOUT).unwrap()).ip(), + &broker_ipv4_address + ); let status = litebox_broker_core::socket::status(&session, handle).unwrap(); assert_eq!(status.status, SocketConnectionStatus::Connected); let local_address = status .local_address .expect("connected socket must expose its local address"); - assert_eq!(*local_address.ip(), Ipv4Addr::LOCALHOST); - assert_ne!(local_address.port(), 0); + assert_eq!(*local_address.ip(), Ipv4Addr::UNSPECIFIED); + assert_eq!(local_address.port(), FIRST_GUEST_EPHEMERAL_PORT); assert_eq!(status.pending_error, None); assert_eq!( litebox_broker_core::socket::get_tcp_option(&session, handle, TcpOptionName::NoDelay,), @@ -2548,8 +6260,20 @@ mod tests { } #[test] - fn reactor_assigns_a_port_to_an_unbound_tcp_listener() { - let provider = Arc::new(LinuxSocketProvider::new(2).unwrap()); + fn unbound_listener_uses_identity_mapping_and_configured_override() { + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings( + 2, + 2, + *host_address.ip(), + &[TcpPortMapping { + broker_port: host_address.port(), + guest_port: FIRST_GUEST_EPHEMERAL_PORT, + }], + ) + .unwrap(), + ); let broker = BrokerCore::new_with_limits( PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) .with_socket_policy(SocketPolicy::Ipv4Loopback), @@ -2560,7 +6284,7 @@ mod tests { let session = broker .create_session(CallerCredential::Unauthenticated) .unwrap(); - let (published, publications) = channel(); + let (published, _publications) = channel(); let (retired, _retirements) = channel(); let readiness = Arc::new(TestReadinessSink { published, retired }); let listener = create_socket(&session, readiness.clone()); @@ -2571,33 +6295,88 @@ mod tests { SocketOutcome::Completed(address) => address, SocketOutcome::Failed(error) => panic!("listen failed: {error:?}"), }; - assert_ne!(local_address.port(), 0); - - let connect_address = SocketAddrV4::new( - if local_address.ip().is_unspecified() { - Ipv4Addr::LOCALHOST - } else { - *local_address.ip() - }, - local_address.port(), + assert_eq!(local_address.port(), FIRST_GUEST_EPHEMERAL_PORT + 1); + assert!(local_address.ip().is_unspecified()); + TcpStream::connect(SocketAddrV4::new(*host_address.ip(), local_address.port())).unwrap(); + drop(TcpListener::bind(host_address).unwrap()); + let mapped_listener = create_socket(&session, readiness); + let mapped_guest_address = + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, FIRST_GUEST_EPHEMERAL_PORT); + assert_eq!( + litebox_broker_core::socket::set_tcp_option( + &session, + mapped_listener, + TcpOptionValue::NoDelay(true), + ), + Ok(()) + ); + assert_eq!( + litebox_broker_core::socket::set_tcp_option( + &session, + mapped_listener, + TcpOptionValue::KeepAlive(true), + ), + Ok(()) + ); + assert_eq!( + litebox_broker_core::socket::bind(&session, mapped_listener, mapped_guest_address,), + Ok(SocketOutcome::Completed(mapped_guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::get_tcp_option( + &session, + mapped_listener, + TcpOptionName::NoDelay, + ), + Ok(TcpOptionValue::NoDelay(true)) ); - let client = TcpStream::connect(connect_address).unwrap(); - wait_for_readiness(&publications, listener, ReadinessFlags::READ); - let accepted = - match litebox_broker_core::socket::accept(&session, listener, readiness).unwrap() { - SocketOutcome::Completed(accepted) => accepted, - SocketOutcome::Failed(error) => panic!("accept failed: {error:?}"), - }; - assert_eq!(accepted.local_address, connect_address); assert_eq!( - accepted.remote_address, - socket_address_v4(client.local_addr().unwrap()) + litebox_broker_core::socket::get_tcp_option( + &session, + mapped_listener, + TcpOptionName::KeepAlive, + ), + Ok(TcpOptionValue::KeepAlive(true)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, mapped_listener, 1), + Ok(SocketOutcome::Completed(mapped_guest_address)) + ); + assert_eq!( + TcpListener::bind(host_address).unwrap_err().kind(), + ErrorKind::AddrInUse + ); + assert_eq!( + create_port_mapping_reservation( + *host_address.ip(), + TcpPortMapping { + broker_port: host_address.port(), + guest_port: FIRST_GUEST_EPHEMERAL_PORT, + }, + true, + true, + ) + .unwrap_err(), + Errno::ADDRINUSE ); + TcpStream::connect(host_address).unwrap(); } #[test] - fn reactor_drives_a_loopback_tcp_listener() { - let provider = Arc::new(LinuxSocketProvider::new(4).unwrap()); + fn reactor_drives_an_external_tcp_listener() { + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_tcp_port_mappings( + 4, + 4, + *host_address.ip(), + &[TcpPortMapping { + broker_port: host_address.port(), + guest_port: FIRST_GUEST_EPHEMERAL_PORT, + }], + ) + .unwrap(), + ); let broker = BrokerCore::new_with_limits( PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) .with_socket_policy(SocketPolicy::Ipv4Loopback), @@ -2612,7 +6391,8 @@ mod tests { let (retired, retirements) = channel(); let readiness = Arc::new(TestReadinessSink { published, retired }); let listener = create_socket(&session, readiness.clone()); - let requested_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0); + let requested_address = + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, FIRST_GUEST_EPHEMERAL_PORT); let local_address = match litebox_broker_core::socket::bind(&session, listener, requested_address).unwrap() { @@ -2620,7 +6400,7 @@ mod tests { SocketOutcome::Failed(error) => panic!("bind failed: {error:?}"), }; assert_eq!(local_address.ip(), requested_address.ip()); - assert_ne!(local_address.port(), 0); + assert_eq!(local_address, requested_address); assert_eq!( litebox_broker_core::socket::listen(&session, listener, 8), Ok(SocketOutcome::Completed(local_address)) @@ -2642,8 +6422,8 @@ mod tests { .contains(ReadinessFlags::READ) ); - let mut first_client = TcpStream::connect(local_address).unwrap(); - let second_client = TcpStream::connect(local_address).unwrap(); + let mut first_client = TcpStream::connect(host_address).unwrap(); + let second_client = TcpStream::connect(host_address).unwrap(); first_client.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); first_client.set_write_timeout(Some(TEST_TIMEOUT)).unwrap(); wait_for_readiness(&publications, listener, ReadinessFlags::READ); @@ -2760,7 +6540,7 @@ mod tests { let server = UdpSocket::bind("127.0.0.1:0").unwrap(); server.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); let server_address = socket_address_v4(server.local_addr().unwrap()); - let provider = Arc::new(LinuxSocketProvider::new(2).unwrap()); + let provider = Arc::new(LinuxSocketProvider::new(2, 2).unwrap()); let socket_policy = SocketPolicy::from_udp_destination_rules(&[ DestinationRule::new( CallerCredential::Unauthenticated, @@ -3039,6 +6819,13 @@ mod tests { assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), handle); } + fn unused_tcp_address() -> SocketAddrV4 { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = socket_address_v4(listener.local_addr().unwrap()); + drop(listener); + address + } + fn socket_address_v4(address: std::net::SocketAddr) -> SocketAddrV4 { let std::net::SocketAddr::V4(address) = address else { panic!("loopback TCP test unexpectedly used IPv6"); diff --git a/litebox_broker_protocol/src/socket.rs b/litebox_broker_protocol/src/socket.rs index 36867db2f..1188e08ae 100644 --- a/litebox_broker_protocol/src/socket.rs +++ b/litebox_broker_protocol/src/socket.rs @@ -305,7 +305,7 @@ impl SocketError { pub enum SocketOutcome { /// The operation completed successfully. Completed(T), - /// The host network stack reported an ordinary socket failure. + /// The broker socket backend reported an ordinary socket failure. Failed(SocketError), } @@ -355,7 +355,7 @@ pub struct BindSocketRequest { /// Response to a socket bind request. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub struct BindSocketResponse { - /// Local address assigned by the host network stack. + /// Guest-visible local address assigned to the socket. pub local_address: SocketAddrV4, } @@ -371,7 +371,7 @@ pub struct ListenSocketRequest { /// Response to a socket listen request. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub struct ListenSocketResponse { - /// Local address assigned by the host network stack. + /// Local address reserved in the guest's broker-managed namespace. pub local_address: SocketAddrV4, } @@ -553,7 +553,7 @@ pub struct SocketStatusRequest { pub struct SocketStatusResponse { /// Current connection state. pub status: SocketConnectionStatus, - /// Local endpoint assigned by the host network stack, if any. + /// Guest-visible local endpoint assigned to the socket, if any. pub local_address: Option, /// Pending asynchronous socket error, consumed by this status query. pub pending_error: Option, diff --git a/litebox_broker_userland/src/main.rs b/litebox_broker_userland/src/main.rs index 61db89ec2..75ff37c4a 100644 --- a/litebox_broker_userland/src/main.rs +++ b/litebox_broker_userland/src/main.rs @@ -19,7 +19,7 @@ use std::time::{Duration, Instant}; use clap::Parser; use litebox_broker_core::{ BrokerCore, BrokerCoreLimits, CallerCredential, DestinationPortRange, DestinationRule, - Ipv4Cidr, ObjectRights, PolicyEngine, SocketPolicy, SocketPolicyError, + Ipv4Cidr, ObjectRights, PolicyEngine, SocketPolicy, SocketPolicyError, socket::TcpPortMapping, }; use litebox_broker_host::{BrokerHostAssociation, ConnectionTermination, setup_connection}; use litebox_broker_platform_linux_userland::LinuxSocketProvider; @@ -85,6 +85,36 @@ impl FromStr for AllowedTcpDestination { } } +/// Command-line description of one broker-to-guest TCP port mapping. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct TcpPortMappingArgument { + broker_port: u16, + guest_port: u16, +} + +impl FromStr for TcpPortMappingArgument { + type Err = String; + + fn from_str(value: &str) -> Result { + let (broker_port, guest_port) = value + .split_once(':') + .ok_or_else(|| "expected BROKER_PORT:GUEST_PORT".to_owned())?; + let broker_port = broker_port + .parse::() + .map_err(|error| format!("invalid broker port: {error}"))?; + let guest_port = guest_port + .parse::() + .map_err(|error| format!("invalid guest port: {error}"))?; + if broker_port == 0 || guest_port == 0 { + return Err("mapped broker and guest ports must be nonzero".to_owned()); + } + Ok(Self { + broker_port, + guest_port, + }) + } +} + #[derive(Parser, Debug)] struct CliArgs { /// Permit outbound TCP connections to a destination CIDR and port range. @@ -93,6 +123,12 @@ struct CliArgs { /// `0.0.0.0/0:1-65535` permits every nonzero IPv4 TCP destination. #[arg(long, value_name = "CIDR:PORT[-PORT]")] allow_tcp_destination: Vec, + /// Host-facing IPv4 address used for TCP listeners and outbound TCP. + #[arg(long, default_value_t = Ipv4Addr::UNSPECIFIED, value_name = "IP")] + broker_ipv4_address: Ipv4Addr, + /// Map one TCP port on the broker IPv4 address to a guest-local TCP port. + #[arg(long, value_name = "BROKER_PORT:GUEST_PORT")] + tcp_port_mapping: Vec, /// Local runner executable to launch. #[arg(long, value_name = "PATH", value_hint = clap::ValueHint::ExecutablePath)] runner: PathBuf, @@ -110,11 +146,17 @@ fn main() -> Result<(), Box> { let control_listener = UnixListener::bind(&control_socket_path)?; control_listener.set_nonblocking(true)?; let limits = BrokerCoreLimits::DEFAULT; + let tcp_port_mappings = configured_tcp_port_mappings(&args.tcp_port_mapping); let broker = BrokerCore::new_with_limits( PolicyEngine::with_host_guaranteed_rights(ObjectRights::all()) .with_socket_policy(configured_socket_policy(&args.allow_tcp_destination)?), limits, - Arc::new(LinuxSocketProvider::new(limits.max_sockets)?), + Arc::new(LinuxSocketProvider::new_with_tcp_port_mappings( + limits.max_sockets, + limits.max_sockets_per_session, + args.broker_ipv4_address, + &tcp_port_mappings, + )?), )?; let mut runner_command = Command::new(&args.runner); @@ -139,6 +181,15 @@ fn main() -> Result<(), Box> { Ok(()) } +fn configured_tcp_port_mappings(tcp: &[TcpPortMappingArgument]) -> Vec { + tcp.iter() + .map(|mapping| TcpPortMapping { + broker_port: mapping.broker_port, + guest_port: mapping.guest_port, + }) + .collect() +} + fn configured_socket_policy( allowed_destinations: &[AllowedTcpDestination], ) -> Result { @@ -563,6 +614,33 @@ mod tests { assert!("203.0.113.0/24:0".parse::().is_err()); } + #[test] + fn tcp_port_mapping_arguments_name_distinct_broker_and_guest_ports() { + let mapping = "8080:80".parse::().unwrap(); + + assert_eq!( + mapping, + TcpPortMappingArgument { + broker_port: 8080, + guest_port: 80, + } + ); + assert!("0:80".parse::().is_err()); + assert!("8080:0".parse::().is_err()); + assert!( + "127.0.0.1:8080:80" + .parse::() + .is_err() + ); + assert_eq!( + configured_tcp_port_mappings(&[mapping]), + vec![TcpPortMapping { + broker_port: 8080, + guest_port: 80, + }] + ); + } + #[test] fn tcp_destination_arguments_replace_the_loopback_default() { assert_eq!( diff --git a/litebox_runner_linux_userland/tests/run.rs b/litebox_runner_linux_userland/tests/run.rs index 16da35fd4..8b14b0593 100644 --- a/litebox_runner_linux_userland/tests/run.rs +++ b/litebox_runner_linux_userland/tests/run.rs @@ -341,6 +341,23 @@ fn spawn_test_broker( control_socket_path: &Path, policy: litebox_broker_core::PolicyEngine, connection_count: usize, +) -> TestBroker { + spawn_test_broker_with_tcp_port_mappings( + control_socket_path, + policy, + connection_count, + std::net::Ipv4Addr::LOCALHOST, + Vec::new(), + ) +} + +#[cfg(all(target_arch = "x86_64", target_os = "linux"))] +fn spawn_test_broker_with_tcp_port_mappings( + control_socket_path: &Path, + policy: litebox_broker_core::PolicyEngine, + connection_count: usize, + broker_ipv4_address: std::net::Ipv4Addr, + tcp_port_mappings: Vec, ) -> TestBroker { let _ = std::fs::remove_file(control_socket_path); @@ -359,8 +376,11 @@ fn spawn_test_broker( policy, limits, std::sync::Arc::new( - litebox_broker_platform_linux_userland::LinuxSocketProvider::new( + litebox_broker_platform_linux_userland::LinuxSocketProvider::new_with_tcp_port_mappings( limits.max_sockets, + limits.max_sockets_per_session, + broker_ipv4_address, + &tcp_port_mappings, ) .expect("failed to create broker test socket provider"), ), @@ -739,23 +759,37 @@ fn test_runner_broker_tcp_server_with_rewriter() { use std::net::{Ipv4Addr, TcpStream}; use std::process::Stdio; + const GUEST_PORT: u16 = 18_080; + let target = common::compile( "./tests/tcp_broker_server.c", "broker_tcp_server_rewriter", false, false, ); + let host_listener = std::net::TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let host_address = match host_listener.local_addr().unwrap() { + std::net::SocketAddr::V4(address) => address, + std::net::SocketAddr::V6(_) => unreachable!("IPv4 bind returned an IPv6 address"), + }; + drop(host_listener); let control_socket_path = unique_test_socket_path("runner-broker-tcp-server-control"); - let broker = spawn_test_broker( + let broker = spawn_test_broker_with_tcp_port_mappings( &control_socket_path, litebox_broker_core::PolicyEngine::with_host_guaranteed_rights( litebox_broker_core::ObjectRights::all(), ) .with_socket_policy(litebox_broker_core::SocketPolicy::Ipv4Loopback), 1, + *host_address.ip(), + vec![litebox_broker_core::socket::TcpPortMapping { + broker_port: host_address.port(), + guest_port: GUEST_PORT, + }], ); let mut child = Runner::new(&target, "broker_tcp_server_rewriter") .broker_socket(&control_socket_path) + .arg(GUEST_PORT.to_string()) .spawn_with_stdio(Stdio::null(), Stdio::piped(), Stdio::inherit()); let stdout = child.stdout.take().unwrap(); let (line_sender, line_receiver) = std::sync::mpsc::channel(); @@ -786,8 +820,9 @@ fn test_runner_broker_tcp_server_with_rewriter() { }; let listen = next_marker("LISTEN "); - let port = listen.split_whitespace().nth(1).unwrap().parse().unwrap(); - let mut first = TcpStream::connect((Ipv4Addr::LOCALHOST, port)).unwrap(); + let port: u16 = listen.split_whitespace().nth(1).unwrap().parse().unwrap(); + assert_eq!(port, GUEST_PORT); + let mut first = TcpStream::connect(host_address).unwrap(); first.set_read_timeout(Some(BROKER_HELPER_TIMEOUT)).unwrap(); first .set_write_timeout(Some(BROKER_HELPER_TIMEOUT)) @@ -813,7 +848,7 @@ fn test_runner_broker_tcp_server_with_rewriter() { child.try_wait().unwrap().is_none(), "blocking accept returned before a client connected" ); - let mut second = TcpStream::connect((Ipv4Addr::LOCALHOST, port)).unwrap(); + let mut second = TcpStream::connect(host_address).unwrap(); second .set_read_timeout(Some(BROKER_HELPER_TIMEOUT)) .unwrap(); diff --git a/litebox_runner_linux_userland/tests/tcp_broker.c b/litebox_runner_linux_userland/tests/tcp_broker.c index 2def56638..64abaa9cc 100644 --- a/litebox_runner_linux_userland/tests/tcp_broker.c +++ b/litebox_runner_linux_userland/tests/tcp_broker.c @@ -44,7 +44,7 @@ int main(int argc, char **argv) { socklen_t connecting_address_length = sizeof(connecting_address); assert(getsockname(fd, (struct sockaddr *)&connecting_address, &connecting_address_length) == 0); - assert(connecting_address.sin_addr.s_addr == htonl(INADDR_LOOPBACK)); + assert(connecting_address.sin_addr.s_addr == htonl(INADDR_ANY)); assert(connecting_address.sin_port != 0); int epoll_fd = epoll_create1(EPOLL_CLOEXEC); assert(epoll_fd >= 0); @@ -63,7 +63,7 @@ int main(int argc, char **argv) { socklen_t address_length = sizeof(local_address); assert(getsockname(fd, (struct sockaddr *)&local_address, &address_length) == 0); assert(local_address.sin_family == AF_INET); - assert(local_address.sin_addr.s_addr == htonl(INADDR_LOOPBACK)); + assert(local_address.sin_addr.s_addr == htonl(INADDR_ANY)); assert(local_address.sin_port != 0); assert(accept(fd, NULL, NULL) == -1); assert(errno == EINVAL); diff --git a/litebox_runner_linux_userland/tests/tcp_broker_server.c b/litebox_runner_linux_userland/tests/tcp_broker_server.c index 2d053e9c2..8d2fb95c8 100644 --- a/litebox_runner_linux_userland/tests/tcp_broker_server.c +++ b/litebox_runner_linux_userland/tests/tcp_broker_server.c @@ -11,6 +11,7 @@ #include #include #include +#include #include #include @@ -58,16 +59,21 @@ static void *blocking_accept_thread(void *argument) { return NULL; } -int main(void) { +int main(int argc, char **argv) { setvbuf(stdout, NULL, _IONBF, 0); + assert(argc == 2); + char *end = NULL; + unsigned long guest_port = strtoul(argv[1], &end, 10); + assert(end != argv[1] && *end == '\0' && guest_port > 0 && + guest_port <= UINT16_MAX); int listener = socket(AF_INET, SOCK_STREAM | SOCK_NONBLOCK | SOCK_CLOEXEC, 0); assert(listener >= 0); struct sockaddr_in local = { .sin_family = AF_INET, - .sin_addr.s_addr = htonl(INADDR_LOOPBACK), - .sin_port = 0, + .sin_addr.s_addr = htonl(INADDR_ANY), + .sin_port = htons((uint16_t)guest_port), }; assert(bind(listener, (const struct sockaddr *)&local, sizeof(local)) == 0); assert(listen(listener, 8) == 0); @@ -75,7 +81,7 @@ int main(void) { socklen_t length = sizeof(local); assert(getsockname(listener, (struct sockaddr *)&local, &length) == 0); assert(length == sizeof(local)); - assert(local.sin_addr.s_addr == htonl(INADDR_LOOPBACK)); + assert(local.sin_addr.s_addr == htonl(INADDR_ANY)); assert(local.sin_port != 0); assert(accept(listener, NULL, NULL) == -1);