From 76c1978fdd00a52042aff07eea2ac6343ae86213 Mon Sep 17 00:00:00 2001 From: Weidong Cui Date: Sat, 3 Oct 2026 18:29:19 -0700 Subject: [PATCH 1/2] Move Unix domain sockets into the broker Unix domain sockets become guest-neutral broker-owned local sockets, so they can be shared across processes and inherited across fork and exec. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: ceea8bb3-229a-4415-a471-bb3b41d4fed0 --- litebox/src/broker/mod.rs | 250 ++- litebox/src/lib.rs | 1 + litebox/src/local_sockets.rs | 480 +++++ litebox/src/process.rs | 8 + litebox_broker_core/src/fs/mod.rs | 2 +- litebox_broker_core/src/fs/service.rs | 27 + litebox_broker_core/src/lib.rs | 31 + litebox_broker_core/src/local_socket.rs | 1359 +++++++++++++ litebox_broker_core/src/local_socket/tests.rs | 859 +++++++++ litebox_broker_core/src/object.rs | 15 +- litebox_broker_core/src/process.rs | 3 + litebox_broker_core/src/socket/tests.rs | 3 + litebox_broker_host/src/lib.rs | 459 ++++- litebox_broker_local/src/lib.rs | 1 + litebox_broker_local/src/local_socket.rs | 391 ++++ litebox_broker_protocol/src/lib.rs | 1 + litebox_broker_protocol/src/local_socket.rs | 532 +++++ litebox_broker_protocol/src/message.rs | 97 +- litebox_broker_protocol/src/readiness.rs | 2 + litebox_broker_protocol/src/socket.rs | 2 +- litebox_broker_protocol/src/wire.rs | 171 +- litebox_broker_protocol/src/wire/fs.rs | 10 +- .../src/wire/local_socket.rs | 412 ++++ litebox_broker_protocol/src/wire/socket.rs | 64 +- litebox_common_linux/src/errno/mod.rs | 45 + litebox_common_linux/src/program_startup.rs | 20 +- .../tests/fork_unix_parent.c | 225 +++ litebox_runner_linux_userland/tests/run.rs | 42 + .../tests/vfork_exec_child.c | 16 + litebox_shim_linux/src/channel.rs | 334 ---- litebox_shim_linux/src/lib.rs | 4 - litebox_shim_linux/src/syscalls/file.rs | 272 ++- litebox_shim_linux/src/syscalls/net.rs | 114 +- .../src/syscalls/test_broker.rs | 2 +- litebox_shim_linux/src/syscalls/unix.rs | 1710 +++-------------- 35 files changed, 5970 insertions(+), 1994 deletions(-) create mode 100644 litebox/src/local_sockets.rs create mode 100644 litebox_broker_core/src/local_socket.rs create mode 100644 litebox_broker_core/src/local_socket/tests.rs create mode 100644 litebox_broker_local/src/local_socket.rs create mode 100644 litebox_broker_protocol/src/local_socket.rs create mode 100644 litebox_broker_protocol/src/wire/local_socket.rs create mode 100644 litebox_runner_linux_userland/tests/fork_unix_parent.c delete mode 100644 litebox_shim_linux/src/channel.rs diff --git a/litebox/src/broker/mod.rs b/litebox/src/broker/mod.rs index 20d58a54c0..acee38c0d9 100644 --- a/litebox/src/broker/mod.rs +++ b/litebox/src/broker/mod.rs @@ -17,6 +17,11 @@ use litebox_broker_protocol::fs::{ FileAccessMode, FileDirectoryEntry, FileError, FileMode, FileOpenFlags, FileSeekWhence, FileStatus, FileStatusFlags, FileUser, MAX_FILE_TRANSFER_SIZE, }; +use litebox_broker_protocol::local_socket::{ + CreateLocalSocketPairResponse, GetLocalSocketOptionsResponse, LOCAL_SOCKET_BUFFER_SIZE, + LocalSocketAddress, LocalSocketError, LocalSocketName, LocalSocketOption, + MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE, MAX_LOCAL_SOCKET_TRANSFER_SIZE, ReceiveLocalSocketResponse, +}; use litebox_broker_protocol::pipe::{CreatePipeResponse, MAX_PIPE_TRANSFER_SIZE}; use litebox_broker_protocol::process::{ CreatedProcess, MAX_CHILD_MEMORY_WRITE_SIZE, MAX_CHILD_OBJECT_DUPLICATES, @@ -30,7 +35,8 @@ use litebox_broker_protocol::socket::{ AcceptSocketResponse, MAX_SOCKET_TRANSFER_SIZE, MAX_UDP_DATAGRAM_SIZE, ReceiveFlags as BrokerReceiveFlags, ReceiveFromFlags as BrokerReceiveFromFlags, ReceiveFromSocketResponse, ReceiveSocketResponse, SendFlags as BrokerSendFlags, ShutdownMode, - SocketConnectionStatus, SocketOutcome, SocketStatusResponse, TcpOptionName, TcpOptionValue, + SocketConnectionStatus, SocketOutcome, SocketStatusResponse, SocketType, TcpOptionName, + TcpOptionValue, }; use litebox_broker_protocol::timer::TimerSpec; use litebox_broker_transport::channel::LocalCallChannel; @@ -357,6 +363,101 @@ pub(crate) trait BrokerControl: Send + Sync { user: FileUser, ) -> core::result::Result, BrokerControlError>; + fn create_local_socket( + &self, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> core::result::Result; + + fn create_local_socket_pair( + &self, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> core::result::Result; + + fn bind_local_socket( + &self, + handle: ObjectHandle, + address: &LocalSocketAddress, + user: FileUser, + mode: FileMode, + ) -> core::result::Result, BrokerControlError>; + + fn listen_local_socket( + &self, + handle: ObjectHandle, + backlog: u32, + ) -> core::result::Result, BrokerControlError>; + + fn connect_local_socket( + &self, + handle: ObjectHandle, + address: &LocalSocketAddress, + user: FileUser, + ) -> core::result::Result, BrokerControlError>; + + fn accept_local_socket( + &self, + handle: ObjectHandle, + flags: FileOpenFlags, + ) -> core::result::Result< + core::result::Result, + BrokerControlError, + >; + + /// Sends `data` to `address`, or to the connected peer if `address` is + /// `None`. + /// + /// Data that does not fit in one transfer fails with + /// [`LocalSocketError::MessageTooLarge`], so stream callers send it in + /// chunks of at most [`LOCAL_SOCKET_BUFFER_SIZE`] bytes. + fn send_local_socket( + &self, + handle: ObjectHandle, + address: Option<&LocalSocketAddress>, + data: &[u8], + user: FileUser, + ) -> core::result::Result, BrokerControlError>; + + /// Receives into at most [`LOCAL_SOCKET_BUFFER_SIZE`] bytes of `data`, as + /// for a non-blocking socket if `nonblocking` is set. + fn receive_local_socket( + &self, + handle: ObjectHandle, + data: &mut [u8], + peek: bool, + nonblocking: bool, + ) -> core::result::Result< + core::result::Result<(ReceiveLocalSocketResponse, LocalSocketName), LocalSocketError>, + BrokerControlError, + >; + + fn shutdown_local_socket( + &self, + handle: ObjectHandle, + mode: ShutdownMode, + ) -> core::result::Result, BrokerControlError>; + + fn local_socket_name( + &self, + handle: ObjectHandle, + peer: bool, + ) -> core::result::Result< + core::result::Result, + BrokerControlError, + >; + + fn set_local_socket_option( + &self, + handle: ObjectHandle, + option: LocalSocketOption, + ) -> core::result::Result<(), BrokerControlError>; + + fn local_socket_options( + &self, + handle: ObjectHandle, + ) -> core::result::Result; + fn close_object(&self, handle: ObjectHandle) -> core::result::Result<(), BrokerControlError>; fn fail_connection(&self); @@ -1066,6 +1167,153 @@ where self.request(|local| local.rmdir_file(lease.sequence(), path, user)) } + fn create_local_socket( + &self, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> core::result::Result { + self.request(|local| local.create_local_socket(socket_type, flags)) + } + + fn create_local_socket_pair( + &self, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> core::result::Result { + self.request(|local| local.create_local_socket_pair(socket_type, flags)) + } + + fn bind_local_socket( + &self, + handle: ObjectHandle, + address: &LocalSocketAddress, + user: FileUser, + mode: FileMode, + ) -> core::result::Result, BrokerControlError> { + let Some(address) = address.encode() else { + return Ok(Err(LocalSocketError::InvalidArgument)); + }; + let lease = self.acquire_shared_buffer(address.len())?; + self.request(|local| { + local.bind_local_socket(handle, lease.sequence(), &address, user, mode) + }) + } + + fn listen_local_socket( + &self, + handle: ObjectHandle, + backlog: u32, + ) -> core::result::Result, BrokerControlError> { + self.request(|local| local.listen_local_socket(handle, backlog)) + } + + fn connect_local_socket( + &self, + handle: ObjectHandle, + address: &LocalSocketAddress, + user: FileUser, + ) -> core::result::Result, BrokerControlError> { + let Some(address) = address.encode() else { + return Ok(Err(LocalSocketError::InvalidArgument)); + }; + let lease = self.acquire_shared_buffer(address.len())?; + self.request(|local| local.connect_local_socket(handle, lease.sequence(), &address, user)) + } + + fn accept_local_socket( + &self, + handle: ObjectHandle, + flags: FileOpenFlags, + ) -> core::result::Result< + core::result::Result, + BrokerControlError, + > { + self.request(|local| local.accept_local_socket(handle, flags)) + } + + fn send_local_socket( + &self, + handle: ObjectHandle, + address: Option<&LocalSocketAddress>, + data: &[u8], + user: FileUser, + ) -> core::result::Result, BrokerControlError> + { + let mut staged = match address { + Some(address) => match address.encode() { + Some(address) => address, + None => return Ok(Err(LocalSocketError::InvalidArgument)), + }, + None => Vec::new(), + }; + let address_length = staged.len(); + if data.len() > MAX_LOCAL_SOCKET_TRANSFER_SIZE as usize - address_length { + return Ok(Err(LocalSocketError::MessageTooLarge)); + } + staged + .try_reserve_exact(data.len()) + .map_err(|_| BrokerControlError::Broker(ErrorCode::OutOfMemory))?; + staged.extend_from_slice(data); + let lease = self.acquire_shared_buffer(staged.len())?; + self.request(|local| { + local.send_local_socket(handle, lease.sequence(), &staged, address_length, user) + }) + } + + fn receive_local_socket( + &self, + handle: ObjectHandle, + data: &mut [u8], + peek: bool, + nonblocking: bool, + ) -> core::result::Result< + core::result::Result<(ReceiveLocalSocketResponse, LocalSocketName), LocalSocketError>, + BrokerControlError, + > { + let capacity = data.len().min(LOCAL_SOCKET_BUFFER_SIZE as usize); + let data = &mut data[..capacity]; + let lease = + self.acquire_shared_buffer(capacity + MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE as usize)?; + self.request(|local| { + local.receive_local_socket(handle, lease.sequence(), data, peek, nonblocking) + }) + } + + fn shutdown_local_socket( + &self, + handle: ObjectHandle, + mode: ShutdownMode, + ) -> core::result::Result, BrokerControlError> { + self.request(|local| local.shutdown_local_socket(handle, mode)) + } + + fn local_socket_name( + &self, + handle: ObjectHandle, + peer: bool, + ) -> core::result::Result< + core::result::Result, + BrokerControlError, + > { + let lease = self.acquire_shared_buffer(MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE as usize)?; + self.request(|local| local.local_socket_name(handle, peer, lease.sequence())) + } + + fn set_local_socket_option( + &self, + handle: ObjectHandle, + option: LocalSocketOption, + ) -> core::result::Result<(), BrokerControlError> { + self.request(|local| local.set_local_socket_option(handle, option)) + } + + fn local_socket_options( + &self, + handle: ObjectHandle, + ) -> core::result::Result { + self.request(|local| local.local_socket_options(handle)) + } + fn close_object(&self, handle: ObjectHandle) -> core::result::Result<(), BrokerControlError> { self.request(|local| local.close_object(handle)) } diff --git a/litebox/src/lib.rs b/litebox/src/lib.rs index 76a6f3ee14..5ffe36289f 100644 --- a/litebox/src/lib.rs +++ b/litebox/src/lib.rs @@ -19,6 +19,7 @@ extern crate alloc; pub mod event; pub mod fd; pub mod fs; +pub mod local_sockets; pub mod mm; pub mod net; pub mod path; diff --git a/litebox/src/local_sockets.rs b/litebox/src/local_sockets.rs new file mode 100644 index 0000000000..642fe76e56 --- /dev/null +++ b/litebox/src/local_sockets.rs @@ -0,0 +1,480 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +//! Broker-owned local sockets, which connect processes through names in the +//! filesystem or in an abstract namespace. +//! +//! The broker owns each socket's state, including its name, connection, +//! queued data, and options, so processes that share a socket through +//! inheritance see one socket. + +use alloc::sync::{Arc, Weak}; +use core::time::Duration; +use litebox_broker_protocol::{ + ObjectHandle, + error::ErrorCode, + fs::{FileMode, FileOpenFlags, FileStatusFlags, FileUser}, + local_socket::{ + LOCAL_SOCKET_BUFFER_SIZE, LocalSocketAddress, LocalSocketError as ProtocolError, + LocalSocketName, LocalSocketOption, LocalSocketOptions, + }, + readiness::ReadinessFlags, + socket::{ShutdownMode, SocketType}, +}; +use litebox_platform::time::TimeProvider; + +use crate::{ + LiteBox, + broker::{ + BrokerControl, BrokerPollableRegistry, + error::{BrokerControlError, BrokerObjectError}, + }, + event::{ + Events, IOPollable, + observer::Observer, + polling::{Pollee, TryOpError}, + wait::{WaitContext, WaitError}, + }, + fs::errors::StatusFlagsError, + process::ProcessError, + sync::RawSyncPrimitivesProvider, +}; + +use errors::LocalSocketError; + +/// Data a [`LocalSocket::receive`] copied into its buffer. +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct Received { + /// Number of bytes copied into the buffer. + pub received: usize, + /// Length of the whole datagram, or `received` for a stream socket. + pub length: usize, + /// Name of the sending socket. + pub source: LocalSocketName, +} + +/// A reference to a broker-owned local socket, which closes its handle when +/// dropped. +pub struct LocalSocket { + broker: Arc, + handle: ObjectHandle, + socket_type: SocketType, + pollable_registry: Arc>, + pollee: Arc>, +} + +impl LiteBox { + /// Creates an unnamed, unconnected local socket. + /// + /// `socket_type` must be [`SocketType::Stream`] or + /// [`SocketType::Datagram`], and `flags` must be within + /// [`FileOpenFlags::STATUS`]. + pub fn create_local_socket( + &self, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> Result, LocalSocketError> { + let broker = self.broker_control().ok_or(LocalSocketError::Io)?; + let handle = broker + .create_local_socket(socket_type, flags) + .map_err(LocalSocketError::from)?; + Ok(LocalSocket::new(self, broker, handle, socket_type)) + } + + /// Creates a pair of unnamed local sockets connected to each other. + /// + /// The arguments are as for [`Self::create_local_socket`]. + pub fn create_local_socket_pair( + &self, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> Result<(LocalSocket, LocalSocket), LocalSocketError> { + let broker = self.broker_control().ok_or(LocalSocketError::Io)?; + let response = broker + .create_local_socket_pair(socket_type, flags) + .map_err(LocalSocketError::from)?; + Ok(( + LocalSocket::new(self, Arc::clone(&broker), response.first, socket_type), + LocalSocket::new(self, broker, response.second, socket_type), + )) + } + + /// Returns the local socket this process inherited from its parent as + /// `handle`. + /// + /// The socket owns `handle`, so callers adopt each handle once and share + /// the socket for every other use. + pub fn adopt_inherited_local_socket( + &self, + handle: ObjectHandle, + ) -> Result, ProcessError> { + let broker = self.broker_control().ok_or(ProcessError::Unavailable)?; + let socket_type = match broker.local_socket_options(handle) { + Ok(response) => response.socket_type, + Err(error) => { + let _ = broker.close_object(handle); + return Err(error.into()); + } + }; + Ok(LocalSocket::new(self, broker, handle, socket_type)) + } +} + +impl LocalSocket { + fn new( + litebox: &LiteBox, + broker: Arc, + handle: ObjectHandle, + socket_type: SocketType, + ) -> Self { + let pollable_registry = litebox.broker_pollable_registry(); + let pollee = Arc::new(Pollee::new()); + pollable_registry.register_pollable(handle, &pollee); + Self { + broker, + handle, + socket_type, + pollable_registry, + pollee, + } + } + + /// The broker handle this socket owns. + pub(crate) fn handle(&self) -> ObjectHandle { + self.handle + } + + /// The socket's type. + pub fn socket_type(&self) -> SocketType { + self.socket_type + } + + /// Binds the socket to `address`. + /// + /// A path address creates a filesystem node with `mode` on behalf of + /// `user`, failing if the path exists. + pub fn bind( + &self, + address: &LocalSocketAddress, + user: FileUser, + mode: FileMode, + ) -> Result<(), LocalSocketError> { + flatten( + self.broker + .bind_local_socket(self.handle, address, user, mode), + ) + } + + /// Makes a bound stream socket accept connections, with up to `backlog` + /// connections waiting beyond the first. + pub fn listen(&self, backlog: u32) -> Result<(), LocalSocketError> { + flatten(self.broker.listen_local_socket(self.handle, backlog)) + } + + /// Connects a stream socket to the listener at `address`, or sets the + /// default destination of a datagram socket. + /// + /// A stream connection waits while the listener's backlog is full, unless + /// `nonblock` is set or the socket is non-blocking, for at most the + /// socket's send timeout. + pub fn connect( + &self, + cx: &WaitContext<'_, Platform>, + address: &LocalSocketAddress, + user: FileUser, + nonblock: bool, + ) -> Result<(), LocalSocketError> { + self.wait(cx, nonblock, send_timeout, || { + self.broker.connect_local_socket(self.handle, address, user) + }) + } + + /// Accepts a connection from a listening stream socket, giving the new + /// socket the status flags `flags`. + /// + /// Waits while no connection is pending, unless `nonblock` is set or the + /// socket is non-blocking, for at most the socket's receive timeout. + pub fn accept( + &self, + cx: &WaitContext<'_, Platform>, + flags: FileOpenFlags, + nonblock: bool, + ) -> Result { + let handle = self.wait(cx, nonblock, receive_timeout, || { + self.broker.accept_local_socket(self.handle, flags) + })?; + let pollee = Arc::new(Pollee::new()); + self.pollable_registry.register_pollable(handle, &pollee); + Ok(Self { + broker: Arc::clone(&self.broker), + handle, + socket_type: SocketType::Stream, + pollable_registry: Arc::clone(&self.pollable_registry), + pollee, + }) + } + + /// Sends `data` to `address`, or to the connected peer if `address` is + /// `None`, returning the number of bytes sent. + /// + /// A datagram socket sends `data` as one datagram. A stream socket sends + /// all of `data`. Either waits while the receiver is full, unless + /// `nonblock` is set or the socket is non-blocking, for at most the + /// socket's send timeout each time. A stream socket that stops early + /// after sending some bytes returns their count instead of failing. + pub fn send( + &self, + cx: &WaitContext<'_, Platform>, + address: Option<&LocalSocketAddress>, + data: &[u8], + user: FileUser, + nonblock: bool, + ) -> Result { + if self.socket_type != SocketType::Stream { + return self.wait(cx, nonblock, send_timeout, || { + self.broker + .send_local_socket(self.handle, address, data, user) + }); + } + let mut sent: usize = 0; + loop { + let end = sent + .saturating_add(LOCAL_SOCKET_BUFFER_SIZE as usize) + .min(data.len()); + let chunk = &data[sent..end]; + let result = self.wait(cx, nonblock, send_timeout, || { + self.broker + .send_local_socket(self.handle, address, chunk, user) + }); + match result { + Ok(written) => sent += written, + Err(_) if sent != 0 => return Ok(sent), + Err(error) => return Err(error), + } + if sent == data.len() { + return Ok(sent); + } + } + } + + /// Receives into `buffer`, leaving the data queued if `peek` is set. + /// + /// A stream socket receives queued bytes. A datagram socket receives one + /// datagram, discarding the bytes that do not fit unless peeking. Waits + /// while no data is queued, unless `nonblock` is set or the socket is + /// non-blocking, for at most the socket's receive timeout. Receiving zero + /// bytes into a nonempty buffer means no more data can arrive, except that + /// a datagram socket that would not wait fails with + /// [`LocalSocketError::WouldBlock`] instead. + pub fn receive( + &self, + cx: &WaitContext<'_, Platform>, + buffer: &mut [u8], + peek: bool, + nonblock: bool, + ) -> Result { + self.wait(cx, nonblock, receive_timeout, || { + self.broker + .receive_local_socket(self.handle, buffer, peek, nonblock) + .map(|result| { + result.map(|(received, source)| Received { + received: received.received as usize, + length: received.length as usize, + source, + }) + }) + }) + } + + /// Shuts down one or both directions of the socket. + pub fn shutdown(&self, mode: ShutdownMode) -> Result<(), LocalSocketError> { + flatten(self.broker.shutdown_local_socket(self.handle, mode)) + } + + /// Returns the socket's name, or its peer's name if `peer` is set. + pub fn name(&self, peer: bool) -> Result { + flatten(self.broker.local_socket_name(self.handle, peer)) + } + + /// Stores one option, which every reference to the socket shares. + pub fn set_option(&self, option: LocalSocketOption) -> Result<(), LocalSocketError> { + Ok(self.broker.set_local_socket_option(self.handle, option)?) + } + + /// Returns the socket's stored options. + pub fn options(&self) -> Result { + Ok(self.broker.local_socket_options(self.handle)?.options) + } + + /// Returns the socket's access mode and status flags. + pub fn get_status_flags(&self) -> Result { + Ok(self.broker.get_status_flags(self.handle)?) + } + + /// Changes the status flags in `mask`, within [`FileOpenFlags::STATUS`], + /// to their values in `flags`. + /// + /// Every reference to the socket sees the change. + pub fn set_status_flags( + &self, + mask: FileOpenFlags, + flags: FileOpenFlags, + ) -> Result<(), StatusFlagsError> { + Ok(self.broker.set_status_flags(self.handle, mask, flags)?) + } + + /// Runs `op`, waiting while the broker reports it would block, unless + /// `nonblock` is set, for at most the stored timeout `timeout` selects. + /// + /// The broker publishes a blocked socket's readiness when the operation + /// may succeed, so any readiness notification retries it. The timeout is + /// read only once `op` would block, so operations that need not wait cost + /// one request. + fn wait( + &self, + cx: &WaitContext<'_, Platform>, + nonblock: bool, + timeout: fn(&LocalSocketOptions) -> Option, + mut op: impl FnMut() -> Result, BrokerControlError>, + ) -> Result { + let mut attempt = || match op() { + Ok(Ok(value)) => Ok(value), + Ok(Err(error)) => Err(TryOpError::Other(LocalSocketError::Socket(error))), + Err(BrokerControlError::Broker(ErrorCode::NonBlockingWouldBlock)) => { + Err(TryOpError::Other(LocalSocketError::WouldBlock)) + } + Err(error) => match BrokerObjectError::from(error) { + BrokerObjectError::WouldBlock => Err(TryOpError::TryAgain), + error => Err(TryOpError::Other(error.into())), + }, + }; + match attempt() { + Ok(value) => return Ok(value), + Err(TryOpError::TryAgain) if !nonblock => {} + Err(TryOpError::TryAgain) => return Err(LocalSocketError::WouldBlock), + Err(TryOpError::Other(error)) => return Err(error), + Err(TryOpError::WaitError(error)) => return Err(LocalSocketError::WaitError(error)), + } + let timeout = timeout(&self.options()?); + self.pollee + .wait( + &cx.with_timeout(timeout), + false, + Events::IN | Events::OUT, + attempt, + ) + .map_err(|error| match error { + TryOpError::TryAgain | TryOpError::WaitError(WaitError::TimedOut) => { + LocalSocketError::WouldBlock + } + TryOpError::WaitError(WaitError::Interrupted) if timeout.is_some() => { + LocalSocketError::Interrupted + } + TryOpError::WaitError(error) => LocalSocketError::WaitError(error), + TryOpError::Other(error) => error, + }) + } +} + +impl IOPollable for LocalSocket { + fn register_observer(&self, observer: Weak>, filter: Events) { + self.pollee.register_observer(observer, filter); + } + + fn check_io_events(&self) -> Events { + self.broker + .check_readiness(self.handle) + .map_or(Events::ERR, local_socket_events) + } +} + +impl Drop for LocalSocket { + fn drop(&mut self) { + self.pollable_registry.unregister_pollable(self.handle); + let _ = self.broker.close_object(self.handle); + } +} + +fn receive_timeout(options: &LocalSocketOptions) -> Option { + options.receive_timeout +} + +fn send_timeout(options: &LocalSocketOptions) -> Option { + options.send_timeout +} + +/// Maps local socket readiness to events: a socket that can no longer +/// receive reports [`Events::RDHUP`], and one that can neither send nor +/// receive, or is a stream socket that is not connected, reports +/// [`Events::HUP`]. +fn local_socket_events(readiness: ReadinessFlags) -> Events { + let mut events = Events::empty(); + events.set(Events::IN, readiness.contains(ReadinessFlags::READ)); + events.set(Events::OUT, readiness.contains(ReadinessFlags::WRITE)); + events.set(Events::RDHUP, readiness.contains(ReadinessFlags::HANGUP)); + events.set(Events::HUP, readiness.contains(ReadinessFlags::CLOSED)); + events.set(Events::ERR, readiness.contains(ReadinessFlags::ERROR)); + events +} + +fn flatten( + result: Result, BrokerControlError>, +) -> Result { + result?.map_err(LocalSocketError::Socket) +} + +impl From for LocalSocketError { + fn from(error: BrokerControlError) -> Self { + BrokerObjectError::from(error).into() + } +} + +impl From for LocalSocketError { + fn from(error: BrokerObjectError) -> Self { + match error { + BrokerObjectError::ResourceExhausted => Self::ResourceExhausted, + BrokerObjectError::OutOfMemory => Self::OutOfMemory, + BrokerObjectError::PermissionDenied => Self::PermissionDenied, + BrokerObjectError::UnsupportedOperation => Self::Unsupported, + BrokerObjectError::Control + | BrokerObjectError::InvalidObject + | BrokerObjectError::WouldBlock + | BrokerObjectError::PeerClosed => Self::Io, + } + } +} + +pub mod errors { + use thiserror::Error; + + use crate::event::wait::WaitError; + + /// Possible errors from local socket operations. + #[non_exhaustive] + #[derive(Error, Debug)] + pub enum LocalSocketError { + /// The socket operation failed in a way meaningful to the guest. + #[error(transparent)] + Socket(litebox_broker_protocol::local_socket::LocalSocketError), + /// The operation would block and must not wait, or its socket's + /// timeout expired. + #[error("local socket operation would block")] + WouldBlock, + #[error("wait error")] + WaitError(WaitError), + /// A wait bounded by the socket's timeout was interrupted, so the + /// operation cannot restart without restarting the timeout. + #[error("local socket wait with a timeout was interrupted")] + Interrupted, + #[error("local socket resource exhausted")] + ResourceExhausted, + #[error("local socket memory allocation failed")] + OutOfMemory, + #[error("local socket permission denied")] + PermissionDenied, + #[error("local socket type or flags are unsupported")] + Unsupported, + #[error("local socket broker I/O failed")] + Io, + } +} diff --git a/litebox/src/process.rs b/litebox/src/process.rs index b63a989d30..4936615d3b 100644 --- a/litebox/src/process.rs +++ b/litebox/src/process.rs @@ -21,6 +21,7 @@ use crate::broker::{ }; use crate::event::{Events, IOPollable, observer::Observer, polling::Pollee}; use crate::fs::FileFd; +use crate::local_sockets::LocalSocket; use crate::pipes::PipeFd; use crate::sync::RawSyncPrimitivesProvider; @@ -115,6 +116,8 @@ pub enum InheritableFd { File(Arc), /// A pipe end, adopted with [`LiteBox::adopt_inherited_pipe`]. Pipe(Arc>), + /// A local socket, adopted with [`LiteBox::adopt_inherited_local_socket`]. + LocalSocket(Arc>), } /// Termination state of a child process. @@ -194,6 +197,7 @@ impl Process { // Holding the objects keeps their handles open until the child has its own. let mut held_files = Vec::new(); let mut held_pipes = Vec::new(); + let mut held_sockets = Vec::new(); let mut fd_handles = Vec::new(); for fd in fds { let handle = match fd { @@ -213,6 +217,10 @@ impl Process { held_pipes.push(pipe); handle } + InheritableFd::LocalSocket(socket) => { + held_sockets.push(Arc::clone(socket)); + socket.handle() + } }; fd_handles.push(handle); } diff --git a/litebox_broker_core/src/fs/mod.rs b/litebox_broker_core/src/fs/mod.rs index 840a08650c..3dcd6fafe5 100644 --- a/litebox_broker_core/src/fs/mod.rs +++ b/litebox_broker_core/src/fs/mod.rs @@ -31,7 +31,7 @@ pub use litebox_broker_protocol::fs::{ FileDirectoryEntry as DirEntry, FileMode as Mode, FileNodeInfo as NodeInfo, FileSeekWhence as SeekWhence, FileStatus, FileType, FileUser as UserInfo, }; -pub(crate) use service::{File, get_status_flags, set_status_flags}; +pub(crate) use service::{File, get_status_flags, open_node, set_status_flags}; pub use service::{ FileResult, FileService, UnsupportedFileService, chmod, chown, handle_status, is_terminal, mkdir, open, path_status, read, read_directory, rmdir, seek, truncate, unlink, write, diff --git a/litebox_broker_core/src/fs/service.rs b/litebox_broker_core/src/fs/service.rs index 74cc3ba63a..a058b47a41 100644 --- a/litebox_broker_core/src/fs/service.rs +++ b/litebox_broker_core/src/fs/service.rs @@ -473,6 +473,33 @@ pub fn open( .map(Ok) } +/// Opens `path` for broker-internal use, returning the open file and its +/// status without installing a reference. +/// +/// Dropping the file can reach the backend, so callers drop it outside their +/// locks. +pub(crate) fn open_node( + process: &BrokerProcess, + path: &str, + user: FileUser, + access: FileAccessMode, + flags: FileOpenFlags, + mode: FileMode, +) -> Result> { + if let Err(error) = validate_path(path) { + return Ok(Err(error)); + } + let file = match process.core.fs.open(path, user, access, flags, mode)? { + Ok(file) => file, + Err(error) => return Ok(Err(error)), + }; + let status = match process.core.fs.handle_status(&file)? { + Ok(status) => status, + Err(error) => return Ok(Err(error)), + }; + Ok(Ok((file, status))) +} + /// Reads bytes from a broker-owned open file. /// /// Fails with [`BrokerError::WouldBlock`] while a file that publishes readiness has nothing to diff --git a/litebox_broker_core/src/lib.rs b/litebox_broker_core/src/lib.rs index e564fb7c2b..05cb930cc3 100644 --- a/litebox_broker_core/src/lib.rs +++ b/litebox_broker_core/src/lib.rs @@ -22,6 +22,7 @@ mod error; pub mod event; pub mod fs; mod id; +pub mod local_socket; mod object; pub mod pipe; mod policy; @@ -80,6 +81,10 @@ pub struct BrokerCoreLimits { pub max_total_pipe_capacity: usize, /// Maximum capacity in bytes reserved by live pipes created by one process. pub max_pipe_capacity_per_process: usize, + /// Maximum total bytes queued in local sockets across all processes. + pub max_total_local_socket_bytes: usize, + /// Maximum bytes queued in local sockets created by one process. + pub max_local_socket_bytes_per_process: usize, /// Maximum live platform socket resources across all processes. pub max_sockets: usize, /// Maximum live platform socket resources owned by one process. @@ -102,6 +107,8 @@ impl BrokerCoreLimits { max_references_per_process: 1024, max_total_pipe_capacity: 64 * 1024 * 1024, max_pipe_capacity_per_process: 16 * 1024 * 1024, + max_total_local_socket_bytes: 64 * 1024 * 1024, + max_local_socket_bytes_per_process: 16 * 1024 * 1024, max_sockets: 1024, max_sockets_per_process: 256, max_threads: 4096, @@ -121,6 +128,8 @@ impl BrokerCoreLimits { max_references_per_process: max_references, max_total_pipe_capacity, max_pipe_capacity_per_process: max_total_pipe_capacity, + max_total_local_socket_bytes: Self::DEFAULT.max_total_local_socket_bytes, + max_local_socket_bytes_per_process: Self::DEFAULT.max_local_socket_bytes_per_process, max_sockets: Self::DEFAULT.max_sockets, max_sockets_per_process: Self::DEFAULT.max_sockets_per_process, max_threads: Self::DEFAULT.max_threads, @@ -146,6 +155,8 @@ impl BrokerCoreLimits { max_references_per_process: max_references, max_total_pipe_capacity, max_pipe_capacity_per_process: max_total_pipe_capacity, + max_total_local_socket_bytes: Self::DEFAULT.max_total_local_socket_bytes, + max_local_socket_bytes_per_process: Self::DEFAULT.max_local_socket_bytes_per_process, max_sockets, max_sockets_per_process, max_threads: Self::DEFAULT.max_threads, @@ -172,6 +183,24 @@ impl BrokerCoreLimits { } } + /// Returns these limits with explicit broker-wide and per-process quotas + /// for bytes queued in local sockets. + /// + /// A per-process quota above the broker-wide limit is accepted; the + /// broker-wide limit still applies. + #[must_use] + pub const fn with_local_socket_limits( + self, + max_total_local_socket_bytes: usize, + max_local_socket_bytes_per_process: usize, + ) -> Self { + Self { + max_total_local_socket_bytes, + max_local_socket_bytes_per_process, + ..self + } + } + /// Returns these limits with an explicit broker process limit. #[must_use] pub const fn with_process_limit(self, max_processes: usize) -> Self { @@ -241,6 +270,7 @@ pub struct BrokerCore { pub(crate) pending_references: Arc, pub(crate) reserved_pipe_capacity: Arc, pub(crate) reserved_sockets: Arc, + pub(crate) local_sockets: Arc, /// Bytes of child memory images held by the broker. pub(crate) reserved_child_image_size: Arc, pub(crate) random_provider: Arc, @@ -303,6 +333,7 @@ impl BrokerCore { pending_references: Arc::new(AtomicUsize::new(0)), reserved_pipe_capacity: Arc::new(AtomicUsize::new(0)), reserved_sockets: Arc::new(AtomicUsize::new(0)), + local_sockets: Arc::new(local_socket::LocalSockets::new(&limits)), reserved_child_image_size: Arc::new(AtomicU64::new(0)), random_provider, socket_provider, diff --git a/litebox_broker_core/src/local_socket.rs b/litebox_broker_core/src/local_socket.rs new file mode 100644 index 0000000000..61ec82cb29 --- /dev/null +++ b/litebox_broker_core/src/local_socket.rs @@ -0,0 +1,1359 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +//! Broker-owned local sockets. +//! +//! Every local socket lives in one broker-wide table behind a single lock, +//! since connecting, sending, and closing change the sockets at both ends of +//! a connection. A socket stays in the table while any reference to it +//! exists, or while it waits in a listener's backlog. Operations call the +//! file service and drop files only after releasing the table lock. +//! +//! Readiness reaches every reference's registration from the moment the +//! reference exists, since another process can change a socket through a +//! connection without sharing a reference to it. + +use alloc::collections::VecDeque; +use alloc::sync::Arc; +use alloc::vec::Vec; +use core::sync::atomic::{AtomicUsize, Ordering}; + +use hashbrown::HashMap; +use litebox_broker_protocol::ObjectHandle; +use litebox_broker_protocol::fs::{ + FileAccessMode, FileError, FileMode, FileOpenFlags, FileStatusFlags, FileUser, +}; +use litebox_broker_protocol::local_socket::{ + LOCAL_SOCKET_BUFFER_SIZE, LocalSocketAddress, LocalSocketError, LocalSocketName, + LocalSocketOption, LocalSocketOptions, MAX_LOCAL_SOCKET_BACKLOG, + MAX_LOCAL_SOCKET_TRANSFER_SIZE, +}; +use litebox_broker_protocol::readiness::ReadinessFlags; +use litebox_broker_protocol::socket::{ShutdownMode, SocketType}; +use spin::{Mutex, MutexGuard, rwlock::RwLock}; + +use crate::fs::File; +use crate::object::{ObjectEntry, ObjectRights}; +use crate::readiness::{ReadinessRegistration, ReadinessSink, ReadinessWatchers}; +use crate::{BrokerCoreLimits, BrokerError, BrokerProcess, Result}; + +/// Guest-visible result of a local socket operation. +pub type LocalSocketResult = core::result::Result; + +/// Bytes a socket queues before senders to it wait. +const CAPACITY: usize = LOCAL_SOCKET_BUFFER_SIZE as usize; + +/// Bytes charged for each queued datagram beyond its data, so empty +/// datagrams still consume queue space and quota. +const DATAGRAM_OVERHEAD: usize = 256; + +/// Bytes charged to a listener for each connection waiting in its backlog, +/// which holds no reference until accepted. +const CONNECTION_OVERHEAD: usize = 256; + +/// Data received from a local socket. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct Received { + /// Received bytes. + pub data: Vec, + /// Length of the whole datagram, or of `data` for a stream socket. + pub length: usize, + /// Name of the sending socket. + pub source: LocalSocketName, +} + +/// Creates an unconnected local socket. +/// +/// The socket starts with the status flags `flags`, which must be within +/// [`FileOpenFlags::STATUS`]. Its reference publishes readiness through +/// `readiness_sink`. +pub fn create( + process: &BrokerProcess, + socket_type: SocketType, + flags: FileOpenFlags, + readiness_sink: &Arc, +) -> Result { + let rights = creation_rights(process, flags)?; + let reference = process.reserve_object_reference(rights)?; + let registration = ReadinessRegistration::new(reference.handle(), Arc::clone(readiness_sink)); + let object = new_object(process, socket_type, flags)?; + object.watch(®istration)?; + reference.commit_with_readiness(ObjectEntry::LocalSocket(object), Some(registration)) +} + +/// Creates a pair of local sockets connected to each other. +/// +/// Both sockets start with the status flags `flags`, which must be within +/// [`FileOpenFlags::STATUS`]. Their references publish readiness through +/// `readiness_sink`. +pub fn create_pair( + process: &BrokerProcess, + socket_type: SocketType, + flags: FileOpenFlags, + readiness_sink: &Arc, +) -> Result<(ObjectHandle, ObjectHandle)> { + let rights = creation_rights(process, flags)?; + let first_reference = process.reserve_object_reference(rights)?; + let second_reference = process.reserve_object_reference(rights)?; + let first_registration = + ReadinessRegistration::new(first_reference.handle(), Arc::clone(readiness_sink)); + let second_registration = + ReadinessRegistration::new(second_reference.handle(), Arc::clone(readiness_sink)); + let first = new_object(process, socket_type, flags)?; + let second = new_object(process, socket_type, flags)?; + { + let mut table = process.core.local_sockets.lock(); + for (id, peer) in [(first.id, second.id), (second.id, first.id)] { + table.socket_mut(id)?.connection = Connection::Connected { + peer: Some(peer), + peer_name: LocalSocketName::Unnamed, + }; + } + } + first.watch(&first_registration)?; + second.watch(&second_registration)?; + let first = first_reference + .commit_with_readiness(ObjectEntry::LocalSocket(first), Some(first_registration))?; + match second_reference + .commit_with_readiness(ObjectEntry::LocalSocket(second), Some(second_registration)) + { + Ok(second) => Ok((first, second)), + Err(error) => { + process.close_object_reference(first)?; + Err(error) + } + } +} + +/// Binds a local socket to `address`. +/// +/// Binding to a path creates a file there with `mode` as `user`, which fails +/// with [`LocalSocketError::AddressInUse`] if the path already exists. +pub fn bind( + process: &BrokerProcess, + handle: ObjectHandle, + address: &LocalSocketAddress, + user: FileUser, + mode: FileMode, +) -> Result> { + let lease = Lease::new(process.authorized_object(handle, ObjectRights::WRITE)?)?; + let path = { + let mut table = lease.lock(); + let socket = table.socket_mut(lease.id)?; + if socket.binding || socket.name != LocalSocketName::Unnamed { + return Ok(Err(LocalSocketError::InvalidArgument)); + } + match address { + LocalSocketAddress::Abstract(_) => { + return table.register_name( + lease.id, + NameKey::of(address, lease.socket_type), + address.name(), + &mut None, + ); + } + LocalSocketAddress::Path { path, .. } => { + socket.binding = true; + path + } + } + }; + + let created = crate::fs::open_node( + process, + path, + user, + FileAccessMode::WriteOnly, + FileOpenFlags::CREATE | FileOpenFlags::EXCLUSIVE, + mode, + ); + // Declared before the table guard, so a marker the socket does not keep + // drops after the lock. + let mut marker; + let mut table = lease.lock(); + table.socket_mut(lease.id)?.binding = false; + let key = match created? { + Ok((file, status)) => { + marker = Some(file); + NameKey::Node { + dev: status.node_info.dev, + ino: status.node_info.ino, + } + } + Err(FileError::AlreadyExists) => return Ok(Err(LocalSocketError::AddressInUse)), + Err(error) => return Ok(Err(LocalSocketError::File(error))), + }; + table.register_name(lease.id, key, address.name(), &mut marker) +} + +/// Starts accepting connections on a bound stream socket, or changes the +/// backlog limit of a listening one. +pub fn listen( + process: &BrokerProcess, + handle: ObjectHandle, + backlog: u32, +) -> Result> { + let lease = Lease::new(process.authorized_object(handle, ObjectRights::WRITE)?)?; + if lease.socket_type != SocketType::Stream { + return Ok(Err(LocalSocketError::Unsupported)); + } + let limit = backlog.min(MAX_LOCAL_SOCKET_BACKLOG) as usize; + let mut table = lease.lock(); + let socket = table.socket_mut(lease.id)?; + if socket.name == LocalSocketName::Unnamed { + return Ok(Err(LocalSocketError::InvalidArgument)); + } + match &mut socket.connection { + Connection::None => { + socket.connection = Connection::Listening { + backlog: VecDeque::new(), + limit, + }; + } + Connection::Listening { limit: current, .. } => *current = limit, + Connection::Connected { .. } => return Ok(Err(LocalSocketError::InvalidArgument)), + } + table.wake_waiters(lease.id); + table.publish(lease.id); + Ok(Ok(())) +} + +/// Connects a local socket to the socket bound to `address`. +/// +/// A stream socket queues a new connection on the listening socket there, +/// failing with [`BrokerError::WouldBlock`] while its backlog is full. A +/// datagram socket sets its default destination. +pub fn connect( + process: &BrokerProcess, + handle: ObjectHandle, + address: &LocalSocketAddress, + user: FileUser, +) -> Result> { + let lease = Lease::new(process.authorized_object(handle, ObjectRights::WRITE)?)?; + // The node stays open until after the table guard drops, so its identity + // cannot be reused while the operation runs. + let (key, _node) = match lookup(process, address, user, lease.socket_type)? { + Ok(found) => found, + Err(error) => return Ok(Err(error)), + }; + let mut table = lease.lock(); + let target = match table.target(&key, lease.socket_type)? { + Ok(target) => target, + Err(error) => return Ok(Err(error)), + }; + match lease.socket_type { + SocketType::Stream => table.connect_stream(lease.id, target, lease.nonblocking), + SocketType::Datagram => { + if !table.accepts_datagrams_from(target, lease.id)? { + return Ok(Err(LocalSocketError::NotPermitted)); + } + let peer_name = table.socket(target)?.name.clone(); + let socket = table.socket_mut(lease.id)?; + let previous = core::mem::replace( + &mut socket.connection, + Connection::Connected { + peer: Some(target), + peer_name, + }, + ); + // Like Linux, connecting to another peer discards the queued + // datagrams. + if matches!(previous, Connection::Connected { peer: Some(old), .. } if old != target) { + table.purge_datagrams(lease.id)?; + } + // Senders blocked on this socket recheck whether it still + // accepts their datagrams. + table.wake_waiters(lease.id); + table.publish(lease.id); + Ok(Ok(())) + } + _ => Err(BrokerError::Internal), + } +} + +/// Accepts the oldest queued connection of a listening stream socket. +/// +/// The accepted socket starts with the status flags `flags`, which must be +/// within [`FileOpenFlags::STATUS`], and its reference publishes readiness +/// through `readiness_sink`. Fails with [`BrokerError::WouldBlock`] while no +/// connection is queued. +pub fn accept( + process: &BrokerProcess, + handle: ObjectHandle, + flags: FileOpenFlags, + readiness_sink: &Arc, +) -> Result> { + let rights = creation_rights(process, flags)?; + let lease = Lease::new(process.authorized_object(handle, ObjectRights::WAIT)?)?; + if lease.socket_type != SocketType::Stream { + return Ok(Err(LocalSocketError::Unsupported)); + } + let reference = process.reserve_object_reference(rights)?; + let registration = ReadinessRegistration::new(reference.handle(), Arc::clone(readiness_sink)); + let accepted = { + let mut table = lease.lock(); + let listener = table.socket_mut(lease.id)?; + let Connection::Listening { backlog, .. } = &mut listener.connection else { + return Ok(Err(LocalSocketError::InvalidArgument)); + }; + let Some(accepted) = backlog.pop_front() else { + // Linux reports a shut-down listener as invalid only to callers + // that would otherwise wait. + return if listener.read_shut && !lease.nonblocking { + Ok(Err(LocalSocketError::InvalidArgument)) + } else { + Err(BrokerError::would_block(lease.nonblocking)) + }; + }; + table.refund(lease.id, CONNECTION_OVERHEAD)?; + table.wake_waiters(lease.id); + accepted + }; + let object = LocalSocketObject::new( + Arc::clone(&lease.sockets), + accepted, + SocketType::Stream, + flags, + ); + object.watch(®istration)?; + reference + .commit_with_readiness(ObjectEntry::LocalSocket(object), Some(registration)) + .map(Ok) +} + +/// Sends bytes from a local socket, to `address` if given or otherwise to +/// its connected peer, and returns how many were sent. +/// +/// A stream socket sends as many bytes as its peer has room for. A datagram +/// socket sends `data` as one datagram. Fails with [`BrokerError::WouldBlock`] +/// while the receiving socket is full. +pub fn send( + process: &BrokerProcess, + handle: ObjectHandle, + address: Option<&LocalSocketAddress>, + data: &[u8], + user: FileUser, +) -> Result> { + if data.len() > MAX_LOCAL_SOCKET_TRANSFER_SIZE as usize { + return Err(BrokerError::ResourceExhausted); + } + let lease = Lease::new(process.authorized_object(handle, ObjectRights::WRITE)?)?; + match lease.socket_type { + SocketType::Stream => { + lease + .lock() + .send_stream(lease.id, address.is_some(), data, lease.nonblocking) + } + SocketType::Datagram => { + // The node stays open until after the table guard drops, as for + // `connect`. + let (key, _node) = + match address.map(|address| lookup(process, address, user, SocketType::Datagram)) { + Some(found) => match found? { + Ok((key, node)) => (Some(key), node), + Err(error) => return Ok(Err(error)), + }, + None => (None, None), + }; + let mut table = lease.lock(); + table.send_datagram(lease.id, key.as_ref(), data, lease.nonblocking) + } + _ => Err(BrokerError::Internal), + } +} + +/// Receives up to `capacity` bytes from a local socket, leaving them queued +/// if `peek` is set. +/// +/// A stream socket receives queued bytes. A datagram socket receives one +/// datagram, discarding the bytes beyond `capacity` unless peeking. Empty +/// data with zero length for a nonzero `capacity` means no more data can +/// arrive. Fails with [`BrokerError::WouldBlock`] while nothing is queued. +/// +/// If `nonblocking` is set, the receive acts as for a non-blocking socket: +/// it fails with [`BrokerError::NonBlockingWouldBlock`], and a datagram +/// socket whose receive direction is shut down fails instead of reporting the +/// end of data, as Linux does. +pub fn receive( + process: &BrokerProcess, + handle: ObjectHandle, + capacity: u32, + peek: bool, + nonblocking: bool, +) -> Result> { + if capacity > MAX_LOCAL_SOCKET_TRANSFER_SIZE { + return Err(BrokerError::ResourceExhausted); + } + let lease = Lease::new(process.authorized_object(handle, ObjectRights::WAIT)?)?; + let nonblocking = nonblocking || lease.nonblocking; + let mut table = lease.lock(); + match lease.socket_type { + SocketType::Stream => table.receive_stream(lease.id, capacity as usize, peek, nonblocking), + SocketType::Datagram => { + table.receive_datagram(lease.id, capacity as usize, peek, nonblocking) + } + _ => Err(BrokerError::Internal), + } +} + +/// Shuts down one or both directions of a local socket. +/// +/// Shutting down a direction of a connected stream socket also shuts down +/// the opposite direction of its peer. +pub fn shutdown( + process: &BrokerProcess, + handle: ObjectHandle, + mode: ShutdownMode, +) -> Result> { + let (read, write) = match mode { + ShutdownMode::Read => (true, false), + ShutdownMode::Write => (false, true), + ShutdownMode::Both => (true, true), + _ => return Err(BrokerError::UnsupportedOperation), + }; + let lease = Lease::new(process.authorized_object(handle, ObjectRights::WRITE)?)?; + let mut table = lease.lock(); + let socket = table.socket_mut(lease.id)?; + socket.read_shut |= read; + socket.write_shut |= write; + let peer = match (&socket.connection, socket.kind) { + ( + Connection::Connected { + peer: Some(peer), .. + }, + SocketType::Stream, + ) => Some(*peer), + _ => None, + }; + if let Some(peer) = peer.and_then(|peer| table.sockets.get_mut(&peer)) { + peer.read_shut |= write; + peer.write_shut |= read; + } + table.publish(lease.id); + if let Some(peer) = peer { + table.publish(peer); + } + // Senders waiting for room now fail. + table.wake_waiters(lease.id); + Ok(Ok(())) +} + +/// Returns the name of a local socket, or of its connected peer if `peer` +/// is set. +pub fn name( + process: &BrokerProcess, + handle: ObjectHandle, + peer: bool, +) -> Result> { + let lease = Lease::new( + process + .authorized_object_with_any_rights(handle, ObjectRights::WAIT | ObjectRights::WRITE)?, + )?; + let table = lease.lock(); + let socket = table.socket(lease.id)?; + if !peer { + return Ok(Ok(socket.name.clone())); + } + match &socket.connection { + Connection::Connected { peer_name, .. } => Ok(Ok(peer_name.clone())), + Connection::None | Connection::Listening { .. } => Ok(Err(LocalSocketError::NotConnected)), + } +} + +/// Stores one option of a local socket. +pub fn set_option( + process: &BrokerProcess, + handle: ObjectHandle, + option: LocalSocketOption, +) -> Result<()> { + let lease = Lease::new(process.authorized_object(handle, ObjectRights::WRITE)?)?; + lease.lock().socket_mut(lease.id)?.options.set(option); + Ok(()) +} + +/// Returns the type and stored options of a local socket. +pub fn options( + process: &BrokerProcess, + handle: ObjectHandle, +) -> Result<(SocketType, LocalSocketOptions)> { + let lease = Lease::new( + process + .authorized_object_with_any_rights(handle, ObjectRights::WAIT | ObjectRights::WRITE)?, + )?; + let options = lease.lock().socket(lease.id)?.options; + Ok((lease.socket_type, options)) +} + +fn creation_rights(process: &BrokerProcess, flags: FileOpenFlags) -> Result { + if !FileOpenFlags::STATUS.contains(flags) { + return Err(BrokerError::UnsupportedOperation); + } + process + .core + .policy + .principal_object_rights(process.caller_credential) +} + +fn new_object( + process: &BrokerProcess, + socket_type: SocketType, + flags: FileOpenFlags, +) -> Result { + if !matches!(socket_type, SocketType::Stream | SocketType::Datagram) { + return Err(BrokerError::UnsupportedOperation); + } + let sockets = &process.core.local_sockets; + let id = sockets.lock().insert(Socket::new( + socket_type, + LocalSocketName::Unnamed, + Arc::clone(&process.local_socket_bytes), + ))?; + Ok(LocalSocketObject::new( + Arc::clone(sockets), + id, + socket_type, + flags, + )) +} + +/// Resolves `address` to the key of the name a `socket_type` socket bound to +/// it holds. +/// +/// A path resolves to its node, which requires write permission like Linux. +/// The node is returned open, so callers keep its key from naming another +/// node while they use it. +fn lookup( + process: &BrokerProcess, + address: &LocalSocketAddress, + user: FileUser, + socket_type: SocketType, +) -> Result)>> { + let LocalSocketAddress::Path { path, .. } = address else { + return Ok(Ok((NameKey::of(address, socket_type), None))); + }; + // A path-only open does not copy the node up an overlay, so write + // permission is checked against its status instead. + let (file, status) = match crate::fs::open_node( + process, + path, + user, + FileAccessMode::ReadOnly, + FileOpenFlags::PATH, + FileMode::empty(), + )? { + Ok(opened) => opened, + Err(error) => return Ok(Err(LocalSocketError::File(error))), + }; + let write = if user.user == status.owner.user { + FileMode::WUSR + } else if user.group == status.owner.group { + FileMode::WGRP + } else { + FileMode::WOTH + }; + if !status.mode.contains(write) { + return Ok(Err(LocalSocketError::File(FileError::AccessNotAllowed))); + } + let key = NameKey::Node { + dev: status.node_info.dev, + ino: status.node_info.ino, + }; + Ok(Ok((key, Some(file)))) +} + +impl ObjectEntry { + fn as_local_socket(&self) -> Result<&LocalSocketObject> { + match self { + Self::LocalSocket(socket) => Ok(socket), + _ => Err(BrokerError::InvalidRights), + } + } +} + +/// One local socket's open state, which its references share. +pub(crate) struct LocalSocketObject { + sockets: Arc, + id: SocketId, + socket_type: SocketType, + /// Whether operations that would block fail instead of waiting. + nonblocking: bool, + /// Whether `O_APPEND` is set, which Linux reports but which has no effect on a socket. + append: bool, +} + +impl LocalSocketObject { + fn new( + sockets: Arc, + id: SocketId, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> Self { + Self { + sockets, + id, + socket_type, + nonblocking: flags.contains(FileOpenFlags::NONBLOCKING), + append: flags.contains(FileOpenFlags::APPEND), + } + } + + /// Publishes this socket's readiness changes through `registration` + /// until every clone of `registration` drops. + pub(crate) fn watch(&self, registration: &ReadinessRegistration) -> Result<()> { + self.sockets + .lock() + .socket_mut(self.id)? + .watchers + .watch(registration) + } + + pub(crate) fn readiness(&self) -> ReadinessFlags { + self.sockets.lock().readiness(self.id) + } + + /// Returns the socket's access mode and status flags. + pub(crate) fn get_status_flags(&self) -> FileStatusFlags { + let mut flags = FileOpenFlags::NONE; + for (set, flag) in [ + (self.nonblocking, FileOpenFlags::NONBLOCKING), + (self.append, FileOpenFlags::APPEND), + ] { + if set { + flags = flags | flag; + } + } + FileStatusFlags { + access: FileAccessMode::ReadWrite, + flags, + } + } + + /// Changes the status flags in `mask` to their values in `flags`. + pub(crate) fn set_status_flags(&mut self, mask: FileOpenFlags, flags: FileOpenFlags) { + if mask.contains(FileOpenFlags::NONBLOCKING) { + self.nonblocking = flags.contains(FileOpenFlags::NONBLOCKING); + } + if mask.contains(FileOpenFlags::APPEND) { + self.append = flags.contains(FileOpenFlags::APPEND); + } + } +} + +impl Drop for LocalSocketObject { + fn drop(&mut self) { + let marker = self.sockets.lock().remove(self.id); + drop(marker); + } +} + +/// An authorized socket whose object stays alive until the lease drops. +/// +/// Callers drop table guards before the lease, since dropping the last +/// reference to the object removes the socket from the table. +struct Lease { + _object: Arc>, + sockets: Arc, + id: SocketId, + socket_type: SocketType, + nonblocking: bool, +} + +impl Lease { + fn new(object: Arc>) -> Result { + let (sockets, id, socket_type, nonblocking) = { + let entry = object.read(); + let socket = entry.as_local_socket()?; + ( + Arc::clone(&socket.sockets), + socket.id, + socket.socket_type, + socket.nonblocking, + ) + }; + Ok(Self { + _object: object, + sockets, + id, + socket_type, + nonblocking, + }) + } + + fn lock(&self) -> MutexGuard<'_, Table> { + self.sockets.lock() + } +} + +/// Every local socket of a broker. +pub(crate) struct LocalSockets(Mutex); + +impl LocalSockets { + pub(crate) fn new(limits: &BrokerCoreLimits) -> Self { + Self(Mutex::new(Table { + sockets: HashMap::new(), + names: HashMap::new(), + next_id: 0, + queued: 0, + max_queued: limits.max_total_local_socket_bytes, + max_queued_per_process: limits.max_local_socket_bytes_per_process, + })) + } + + fn lock(&self) -> MutexGuard<'_, Table> { + self.0.lock() + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +struct SocketId(u64); + +/// Identity of a bound name. +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +enum NameKey { + /// The filesystem node a path name created. + Node { dev: u64, ino: u64 }, + /// An abstract name, which sockets of each type hold separately like + /// Linux. + Abstract(SocketType, Vec), +} + +impl NameKey { + /// Returns the key of an abstract address held by a `socket_type` socket. + fn of(address: &LocalSocketAddress, socket_type: SocketType) -> Self { + match address { + LocalSocketAddress::Abstract(bytes) => Self::Abstract(socket_type, bytes.clone()), + LocalSocketAddress::Path { .. } => unreachable!("path names are keyed by their node"), + } + } +} + +struct Socket { + kind: SocketType, + name: LocalSocketName, + /// Whether a bind is creating the socket's path outside the table lock. + binding: bool, + name_key: Option, + /// File created by binding to a path, held open so the node that + /// identifies the name stays allocated. + marker: Option, + options: LocalSocketOptions, + read_shut: bool, + write_shut: bool, + connection: Connection, + /// Bytes queued for a stream socket. + stream: VecDeque, + /// Datagrams queued for a datagram socket. + datagrams: VecDeque, + /// Bytes charged for queued data, or for a listener's backlog. + queued: usize, + /// Quota of the process that created the socket, charged for its queued + /// data. + quota: Arc, + /// Sockets to wake once this socket's queue drains or it closes. + waiters: Vec, + watchers: ReadinessWatchers, +} + +impl Socket { + fn new(kind: SocketType, name: LocalSocketName, quota: Arc) -> Self { + Self { + kind, + name, + binding: false, + name_key: None, + marker: None, + options: LocalSocketOptions::default(), + read_shut: false, + write_shut: false, + connection: Connection::None, + stream: VecDeque::new(), + datagrams: VecDeque::new(), + queued: 0, + quota, + waiters: Vec::new(), + watchers: ReadinessWatchers::default(), + } + } + + fn is_empty(&self) -> bool { + self.stream.is_empty() && self.datagrams.is_empty() + } +} + +enum Connection { + None, + Listening { + /// Connected sockets not yet accepted. + backlog: VecDeque, + /// Backlog length beyond which connecting waits. + limit: usize, + }, + Connected { + /// The peer, or `None` once a stream peer closes. + peer: Option, + /// Name the peer had when the connection formed. + peer_name: LocalSocketName, + }, +} + +struct Datagram { + data: Vec, + source: LocalSocketName, +} + +impl Datagram { + fn charge(&self) -> usize { + self.data.len() + DATAGRAM_OVERHEAD + } +} + +struct Table { + sockets: HashMap, + names: HashMap, + next_id: u64, + /// Bytes charged for data queued in every socket. + queued: usize, + max_queued: usize, + max_queued_per_process: usize, +} + +impl Table { + fn socket(&self, id: SocketId) -> Result<&Socket> { + self.sockets.get(&id).ok_or(BrokerError::Internal) + } + + fn socket_mut(&mut self, id: SocketId) -> Result<&mut Socket> { + self.sockets.get_mut(&id).ok_or(BrokerError::Internal) + } + + fn insert(&mut self, socket: Socket) -> Result { + let id = SocketId(self.next_id); + self.next_id = self + .next_id + .checked_add(1) + .ok_or(BrokerError::ResourceExhausted)?; + self.sockets + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory)?; + self.sockets.insert(id, socket); + Ok(id) + } + + /// Records that socket `id` holds the name `name` identified by `key`, + /// taking `marker` if it succeeds. + fn register_name( + &mut self, + id: SocketId, + key: NameKey, + name: LocalSocketName, + marker: &mut Option, + ) -> Result> { + if self.names.contains_key(&key) { + return Ok(Err(LocalSocketError::AddressInUse)); + } + self.names + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory)?; + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + socket.name = name; + socket.name_key = Some(key.clone()); + socket.marker = marker.take(); + self.names.insert(key, id); + Ok(Ok(())) + } + + /// Returns the socket holding the name `key`, which must have type + /// `socket_type`. + fn target( + &self, + key: &NameKey, + socket_type: SocketType, + ) -> Result> { + let Some(&target) = self.names.get(key) else { + return Ok(Err(LocalSocketError::ConnectionRefused)); + }; + if self.socket(target)?.kind != socket_type { + return Ok(Err(LocalSocketError::WrongType)); + } + Ok(Ok(target)) + } + + /// Returns whether datagram socket `receiver` accepts datagrams from + /// `sender`, which it does unless connected to another socket. + fn accepts_datagrams_from(&self, receiver: SocketId, sender: SocketId) -> Result { + Ok(match self.socket(receiver)?.connection { + Connection::Connected { + peer: Some(peer), .. + } => peer == sender, + _ => true, + }) + } + + fn connect_stream( + &mut self, + id: SocketId, + target: SocketId, + nonblocking: bool, + ) -> Result> { + let client = self.socket(id)?; + match client.connection { + Connection::None => {} + Connection::Connected { .. } => return Ok(Err(LocalSocketError::AlreadyConnected)), + Connection::Listening { .. } => return Ok(Err(LocalSocketError::InvalidArgument)), + } + let client_name = client.name.clone(); + let listener = self.socket(target)?; + let Connection::Listening { backlog, limit } = &listener.connection else { + return Ok(Err(LocalSocketError::ConnectionRefused)); + }; + if listener.read_shut { + return Ok(Err(LocalSocketError::ConnectionRefused)); + } + if backlog.len() > *limit { + self.add_waiter(target, id)?; + return Err(BrokerError::would_block(nonblocking)); + } + let listener_name = listener.name.clone(); + let mut server = Socket::new( + SocketType::Stream, + listener_name.clone(), + Arc::clone(&listener.quota), + ); + server.connection = Connection::Connected { + peer: Some(id), + peer_name: client_name, + }; + let Connection::Listening { backlog, .. } = &mut self.socket_mut(target)?.connection else { + return Err(BrokerError::Internal); + }; + backlog + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory)?; + self.charge(target, CONNECTION_OVERHEAD)?; + let server = match self.insert(server) { + Ok(server) => server, + Err(error) => { + self.refund(target, CONNECTION_OVERHEAD)?; + return Err(error); + } + }; + let Connection::Listening { backlog, .. } = &mut self.socket_mut(target)?.connection else { + return Err(BrokerError::Internal); + }; + backlog.push_back(server); + self.socket_mut(id)?.connection = Connection::Connected { + peer: Some(server), + peer_name: listener_name, + }; + self.publish(target); + self.publish(id); + Ok(Ok(())) + } + + fn send_stream( + &mut self, + id: SocketId, + addressed: bool, + data: &[u8], + nonblocking: bool, + ) -> Result> { + let socket = self.socket(id)?; + let Connection::Connected { peer, .. } = socket.connection else { + return Ok(Err(if addressed { + LocalSocketError::Unsupported + } else { + LocalSocketError::NotConnected + })); + }; + if addressed { + return Ok(Err(LocalSocketError::AlreadyConnected)); + } + if socket.write_shut { + return Ok(Err(LocalSocketError::BrokenPipe)); + } + if data.is_empty() { + return Ok(Ok(0)); + } + let Some(peer) = peer.filter(|peer| self.sockets.get(peer).is_some_and(|p| !p.read_shut)) + else { + return Ok(Err(LocalSocketError::BrokenPipe)); + }; + let available = CAPACITY.saturating_sub(self.socket(peer)?.queued); + if available == 0 { + return Err(BrokerError::would_block(nonblocking)); + } + let length = available.min(data.len()); + // Charging first keeps a rejected send from growing the queue. + self.charge(peer, length)?; + let receiver = self.socket_mut(peer)?; + if receiver.stream.try_reserve(length).is_err() { + self.refund(peer, length)?; + return Err(BrokerError::OutOfMemory); + } + receiver.stream.extend(&data[..length]); + self.publish(peer); + Ok(Ok(length)) + } + + fn send_datagram( + &mut self, + id: SocketId, + key: Option<&NameKey>, + data: &[u8], + nonblocking: bool, + ) -> Result> { + let target = match key { + Some(key) => match self.target(key, SocketType::Datagram)? { + Ok(target) => target, + Err(error) => return Ok(Err(error)), + }, + None => match self.socket(id)?.connection { + Connection::Connected { + peer: Some(peer), .. + } => { + if !self.sockets.contains_key(&peer) { + self.socket_mut(id)?.connection = Connection::None; + // Like Linux, disconnecting from the closed peer + // discards the queued datagrams. + self.purge_datagrams(id)?; + self.wake_waiters(id); + self.publish(id); + return Ok(Err(LocalSocketError::ConnectionRefused)); + } + peer + } + _ => return Ok(Err(LocalSocketError::NotConnected)), + }, + }; + if data.len() > CAPACITY { + return Ok(Err(LocalSocketError::MessageTooLarge)); + } + let socket = self.socket(id)?; + if socket.write_shut { + return Ok(Err(LocalSocketError::BrokenPipe)); + } + let source = socket.name.clone(); + if !self.accepts_datagrams_from(target, id)? { + return Ok(Err(LocalSocketError::NotPermitted)); + } + let receiver = self.socket(target)?; + if receiver.read_shut { + return Ok(Err(LocalSocketError::BrokenPipe)); + } + // A datagram may overrun the capacity, so a socket that is not full + // accepts any datagram, as its readiness reports. + if receiver.queued >= CAPACITY { + self.add_waiter(target, id)?; + return Err(BrokerError::would_block(nonblocking)); + } + let mut copy = Vec::new(); + copy.try_reserve_exact(data.len()) + .map_err(|_| BrokerError::OutOfMemory)?; + copy.extend_from_slice(data); + let datagram = Datagram { data: copy, source }; + self.charge(target, datagram.charge())?; + let receiver = self.socket_mut(target)?; + if receiver.datagrams.try_reserve(1).is_err() { + self.refund(target, datagram.charge())?; + return Err(BrokerError::OutOfMemory); + } + receiver.datagrams.push_back(datagram); + self.publish(target); + Ok(Ok(data.len())) + } + + fn receive_stream( + &mut self, + id: SocketId, + capacity: usize, + peek: bool, + nonblocking: bool, + ) -> Result> { + let socket = self.socket_mut(id)?; + let Connection::Connected { peer_name, .. } = &socket.connection else { + return Ok(Err(LocalSocketError::InvalidArgument)); + }; + let source = peer_name.clone(); + if capacity == 0 { + return Ok(Ok(Received { + source, + ..Received::default() + })); + } + if socket.stream.is_empty() { + return if socket.read_shut { + Ok(Ok(Received { + source, + ..Received::default() + })) + } else { + Err(BrokerError::would_block(nonblocking)) + }; + } + let length = capacity.min(socket.stream.len()); + let mut data = Vec::new(); + data.try_reserve_exact(length) + .map_err(|_| BrokerError::OutOfMemory)?; + if peek { + data.extend(socket.stream.iter().take(length)); + } else { + data.extend(socket.stream.drain(..length)); + if socket.stream.is_empty() { + socket.stream = VecDeque::new(); + } + self.refund(id, length)?; + self.drained(id); + } + Ok(Ok(Received { + data, + length, + source, + })) + } + + fn receive_datagram( + &mut self, + id: SocketId, + capacity: usize, + peek: bool, + nonblocking: bool, + ) -> Result> { + let socket = self.socket_mut(id)?; + let Some(datagram) = socket.datagrams.front() else { + return if socket.read_shut && !nonblocking { + Ok(Ok(Received::default())) + } else { + Err(BrokerError::would_block(nonblocking)) + }; + }; + let length = datagram.data.len(); + let mut data = Vec::new(); + data.try_reserve_exact(capacity.min(length)) + .map_err(|_| BrokerError::OutOfMemory)?; + data.extend_from_slice(&datagram.data[..capacity.min(length)]); + let source = datagram.source.clone(); + if !peek { + let datagram = socket.datagrams.pop_front().ok_or(BrokerError::Internal)?; + self.refund(id, datagram.charge())?; + self.drained(id); + } + Ok(Ok(Received { + data, + length, + source, + })) + } + + /// Charges `bytes` of newly queued data to socket `id`. + fn charge(&mut self, id: SocketId, bytes: usize) -> Result<()> { + let queued = self + .queued + .checked_add(bytes) + .filter(|queued| *queued <= self.max_queued) + .ok_or(BrokerError::ResourceExhausted)?; + let max_queued_per_process = self.max_queued_per_process; + let socket = self.socket_mut(id)?; + socket + .quota + .try_update(Ordering::Relaxed, Ordering::Relaxed, |charged| { + charged + .checked_add(bytes) + .filter(|charged| *charged <= max_queued_per_process) + }) + .map_err(|_| BrokerError::ResourceExhausted)?; + socket.queued += bytes; + self.queued = queued; + Ok(()) + } + + /// Releases the charge for `bytes` of data dequeued from socket `id`. + fn refund(&mut self, id: SocketId, bytes: usize) -> Result<()> { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + Self::release(&mut self.queued, socket, bytes); + Ok(()) + } + + /// Discards the datagrams queued on socket `id`, which is disconnecting + /// from its peer. + fn purge_datagrams(&mut self, id: SocketId) -> Result<()> { + let socket = self.socket_mut(id)?; + socket.datagrams.clear(); + let queued = socket.queued; + self.refund(id, queued) + } + + fn release(total: &mut usize, socket: &mut Socket, bytes: usize) { + socket.queued = socket + .queued + .checked_sub(bytes) + .expect("a socket's charge must cover its queued data"); + *total = total + .checked_sub(bytes) + .expect("the broker's local socket charge must cover every socket's"); + socket + .quota + .try_update(Ordering::Relaxed, Ordering::Relaxed, |charged| { + charged.checked_sub(bytes) + }) + .expect("a process's local socket charge must cover its sockets'"); + } + + /// Wakes the sockets that may send to socket `id` after its queue shrank. + fn drained(&mut self, id: SocketId) { + if let Some(Socket { + connection: Connection::Connected { + peer: Some(peer), .. + }, + .. + }) = self.sockets.get(&id) + { + let peer = *peer; + self.publish(peer); + } + self.wake_waiters(id); + } + + /// Records that socket `waiter` waits for socket `target` to drain or + /// accept a connection. + fn add_waiter(&mut self, target: SocketId, waiter: SocketId) -> Result<()> { + let Self { sockets, .. } = self; + let Some(mut waiters) = sockets + .get_mut(&target) + .map(|target| core::mem::take(&mut target.waiters)) + else { + return Ok(()); + }; + // Pruning closed waiters bounds the list by the live sockets. + waiters.retain(|waiter| sockets.contains_key(waiter)); + let reserved = if waiters.contains(&waiter) { + Ok(()) + } else { + waiters + .try_reserve(1) + .map(|()| waiters.push(waiter)) + .map_err(|_| BrokerError::OutOfMemory) + }; + if let Some(target) = sockets.get_mut(&target) { + target.waiters = waiters; + } + reserved + } + + fn wake_waiters(&mut self, id: SocketId) { + let Some(waiters) = self + .sockets + .get_mut(&id) + .map(|socket| core::mem::take(&mut socket.waiters)) + else { + return; + }; + for waiter in waiters { + self.publish(waiter); + } + } + + /// Wakes the watchers of socket `id` after a change that may make it ready. + /// + /// Callers hold the table lock, so publications follow the changes in + /// order. + fn publish(&mut self, id: SocketId) { + let readiness = self.readiness(id); + if let Some(socket) = self.sockets.get(&id) { + socket.watchers.publish(readiness); + } + } + + /// Returns the readiness of socket `id`. + /// + /// A datagram socket that cannot send because its peer is full starts + /// waiting for the peer to drain, as Linux does when polled. + fn readiness(&mut self, id: SocketId) -> ReadinessFlags { + let Some(socket) = self.sockets.get(&id) else { + return ReadinessFlags::default(); + }; + let socket_type = socket.kind; + let mut readiness = ReadinessFlags::default(); + if socket.read_shut { + readiness = readiness | ReadinessFlags::READ | ReadinessFlags::HANGUP; + if socket.write_shut { + readiness = readiness | ReadinessFlags::CLOSED; + } + } + if !socket.is_empty() { + readiness = readiness | ReadinessFlags::READ; + } + let peer = match &socket.connection { + Connection::Listening { backlog, .. } => { + if !backlog.is_empty() { + readiness = readiness | ReadinessFlags::READ; + } + return readiness; + } + Connection::None => { + if socket_type == SocketType::Stream { + readiness = readiness | ReadinessFlags::CLOSED; + } + None + } + Connection::Connected { peer, .. } => *peer, + }; + let full_peer = peer.filter(|peer| { + self.sockets + .get(peer) + .is_some_and(|peer| !peer.read_shut && peer.queued >= CAPACITY) + }); + match full_peer { + Some(peer) => { + if socket_type == SocketType::Datagram { + // Failing to record the waiter only loses a wakeup the + // waiter's next send would record. + let _ = self.add_waiter(peer, id); + } + readiness + } + None => readiness | ReadinessFlags::WRITE, + } + } + + /// Removes socket `id` once nothing references it, returning the file + /// that held its path name for the caller to drop after the table lock. + fn remove(&mut self, id: SocketId) -> Option { + let mut socket = self.sockets.remove(&id)?; + let queued = socket.queued; + Self::release(&mut self.queued, &mut socket, queued); + if let Some(key) = socket.name_key.take() + && self.names.get(&key) == Some(&id) + { + self.names.remove(&key); + } + match core::mem::replace(&mut socket.connection, Connection::None) { + Connection::Listening { backlog, .. } => { + // Unaccepted sockets have no references, so no names. + for accepted in backlog { + let marker = self.remove(accepted); + debug_assert!(marker.is_none()); + } + } + Connection::Connected { + peer: Some(peer), .. + } if socket.kind == SocketType::Stream => { + if let Some(peer_socket) = self.sockets.get_mut(&peer) { + peer_socket.read_shut = true; + peer_socket.write_shut = true; + if let Connection::Connected { peer, .. } = &mut peer_socket.connection { + *peer = None; + } + } + self.publish(peer); + } + Connection::None | Connection::Connected { .. } => {} + } + for waiter in core::mem::take(&mut socket.waiters) { + self.publish(waiter); + } + socket.marker.take() + } +} + +#[cfg(test)] +mod tests; diff --git a/litebox_broker_core/src/local_socket/tests.rs b/litebox_broker_core/src/local_socket/tests.rs new file mode 100644 index 0000000000..90c42f6c7e --- /dev/null +++ b/litebox_broker_core/src/local_socket/tests.rs @@ -0,0 +1,859 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +use super::*; +use crate::fs::in_mem::InMem; +use crate::fs::inode_allocator::InodeAllocator; +use crate::fs::resolver::Resolver; +use crate::readiness::tests::TestReadinessSink; +use crate::test_platform::TestPlatform; +use crate::test_support::TestBrokerCoreBuilder; +use crate::{BrokerCore, CallerCredential, PolicyEngine}; +use core::time::Duration; +use std::vec; + +const USER: FileUser = FileUser::ROOT; + +struct Fixture { + broker: BrokerCore, + process: Arc, + sink: Arc, + readiness_sink: Arc, +} + +impl Fixture { + fn new() -> Self { + Self::with_limits(BrokerCoreLimits::DEFAULT) + } + + fn with_limits(limits: BrokerCoreLimits) -> Self { + let fs = Resolver::::new(InMem::::new( + InodeAllocator::standalone(), + )); + let broker = TestBrokerCoreBuilder::new(PolicyEngine::with_unauthenticated_rights( + ObjectRights::all(), + )) + .with_limits(limits) + .with_file_service(Arc::new(fs)) + .build() + .unwrap(); + let process = broker + .create_process(CallerCredential::Unauthenticated, None) + .unwrap(); + let sink = Arc::new(TestReadinessSink::default()); + Self { + broker, + process, + readiness_sink: sink.clone(), + sink, + } + } + + fn create(&self, socket_type: SocketType) -> ObjectHandle { + create( + &self.process, + socket_type, + FileOpenFlags::NONE, + &self.readiness_sink, + ) + .unwrap() + } + + fn pair(&self, socket_type: SocketType) -> (ObjectHandle, ObjectHandle) { + create_pair( + &self.process, + socket_type, + FileOpenFlags::NONE, + &self.readiness_sink, + ) + .unwrap() + } + + fn bind(&self, handle: ObjectHandle, address: &LocalSocketAddress) -> LocalSocketResult<()> { + bind(&self.process, handle, address, USER, FileMode::RWXU).unwrap() + } + + fn listener(&self, address: &LocalSocketAddress, backlog: u32) -> ObjectHandle { + let listener = self.create(SocketType::Stream); + self.bind(listener, address).unwrap(); + listen(&self.process, listener, backlog).unwrap().unwrap(); + listener + } + + fn connect( + &self, + handle: ObjectHandle, + address: &LocalSocketAddress, + ) -> Result> { + connect(&self.process, handle, address, USER) + } + + fn accept(&self, listener: ObjectHandle) -> Result> { + accept( + &self.process, + listener, + FileOpenFlags::NONE, + &self.readiness_sink, + ) + } + + fn send(&self, handle: ObjectHandle, data: &[u8]) -> Result> { + send(&self.process, handle, None, data, USER) + } + + fn send_to( + &self, + handle: ObjectHandle, + address: &LocalSocketAddress, + data: &[u8], + ) -> Result> { + send(&self.process, handle, Some(address), data, USER) + } + + fn receive(&self, handle: ObjectHandle, capacity: u32) -> Result> { + receive(&self.process, handle, capacity, false, false) + } + + fn receive_data(&self, handle: ObjectHandle, capacity: u32) -> Vec { + self.receive(handle, capacity).unwrap().unwrap().data + } + + fn readiness(&self, handle: ObjectHandle) -> ReadinessFlags { + self.process.check_readiness(handle).unwrap() + } + + fn take_republished(&self) -> Vec { + core::mem::take(&mut *self.sink.republished.lock().unwrap()) + .into_iter() + .map(|(handle, _)| handle) + .collect() + } + + fn close(&self, handle: ObjectHandle) { + self.process.close_object_reference(handle).unwrap(); + } + + /// Returns the broker-wide and process charges for queued bytes. + fn queued(&self) -> (usize, usize) { + ( + self.process.core.local_sockets.lock().queued, + self.process.local_socket_bytes.load(Ordering::Relaxed), + ) + } +} + +fn path(path: &str) -> LocalSocketAddress { + LocalSocketAddress::Path { + path: path.into(), + name: path.as_bytes().to_vec(), + } +} + +fn abstract_name(name: &[u8]) -> LocalSocketAddress { + LocalSocketAddress::Abstract(name.to_vec()) +} + +#[test] +fn stream_pairs_carry_bytes_in_both_directions() { + let fixture = Fixture::new(); + let (first, second) = fixture.pair(SocketType::Stream); + assert!(!fixture.readiness(second).contains(ReadinessFlags::READ)); + assert_eq!(fixture.send(first, b"hello").unwrap(), Ok(5)); + assert!(fixture.take_republished().contains(&second)); + assert!(fixture.readiness(second).contains(ReadinessFlags::READ)); + + let peeked = receive(&fixture.process, second, 3, true, false) + .unwrap() + .unwrap(); + assert_eq!(peeked.data, b"hel"); + assert_eq!(peeked.source, LocalSocketName::Unnamed); + assert_eq!(fixture.receive_data(second, 3), b"hel"); + assert_eq!(fixture.receive_data(second, 10), b"lo"); + assert_eq!(fixture.receive(second, 10), Err(BrokerError::WouldBlock)); + assert_eq!(fixture.send(first, b""), Ok(Ok(0))); + + assert_eq!(fixture.send(second, b"back").unwrap(), Ok(4)); + assert_eq!(fixture.receive_data(first, 10), b"back"); + assert_eq!(fixture.queued(), (0, 0)); +} + +#[test] +fn status_flags_belong_to_the_shared_object() { + let fixture = Fixture::new(); + let (first, second) = create_pair( + &fixture.process, + SocketType::Datagram, + FileOpenFlags::NONBLOCKING, + &fixture.readiness_sink, + ) + .unwrap(); + assert_eq!( + fixture.receive(first, 1), + Err(BrokerError::NonBlockingWouldBlock) + ); + assert_eq!( + fixture.process.get_status_flags(second).unwrap(), + FileStatusFlags { + access: FileAccessMode::ReadWrite, + flags: FileOpenFlags::NONBLOCKING, + } + ); + assert_eq!( + create( + &fixture.process, + SocketType::Stream, + FileOpenFlags::CREATE, + &fixture.readiness_sink, + ), + Err(BrokerError::UnsupportedOperation) + ); +} + +#[test] +fn path_listeners_accept_connections_with_names() { + let fixture = Fixture::new(); + let address = path("/server"); + let listener = fixture.listener(&address, 8); + assert_eq!(fixture.readiness(listener), ReadinessFlags::default()); + + let client = fixture.create(SocketType::Stream); + fixture.bind(client, &abstract_name(b"client")).unwrap(); + assert_eq!(fixture.accept(listener), Err(BrokerError::WouldBlock)); + assert_eq!(fixture.connect(client, &address).unwrap(), Ok(())); + assert!(fixture.take_republished().contains(&listener)); + assert!(fixture.readiness(listener).contains(ReadinessFlags::READ)); + assert_eq!( + fixture.connect(client, &address).unwrap(), + Err(LocalSocketError::AlreadyConnected) + ); + + let server = fixture.accept(listener).unwrap().unwrap(); + assert_eq!( + name(&fixture.process, server, false).unwrap(), + Ok(address.name()) + ); + assert_eq!( + name(&fixture.process, server, true).unwrap(), + Ok(LocalSocketName::Abstract(b"client".to_vec())) + ); + assert_eq!( + name(&fixture.process, client, true).unwrap(), + Ok(address.name()) + ); + assert_eq!( + name(&fixture.process, listener, true).unwrap(), + Err(LocalSocketError::NotConnected) + ); + + assert_eq!(fixture.send(client, b"ping").unwrap(), Ok(4)); + let received = fixture.receive(server, 16).unwrap().unwrap(); + assert_eq!(received.data, b"ping"); + assert_eq!( + received.source, + LocalSocketName::Abstract(b"client".to_vec()) + ); + assert_eq!(fixture.send(server, b"pong").unwrap(), Ok(4)); + assert_eq!(fixture.receive_data(client, 16), b"pong"); +} + +#[test] +fn path_names_stay_taken_until_unlinked() { + let fixture = Fixture::new(); + let address = path("/server"); + let listener = fixture.listener(&address, 8); + let other = fixture.create(SocketType::Stream); + assert_eq!( + fixture.bind(other, &address), + Err(LocalSocketError::AddressInUse) + ); + assert_eq!( + fixture.bind(listener, &path("/again")), + Err(LocalSocketError::InvalidArgument) + ); + + fixture.close(listener); + let client = fixture.create(SocketType::Stream); + assert_eq!( + fixture.connect(client, &address).unwrap(), + Err(LocalSocketError::ConnectionRefused) + ); + assert_eq!( + fixture.bind(other, &address), + Err(LocalSocketError::AddressInUse) + ); + crate::fs::unlink(&fixture.process, "/server", USER) + .unwrap() + .unwrap(); + assert_eq!( + fixture.connect(client, &address).unwrap(), + Err(LocalSocketError::File(FileError::NoSuchFileOrDirectory)) + ); + assert_eq!(fixture.bind(other, &address), Ok(())); + assert_eq!( + fixture.connect(client, &path("/")).unwrap(), + Err(LocalSocketError::ConnectionRefused) + ); +} + +#[test] +fn abstract_names_are_released_on_close() { + let fixture = Fixture::new(); + let address = abstract_name(b"name"); + let first = fixture.create(SocketType::Datagram); + let second = fixture.create(SocketType::Datagram); + assert_eq!(fixture.bind(first, &address), Ok(())); + assert_eq!( + fixture.bind(second, &address), + Err(LocalSocketError::AddressInUse) + ); + fixture.close(first); + assert_eq!(fixture.bind(second, &address), Ok(())); +} + +#[test] +fn abstract_names_are_held_per_socket_type() { + let fixture = Fixture::new(); + let address = abstract_name(b"name"); + let datagram = fixture.create(SocketType::Datagram); + assert_eq!(fixture.bind(datagram, &address), Ok(())); + let client = fixture.create(SocketType::Stream); + assert_eq!( + fixture.connect(client, &address).unwrap(), + Err(LocalSocketError::ConnectionRefused) + ); + + let listener = fixture.listener(&address, 1); + assert_eq!(fixture.connect(client, &address).unwrap(), Ok(())); + assert!(fixture.accept(listener).unwrap().is_ok()); + let sender = fixture.create(SocketType::Datagram); + assert_eq!(fixture.send_to(sender, &address, b"x").unwrap(), Ok(1)); + assert_eq!(fixture.receive_data(datagram, 1), b"x"); +} + +#[test] +fn path_connections_require_write_permission() { + let fixture = Fixture::new(); + let address = path("/server"); + let listener = fixture.listener(&address, 1); + let other = FileUser { + user: 1000, + group: 1000, + }; + let client = fixture.create(SocketType::Stream); + assert_eq!( + connect(&fixture.process, client, &address, other).unwrap(), + Err(LocalSocketError::File(FileError::AccessNotAllowed)) + ); + crate::fs::chmod( + &fixture.process, + "/server", + USER, + FileMode::RWXU | FileMode::WOTH, + ) + .unwrap() + .unwrap(); + assert_eq!( + connect(&fixture.process, client, &address, other).unwrap(), + Ok(()) + ); + assert!(fixture.accept(listener).unwrap().is_ok()); + + let (file, _) = crate::fs::open_node( + &fixture.process, + "/file", + USER, + FileAccessMode::WriteOnly, + FileOpenFlags::CREATE, + FileMode::RWXU, + ) + .unwrap() + .unwrap(); + drop(file); + let client = fixture.create(SocketType::Stream); + assert_eq!( + fixture.connect(client, &path("/file")).unwrap(), + Err(LocalSocketError::ConnectionRefused) + ); +} + +#[test] +fn unconnected_streams_reject_data_operations() { + let fixture = Fixture::new(); + let unnamed = fixture.create(SocketType::Stream); + assert_eq!( + listen(&fixture.process, unnamed, 1).unwrap(), + Err(LocalSocketError::InvalidArgument) + ); + let datagram = fixture.create(SocketType::Datagram); + assert_eq!( + listen(&fixture.process, datagram, 1).unwrap(), + Err(LocalSocketError::Unsupported) + ); + let (connected, _) = fixture.pair(SocketType::Stream); + assert_eq!( + listen(&fixture.process, connected, 1).unwrap(), + Err(LocalSocketError::InvalidArgument) + ); + assert_eq!( + fixture.receive(unnamed, 1).unwrap(), + Err(LocalSocketError::InvalidArgument) + ); + assert_eq!( + fixture.send(unnamed, b"x").unwrap(), + Err(LocalSocketError::NotConnected) + ); + assert_eq!( + fixture + .send_to(unnamed, &abstract_name(b"x"), b"x") + .unwrap(), + Err(LocalSocketError::Unsupported) + ); + assert_eq!( + fixture.readiness(unnamed), + ReadinessFlags::WRITE | ReadinessFlags::CLOSED + ); +} + +#[test] +fn full_backlogs_wait_for_accept() { + let fixture = Fixture::new(); + let address = abstract_name(b"server"); + let listener = fixture.listener(&address, 0); + let first = fixture.create(SocketType::Stream); + let second = fixture.create(SocketType::Stream); + assert_eq!(fixture.connect(first, &address).unwrap(), Ok(())); + assert_eq!( + fixture.connect(second, &address), + Err(BrokerError::WouldBlock) + ); + fixture.take_republished(); + + let server = fixture.accept(listener).unwrap().unwrap(); + assert!(fixture.take_republished().contains(&second)); + assert_eq!(fixture.connect(second, &address).unwrap(), Ok(())); + fixture.close(server); + assert_eq!( + fixture.send(first, b"x").unwrap(), + Err(LocalSocketError::BrokenPipe) + ); +} + +#[test] +fn closing_a_listener_resets_unaccepted_connections() { + let fixture = Fixture::new(); + let address = abstract_name(b"server"); + let listener = fixture.listener(&address, 8); + let client = fixture.create(SocketType::Stream); + assert_eq!(fixture.connect(client, &address).unwrap(), Ok(())); + assert_eq!(fixture.send(client, b"lost").unwrap(), Ok(4)); + fixture.close(listener); + assert_eq!(fixture.queued(), (0, 0)); + assert!( + fixture + .readiness(client) + .contains(ReadinessFlags::READ | ReadinessFlags::HANGUP | ReadinessFlags::CLOSED) + ); + assert_eq!(fixture.receive_data(client, 8), b""); + assert_eq!( + fixture.send(client, b"x").unwrap(), + Err(LocalSocketError::BrokenPipe) + ); +} + +#[test] +fn shutdown_reaches_the_stream_peer() { + let fixture = Fixture::new(); + let (first, second) = fixture.pair(SocketType::Stream); + assert_eq!(fixture.send(first, b"tail").unwrap(), Ok(4)); + shutdown(&fixture.process, first, ShutdownMode::Write) + .unwrap() + .unwrap(); + assert!(fixture.take_republished().contains(&second)); + assert_eq!( + fixture.send(first, b"x").unwrap(), + Err(LocalSocketError::BrokenPipe) + ); + assert!( + fixture + .readiness(second) + .contains(ReadinessFlags::READ | ReadinessFlags::HANGUP) + ); + assert_eq!(fixture.receive_data(second, 8), b"tail"); + assert_eq!(fixture.receive_data(second, 8), b""); + assert_eq!(fixture.send(second, b"ok").unwrap(), Ok(2)); + assert_eq!(fixture.receive_data(first, 8), b"ok"); + + shutdown(&fixture.process, first, ShutdownMode::Read) + .unwrap() + .unwrap(); + assert!(fixture.readiness(first).contains(ReadinessFlags::CLOSED)); + assert_eq!( + fixture.send(second, b"x").unwrap(), + Err(LocalSocketError::BrokenPipe) + ); + assert_eq!( + shutdown(&fixture.process, first, ShutdownMode::Abort), + Err(BrokerError::UnsupportedOperation) + ); +} + +#[test] +fn a_read_shut_datagram_socket_ends_only_blocking_receives() { + let fixture = Fixture::new(); + let (first, second) = fixture.pair(SocketType::Datagram); + assert_eq!(fixture.send(second, b"queued").unwrap(), Ok(6)); + shutdown(&fixture.process, first, ShutdownMode::Read) + .unwrap() + .unwrap(); + assert_eq!( + fixture.send(second, b"x").unwrap(), + Err(LocalSocketError::BrokenPipe) + ); + assert_eq!( + receive(&fixture.process, first, 8, false, true) + .unwrap() + .unwrap() + .data, + b"queued" + ); + assert_eq!( + receive(&fixture.process, first, 8, false, true), + Err(BrokerError::NonBlockingWouldBlock) + ); + assert_eq!(fixture.receive(first, 8), Ok(Ok(Received::default()))); + // Shutting down the peer's sending direction leaves this socket open. + shutdown(&fixture.process, first, ShutdownMode::Write) + .unwrap() + .unwrap(); + assert_eq!(fixture.receive(second, 8), Err(BrokerError::WouldBlock)); +} + +#[test] +fn closing_a_stream_hangs_up_its_peer() { + let fixture = Fixture::new(); + let (first, second) = fixture.pair(SocketType::Stream); + assert_eq!(fixture.send(second, b"unread").unwrap(), Ok(6)); + fixture.close(first); + assert!(fixture.take_republished().contains(&second)); + assert_eq!(fixture.queued(), (0, 0)); + assert_eq!( + fixture.readiness(second), + ReadinessFlags::READ + | ReadinessFlags::WRITE + | ReadinessFlags::HANGUP + | ReadinessFlags::CLOSED + ); + assert_eq!(fixture.receive_data(second, 8), b""); + assert_eq!( + fixture.send(second, b"x").unwrap(), + Err(LocalSocketError::BrokenPipe) + ); +} + +#[test] +fn full_streams_wait_for_the_reader() { + let fixture = Fixture::new(); + let (first, second) = fixture.pair(SocketType::Stream); + let data = vec![7; CAPACITY + 1]; + assert_eq!(fixture.send(first, &data).unwrap(), Ok(CAPACITY)); + assert_eq!(fixture.send(first, b"x"), Err(BrokerError::WouldBlock)); + assert!(!fixture.readiness(first).contains(ReadinessFlags::WRITE)); + fixture.take_republished(); + + assert_eq!(fixture.receive_data(second, 1).len(), 1); + assert!(fixture.take_republished().contains(&first)); + assert!(fixture.readiness(first).contains(ReadinessFlags::WRITE)); + assert_eq!(fixture.send(first, &data).unwrap(), Ok(1)); + assert_eq!(fixture.queued(), (CAPACITY, CAPACITY)); + assert_eq!( + fixture.send(first, &vec![0; MAX_LOCAL_SOCKET_TRANSFER_SIZE as usize + 1]), + Err(BrokerError::ResourceExhausted) + ); +} + +#[test] +fn datagrams_keep_boundaries_and_sources() { + let fixture = Fixture::new(); + let receiver_address = abstract_name(b"receiver"); + let receiver = fixture.create(SocketType::Datagram); + fixture.bind(receiver, &receiver_address).unwrap(); + let named = fixture.create(SocketType::Datagram); + fixture.bind(named, &path("/named")).unwrap(); + let unnamed = fixture.create(SocketType::Datagram); + + assert_eq!( + fixture.send_to(named, &receiver_address, b"abc").unwrap(), + Ok(3) + ); + assert_eq!( + fixture.send_to(unnamed, &receiver_address, b"").unwrap(), + Ok(0) + ); + assert_eq!( + fixture + .send_to(unnamed, &receiver_address, b"defgh") + .unwrap(), + Ok(5) + ); + + assert_eq!( + fixture.receive(receiver, 2).unwrap().unwrap(), + Received { + data: b"ab".to_vec(), + length: 3, + source: path("/named").name(), + } + ); + assert_eq!( + fixture.receive(receiver, 8).unwrap().unwrap(), + Received::default() + ); + let peeked = receive(&fixture.process, receiver, 8, true, false) + .unwrap() + .unwrap(); + assert_eq!(peeked.data, b"defgh"); + assert_eq!(fixture.receive_data(receiver, 8), b"defgh"); + assert_eq!(fixture.receive(receiver, 8), Err(BrokerError::WouldBlock)); + assert_eq!(fixture.queued(), (0, 0)); + + assert_eq!( + fixture.send(unnamed, b"x").unwrap(), + Err(LocalSocketError::NotConnected) + ); + assert_eq!( + fixture + .send_to(unnamed, &receiver_address, &vec![0; CAPACITY + 1]) + .unwrap(), + Err(LocalSocketError::MessageTooLarge) + ); +} + +#[test] +fn connected_datagram_sockets_only_accept_their_peer() { + let fixture = Fixture::new(); + let receiver_address = abstract_name(b"receiver"); + let receiver = fixture.create(SocketType::Datagram); + fixture.bind(receiver, &receiver_address).unwrap(); + let peer_address = abstract_name(b"peer"); + let peer = fixture.create(SocketType::Datagram); + fixture.bind(peer, &peer_address).unwrap(); + let other = fixture.create(SocketType::Datagram); + + assert_eq!(fixture.connect(receiver, &peer_address).unwrap(), Ok(())); + assert_eq!( + fixture.send_to(other, &receiver_address, b"x").unwrap(), + Err(LocalSocketError::NotPermitted) + ); + assert_eq!( + fixture.connect(other, &receiver_address).unwrap(), + Err(LocalSocketError::NotPermitted) + ); + assert_eq!( + fixture.send_to(peer, &receiver_address, b"x").unwrap(), + Ok(1) + ); + assert_eq!(fixture.send(receiver, b"y").unwrap(), Ok(1)); + assert_eq!( + fixture.receive(peer, 1).unwrap().unwrap().source, + receiver_address.name() + ); + assert_eq!( + name(&fixture.process, receiver, true).unwrap(), + Ok(peer_address.name()) + ); + + fixture.close(peer); + assert_ne!(fixture.queued(), (0, 0)); + assert_eq!( + fixture.send(receiver, b"z").unwrap(), + Err(LocalSocketError::ConnectionRefused) + ); + assert_eq!(fixture.queued(), (0, 0)); + assert_eq!( + fixture.send(receiver, b"z").unwrap(), + Err(LocalSocketError::NotConnected) + ); +} + +#[test] +fn reconnecting_a_datagram_socket_drops_queued_datagrams() { + let fixture = Fixture::new(); + let receiver_address = abstract_name(b"receiver"); + let receiver = fixture.create(SocketType::Datagram); + fixture.bind(receiver, &receiver_address).unwrap(); + let first_address = abstract_name(b"first"); + let first = fixture.create(SocketType::Datagram); + fixture.bind(first, &first_address).unwrap(); + let second_address = abstract_name(b"second"); + let second = fixture.create(SocketType::Datagram); + fixture.bind(second, &second_address).unwrap(); + + assert_eq!(fixture.connect(receiver, &first_address).unwrap(), Ok(())); + assert_eq!( + fixture.send_to(first, &receiver_address, b"x").unwrap(), + Ok(1) + ); + assert_eq!(fixture.connect(receiver, &first_address).unwrap(), Ok(())); + assert_eq!( + fixture.queued(), + (1 + DATAGRAM_OVERHEAD, 1 + DATAGRAM_OVERHEAD) + ); + + assert_eq!(fixture.connect(receiver, &second_address).unwrap(), Ok(())); + assert_eq!(fixture.queued(), (0, 0)); + assert!(!fixture.readiness(receiver).contains(ReadinessFlags::READ)); + assert_eq!( + receive(&fixture.process, receiver, 1, false, true), + Err(BrokerError::NonBlockingWouldBlock) + ); +} + +#[test] +fn full_datagram_receivers_wake_connected_senders() { + let fixture = Fixture::new(); + let (first, second) = fixture.pair(SocketType::Datagram); + let datagram = vec![0; CAPACITY / 2]; + assert_eq!(fixture.send(first, &datagram).unwrap(), Ok(datagram.len())); + assert_eq!(fixture.send(first, &datagram).unwrap(), Ok(datagram.len())); + assert_eq!(fixture.send(first, b"x"), Err(BrokerError::WouldBlock)); + assert!(!fixture.readiness(first).contains(ReadinessFlags::WRITE)); + fixture.take_republished(); + + assert_eq!(fixture.receive_data(second, 1).len(), 1); + assert!(fixture.take_republished().contains(&first)); + assert!(fixture.readiness(first).contains(ReadinessFlags::WRITE)); + assert_eq!(fixture.send(first, b"x").unwrap(), Ok(1)); +} + +#[test] +fn connecting_a_full_datagram_receiver_wakes_rejected_senders() { + let fixture = Fixture::new(); + let receiver_address = abstract_name(b"receiver"); + let receiver = fixture.create(SocketType::Datagram); + fixture.bind(receiver, &receiver_address).unwrap(); + let peer_address = abstract_name(b"peer"); + let peer = fixture.create(SocketType::Datagram); + fixture.bind(peer, &peer_address).unwrap(); + let sender = fixture.create(SocketType::Datagram); + let datagram = vec![0; CAPACITY / 2]; + for _ in 0..2 { + assert_eq!( + fixture + .send_to(sender, &receiver_address, &datagram) + .unwrap(), + Ok(datagram.len()) + ); + } + assert_eq!( + fixture.send_to(sender, &receiver_address, b"x"), + Err(BrokerError::WouldBlock) + ); + fixture.take_republished(); + + assert_eq!(fixture.connect(receiver, &peer_address).unwrap(), Ok(())); + assert!(fixture.take_republished().contains(&sender)); + assert_eq!( + fixture.send_to(sender, &receiver_address, b"x").unwrap(), + Err(LocalSocketError::NotPermitted) + ); +} + +#[test] +fn queued_bytes_are_limited_and_refunded() { + let fixture = + Fixture::with_limits(BrokerCoreLimits::DEFAULT.with_local_socket_limits(usize::MAX, 1000)); + let (first, second) = fixture.pair(SocketType::Stream); + assert_eq!(fixture.send(first, &[0; 600]).unwrap(), Ok(600)); + assert_eq!( + fixture.send(first, &[0; 600]), + Err(BrokerError::ResourceExhausted) + ); + assert_eq!(fixture.queued(), (600, 600)); + assert_eq!(fixture.receive_data(second, 100).len(), 100); + assert_eq!(fixture.queued(), (500, 500)); + fixture.close(second); + assert_eq!(fixture.queued(), (0, 0)); + fixture.close(first); + + let (first, second) = fixture.pair(SocketType::Datagram); + assert_eq!(fixture.send(first, &[0; 600]).unwrap(), Ok(600)); + assert_eq!( + fixture.queued(), + (600 + DATAGRAM_OVERHEAD, 600 + DATAGRAM_OVERHEAD) + ); + assert_eq!( + fixture.send(first, &[0; 200]), + Err(BrokerError::ResourceExhausted) + ); + fixture.close(first); + fixture.close(second); + assert_eq!(fixture.queued(), (0, 0)); + assert!(fixture.process.core.local_sockets.lock().sockets.is_empty()); +} + +#[test] +fn unaccepted_connections_are_charged_to_the_listener() { + let fixture = Fixture::with_limits( + BrokerCoreLimits::DEFAULT.with_local_socket_limits(usize::MAX, CONNECTION_OVERHEAD), + ); + let address = abstract_name(b"server"); + let listener = fixture.listener(&address, 8); + let first = fixture.create(SocketType::Stream); + assert_eq!(fixture.connect(first, &address).unwrap(), Ok(())); + assert_eq!(fixture.queued(), (CONNECTION_OVERHEAD, CONNECTION_OVERHEAD)); + let second = fixture.create(SocketType::Stream); + assert_eq!( + fixture.connect(second, &address), + Err(BrokerError::ResourceExhausted) + ); + + let accepted = fixture.accept(listener).unwrap().unwrap(); + assert_eq!(fixture.queued(), (0, 0)); + assert_eq!(fixture.connect(second, &address).unwrap(), Ok(())); + fixture.close(listener); + assert_eq!(fixture.queued(), (0, 0)); + fixture.close(accepted); +} + +#[test] +fn duplicated_references_share_one_socket() { + let fixture = Fixture::new(); + let child = fixture + .broker + .create_process( + CallerCredential::Unauthenticated, + Some(fixture.process.id()), + ) + .unwrap(); + let (first, second) = fixture.pair(SocketType::Stream); + let duplicated = fixture + .process + .duplicate_object_reference_to(first, &child, ObjectRights::all()) + .unwrap(); + fixture.close(first); + assert_eq!(fixture.send(second, b"child").unwrap(), Ok(5)); + assert_eq!( + receive(&child, duplicated, 8, false, false) + .unwrap() + .unwrap() + .data, + b"child" + ); + child.close_object_reference(duplicated).unwrap(); + assert_eq!(fixture.receive_data(second, 8), b""); +} + +#[test] +fn options_are_stored_with_the_socket() { + let fixture = Fixture::new(); + let socket = fixture.create(SocketType::Datagram); + set_option( + &fixture.process, + socket, + LocalSocketOption::ReceiveTimeout(Some(Duration::from_secs(2))), + ) + .unwrap(); + let (socket_type, options) = options(&fixture.process, socket).unwrap(); + assert_eq!(socket_type, SocketType::Datagram); + assert_eq!(options.receive_timeout, Some(Duration::from_secs(2))); +} diff --git a/litebox_broker_core/src/object.rs b/litebox_broker_core/src/object.rs index 6b891bdd4e..c4aff3ad09 100644 --- a/litebox_broker_core/src/object.rs +++ b/litebox_broker_core/src/object.rs @@ -12,6 +12,7 @@ use spin::rwlock::RwLock; use crate::event::EventObject; use crate::fs::File; +use crate::local_socket::LocalSocketObject; use crate::pipe::PipeObject; use crate::process::ProcessObject; use crate::readiness::{ReadinessRegistration, ReadinessSink}; @@ -48,6 +49,7 @@ pub(crate) struct ObjectReference { pub(crate) enum ObjectEntry { Event(EventObject), File(File), + LocalSocket(LocalSocketObject), Pipe(PipeObject), Socket(SocketObject), Process(ProcessObject), @@ -66,7 +68,7 @@ impl ObjectEntry { /// other processes must also implement [`Self::watch`]. pub(crate) fn is_duplicable(&self) -> bool { match self { - Self::Event(_) | Self::File(_) | Self::Pipe(_) => true, + Self::Event(_) | Self::File(_) | Self::LocalSocket(_) | Self::Pipe(_) => true, Self::Socket(_) | Self::Process(_) | Self::Signals(_) | Self::Timer(_) => false, } } @@ -81,6 +83,11 @@ impl ObjectEntry { ) -> Result> { match self { Self::File(file) => file.watch(handle, readiness_sink), + Self::LocalSocket(socket) => { + let registration = ReadinessRegistration::new(handle, Arc::clone(readiness_sink)); + socket.watch(®istration)?; + Ok(Some(registration)) + } Self::Pipe(pipe) => { let registration = ReadinessRegistration::new(handle, Arc::clone(readiness_sink)); pipe.watch(®istration)?; @@ -104,6 +111,7 @@ pub(crate) fn readiness(object: &RwLock) -> Result match &*object { ObjectEntry::Event(event) => return Ok(event.readiness()), ObjectEntry::File(file) => return file.readiness(), + ObjectEntry::LocalSocket(socket) => return Ok(socket.readiness()), ObjectEntry::Pipe(pipe) => return Ok(pipe.readiness()), ObjectEntry::Process(process) => return Ok(process.readiness()), ObjectEntry::Signals(signals) => return Ok(signals.readiness()), @@ -121,6 +129,7 @@ pub(crate) fn get_status_flags( ) -> Result { let file = match &*object.read() { ObjectEntry::File(file) => file.clone(), + ObjectEntry::LocalSocket(socket) => return Ok(socket.get_status_flags()), ObjectEntry::Pipe(pipe) => return Ok(pipe.get_status_flags()), ObjectEntry::Event(_) | ObjectEntry::Socket(_) @@ -140,6 +149,10 @@ pub(crate) fn set_status_flags( ) -> Result<()> { let file = match &mut *object.write() { ObjectEntry::File(file) => file.clone(), + ObjectEntry::LocalSocket(socket) => { + socket.set_status_flags(mask, flags); + return Ok(()); + } ObjectEntry::Pipe(pipe) => { pipe.set_status_flags(mask, flags); return Ok(()); diff --git a/litebox_broker_core/src/process.rs b/litebox_broker_core/src/process.rs index b6adc11ba4..9837993e25 100644 --- a/litebox_broker_core/src/process.rs +++ b/litebox_broker_core/src/process.rs @@ -226,6 +226,8 @@ pub struct BrokerProcess { threads: Mutex>, /// Pipe capacity charged to this process by live pipe objects. pub(crate) reserved_pipe_capacity: Arc, + /// Bytes queued in local sockets created by this process. + pub(crate) local_socket_bytes: Arc, /// Socket quota held by pending, live, and closing in-flight resources. pub(crate) reserved_sockets: Arc, /// Signals sent to this process that it has not taken. @@ -370,6 +372,7 @@ impl BrokerProcess { }), threads: Mutex::new(HashSet::new()), reserved_pipe_capacity: Arc::new(AtomicUsize::new(0)), + local_socket_bytes: Arc::new(AtomicUsize::new(0)), reserved_sockets: Arc::new(AtomicUsize::new(0)), signals: Arc::new(ProcessSignals::new()), cancellation: AssociationCancellation::default(), diff --git a/litebox_broker_core/src/socket/tests.rs b/litebox_broker_core/src/socket/tests.rs index bd9b1f2bd2..978ddf99f1 100644 --- a/litebox_broker_core/src/socket/tests.rs +++ b/litebox_broker_core/src/socket/tests.rs @@ -1325,6 +1325,9 @@ fn test_broker_with_policy( .with_socket_policy(*socket_policy), ), limits: crate::BrokerCoreLimits::new_with_all_limits(16, 4, 8, 8), + local_sockets: Arc::new(crate::local_socket::LocalSockets::new( + &crate::BrokerCoreLimits::DEFAULT, + )), ids: Arc::new(spin::Mutex::new( crate::id::IdAllocator::new(crate::id::MAX_ALLOCATED_ID).unwrap(), )), diff --git a/litebox_broker_host/src/lib.rs b/litebox_broker_host/src/lib.rs index 560a457164..6ea5629215 100644 --- a/litebox_broker_host/src/lib.rs +++ b/litebox_broker_host/src/lib.rs @@ -23,6 +23,7 @@ extern crate std; use alloc::{boxed::Box, sync::Arc, vec::Vec}; +use litebox_broker_core::local_socket::LocalSocketResult; use litebox_broker_core::readiness::ReadinessSink; use litebox_broker_core::{ BrokerCore, BrokerError, BrokerProcess, CallerCredential, ChildImage, ProcessImage, @@ -37,10 +38,17 @@ use litebox_broker_protocol::fs::{ SeekFileResponse, TruncateFileRequest, UnlinkFileRequest, WriteFileRequest, WriteFileResponse, encode_directory_entries_chunk, }; +use litebox_broker_protocol::local_socket::{ + AcceptLocalSocketResponse, CreateLocalSocketPairResponse, CreateLocalSocketResponse, + GetLocalSocketNameResponse, GetLocalSocketOptionsResponse, LocalSocketAddress, + LocalSocketError, MAX_ENCODED_LOCAL_SOCKET_ADDRESS_SIZE, MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE, + MAX_LOCAL_SOCKET_TRANSFER_SIZE, ReceiveLocalSocketResponse, SendLocalSocketResponse, +}; use litebox_broker_protocol::message::{ BrokerHandshakeResponse, BrokerOperation, BrokerRequest, BrokerResponse, BrokerResult, - EventRequest, EventResponse, FileRequest, FileResponse, PipeRequest, PipeResponse, - SignalRequest, SignalResponse, SocketRequest, SocketResponse, TimerRequest, TimerResponse, + EventRequest, EventResponse, FileRequest, FileResponse, LocalSocketRequest, + LocalSocketResponse, PipeRequest, PipeResponse, SignalRequest, SignalResponse, SocketRequest, + SocketResponse, TimerRequest, TimerResponse, }; use litebox_broker_protocol::pipe::{ CreatePipeResponse, MAX_PIPE_TRANSFER_SIZE, ReadPipeResponse, WritePipeResponse, @@ -564,6 +572,10 @@ fn handle_request( handle_socket_request(process, request, shared_buffers, readiness_sink) .map(BrokerResult::Socket) } + BrokerOperation::LocalSocket(request) => { + handle_local_socket_request(process, request, shared_buffers, readiness_sink) + .map(BrokerResult::LocalSocket) + } BrokerOperation::FillRandom(buffer) => { validate_shared_buffer(buffer, MAX_RANDOM_TRANSFER_SIZE)?; let length = buffer.length() as usize; @@ -1342,6 +1354,217 @@ fn handle_pipe_request( } } +fn handle_local_socket_request( + process: &BrokerProcess, + request: LocalSocketRequest, + shared_buffers: &SharedBufferPool, + readiness_sink: &Arc, +) -> RequestResult { + Ok( + local_socket_operation(process, request, shared_buffers, readiness_sink)? + .unwrap_or_else(LocalSocketResponse::Failed), + ) +} + +fn local_socket_operation( + process: &BrokerProcess, + request: LocalSocketRequest, + shared_buffers: &SharedBufferPool, + readiness_sink: &Arc, +) -> RequestResult> { + use litebox_broker_core::local_socket; + + match request { + LocalSocketRequest::Create(request) => { + let handle = + local_socket::create(process, request.socket_type, request.flags, readiness_sink)?; + Ok(Ok(LocalSocketResponse::Create(CreateLocalSocketResponse { + handle, + }))) + } + LocalSocketRequest::CreatePair(request) => { + let (first, second) = local_socket::create_pair( + process, + request.socket_type, + request.flags, + readiness_sink, + )?; + Ok(Ok(LocalSocketResponse::CreatePair( + CreateLocalSocketPairResponse { first, second }, + ))) + } + LocalSocketRequest::Bind(request) => { + let address = match read_local_socket_address(shared_buffers, request.address)? { + Ok(address) => address, + Err(error) => return Ok(Err(error)), + }; + Ok(local_socket::bind( + process, + request.handle, + &address, + request.user, + request.mode, + )? + .map(|()| LocalSocketResponse::Bind)) + } + LocalSocketRequest::Listen(request) => { + Ok( + local_socket::listen(process, request.handle, request.backlog)? + .map(|()| LocalSocketResponse::Listen), + ) + } + LocalSocketRequest::Connect(request) => { + let address = match read_local_socket_address(shared_buffers, request.address)? { + Ok(address) => address, + Err(error) => return Ok(Err(error)), + }; + Ok( + local_socket::connect(process, request.handle, &address, request.user)? + .map(|()| LocalSocketResponse::Connect), + ) + } + LocalSocketRequest::Accept(request) => { + Ok( + local_socket::accept(process, request.handle, request.flags, readiness_sink)?.map( + |handle| LocalSocketResponse::Accept(AcceptLocalSocketResponse { handle }), + ), + ) + } + LocalSocketRequest::Send(request) => { + let buffer = read_shared_buffer( + shared_buffers, + request.buffer, + MAX_LOCAL_SOCKET_TRANSFER_SIZE, + )?; + let (address, data) = buffer + .split_at_checked(request.address_length as usize) + .ok_or(RequestFailure::Abort(ErrorCode::MalformedRequest))?; + let address = if address.is_empty() { + None + } else { + match LocalSocketAddress::decode(address) { + Ok(address) => Some(address), + Err(_) => return Ok(Err(LocalSocketError::InvalidArgument)), + } + }; + Ok(local_socket::send( + process, + request.handle, + address.as_ref(), + data, + request.user, + )? + .map(|sent| { + LocalSocketResponse::Send(SendLocalSocketResponse { + sent: u32::try_from(sent).expect("sent bytes fit in the request"), + }) + })) + } + LocalSocketRequest::Receive(request) => { + validate_shared_buffer(request.buffer, MAX_LOCAL_SOCKET_TRANSFER_SIZE)?; + if request + .capacity + .checked_add(MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE) + .is_none_or(|length| length > request.buffer.length()) + { + return Err(RequestFailure::Abort(ErrorCode::MalformedRequest)); + } + let received = match local_socket::receive( + process, + request.handle, + request.capacity, + request.peek, + request.nonblocking, + )? { + Ok(received) => received, + Err(error) => return Ok(Err(error)), + }; + let source = received + .source + .encode() + .ok_or(RequestFailure::Abort(ErrorCode::Internal))?; + let length = u32::try_from(received.length) + .map_err(|_| RequestFailure::Abort(ErrorCode::Internal))?; + let mut output = received.data; + let received = u32::try_from(output.len()).expect("received bytes fit the capacity"); + output + .try_reserve_exact(source.len()) + .map_err(|_| RequestFailure::Respond(ErrorCode::OutOfMemory))?; + output.extend_from_slice(&source); + write_shared_buffer( + shared_buffers, + request.buffer, + &output, + MAX_LOCAL_SOCKET_TRANSFER_SIZE, + )?; + Ok(Ok(LocalSocketResponse::Receive( + ReceiveLocalSocketResponse { + received, + length, + source_length: u32::try_from(source.len()).expect("encoded names fit in u32"), + }, + ))) + } + LocalSocketRequest::Shutdown(request) => { + Ok( + local_socket::shutdown(process, request.handle, request.mode)? + .map(|()| LocalSocketResponse::Shutdown), + ) + } + LocalSocketRequest::GetName(request) => { + validate_shared_buffer(request.buffer, MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE)?; + let name = match local_socket::name(process, request.handle, request.peer)? { + Ok(name) => name, + Err(error) => return Ok(Err(error)), + }; + let encoded = name + .encode() + .ok_or(RequestFailure::Abort(ErrorCode::Internal))?; + if encoded.len() > request.buffer.length() as usize { + return Err(RequestFailure::Abort(ErrorCode::MalformedRequest)); + } + write_shared_buffer( + shared_buffers, + request.buffer, + &encoded, + MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE, + )?; + Ok(Ok(LocalSocketResponse::GetName( + GetLocalSocketNameResponse { + length: u32::try_from(encoded.len()).expect("encoded names fit in u32"), + }, + ))) + } + LocalSocketRequest::SetOption(request) => { + local_socket::set_option(process, request.handle, request.option)?; + Ok(Ok(LocalSocketResponse::SetOption)) + } + LocalSocketRequest::GetOptions(handle) => { + let (socket_type, options) = local_socket::options(process, handle)?; + Ok(Ok(LocalSocketResponse::GetOptions( + GetLocalSocketOptionsResponse { + socket_type, + options, + }, + ))) + } + } +} + +/// Reads an encoded local socket address, reporting an undecodable one as +/// an invalid argument. +fn read_local_socket_address( + shared_buffers: &SharedBufferPool, + buffer: SharedBufferSequence, +) -> RequestResult> { + let encoded = read_shared_buffer( + shared_buffers, + buffer, + MAX_ENCODED_LOCAL_SOCKET_ADDRESS_SIZE, + )?; + Ok(LocalSocketAddress::decode(&encoded).map_err(|_| LocalSocketError::InvalidArgument)) +} + fn handle_event_request( process: &BrokerProcess, request: EventRequest, @@ -1445,6 +1668,11 @@ mod tests { FileUser, OpenFileRequest, ReadDirectoryRequest, ReadFileRequest, SeekFileRequest, SetStatusFlagsRequest, WriteFileRequest, decode_directory_entries, }; + use litebox_broker_protocol::local_socket::{ + AcceptLocalSocketRequest, BindLocalSocketRequest, ConnectLocalSocketRequest, + CreateLocalSocketRequest, GetLocalSocketNameRequest, ListenLocalSocketRequest, + LocalSocketName, ReceiveLocalSocketRequest, SendLocalSocketRequest, + }; use litebox_broker_protocol::message::BrokerHandshakeRequest; use litebox_broker_protocol::pipe::{CreatePipeRequest, ReadPipeRequest, WritePipeRequest}; use litebox_broker_protocol::random::MAX_RANDOM_TRANSFER_SIZE; @@ -1802,6 +2030,7 @@ mod tests { active_requests_operate_timers(&broker, &timer_provider); association_shared_buffer_sequences_stage_pipe_data(&broker); association_shared_buffer_sequences_stage_socket_data(&broker); + association_shared_buffer_sequences_stage_local_socket_data(&broker); association_shared_buffer_sequence_stages_random_data(&broker); active_request_queries_file_terminal(&broker, &stdio_provider); active_requests_change_file_status_flags(&broker); @@ -2660,6 +2889,232 @@ mod tests { assert_eq!(second_slot, [9]); } + fn association_shared_buffer_sequences_stage_local_socket_data(broker: &BrokerCore) { + let process = broker + .create_process(CallerCredential::Unauthenticated, None) + .unwrap(); + let shared_buffers = test_shared_buffers(); + let request = |operation| { + let BrokerResult::LocalSocket(response) = + handle_test_request_with_buffers(&process, operation, &shared_buffers) + else { + panic!("expected a local socket response"); + }; + response + }; + let create = |socket_type| { + let LocalSocketResponse::Create(response) = request(BrokerOperation::LocalSocket( + LocalSocketRequest::Create(CreateLocalSocketRequest { + socket_type, + flags: FileOpenFlags::NONE, + }), + )) else { + panic!("expected local socket creation"); + }; + response.handle + }; + let stage = |slot, data: &[u8]| { + shared_buffers + .write(SharedBufferSlotIndex(slot), data) + .unwrap(); + single_slot_sequence(slot, u32::try_from(data.len()).unwrap()) + }; + let receive = |handle, capacity: u32| { + let response = request(BrokerOperation::LocalSocket(LocalSocketRequest::Receive( + ReceiveLocalSocketRequest { + handle, + buffer: single_slot_sequence(9, capacity + MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE), + capacity, + peek: false, + nonblocking: false, + }, + ))); + let LocalSocketResponse::Receive(response) = response else { + panic!("expected a receive response, got {response:?}"); + }; + let mut output = std::vec![0; (response.received + response.source_length) as usize]; + shared_buffers + .read(SharedBufferSlotIndex(9), &mut output) + .unwrap(); + let source = output.split_off(response.received as usize); + ( + output, + response.length, + LocalSocketName::decode(&source).unwrap(), + ) + }; + + // A stream pair moves staged bytes to its peer. + let LocalSocketResponse::CreatePair(pair) = request(BrokerOperation::LocalSocket( + LocalSocketRequest::CreatePair(CreateLocalSocketRequest { + socket_type: SocketType::Stream, + flags: FileOpenFlags::NONE, + }), + )) else { + panic!("expected local socket pair creation"); + }; + assert_eq!( + request(BrokerOperation::LocalSocket(LocalSocketRequest::Send( + SendLocalSocketRequest { + handle: pair.first, + buffer: stage(2, &[1, 2, 3]), + address_length: 0, + user: ROOT, + }, + ))), + LocalSocketResponse::Send(SendLocalSocketResponse { sent: 3 }) + ); + assert_eq!( + receive(pair.second, 8), + (std::vec![1, 2, 3], 3, LocalSocketName::Unnamed) + ); + + // A path-bound listener accepts a connection made through the same + // path. + let path = LocalSocketAddress::Path { + path: "/server.sock".into(), + name: b"server.sock".to_vec(), + }; + let listener = create(SocketType::Stream); + assert_eq!( + request(BrokerOperation::LocalSocket(LocalSocketRequest::Bind( + BindLocalSocketRequest { + handle: listener, + address: stage(3, &path.encode().unwrap()), + user: ROOT, + mode: FileMode::from_bits(0o755).unwrap(), + }, + ))), + LocalSocketResponse::Bind + ); + assert_eq!( + request(BrokerOperation::LocalSocket(LocalSocketRequest::Listen( + ListenLocalSocketRequest { + handle: listener, + backlog: 1, + }, + ))), + LocalSocketResponse::Listen + ); + let client = create(SocketType::Stream); + assert_eq!( + request(BrokerOperation::LocalSocket(LocalSocketRequest::Connect( + ConnectLocalSocketRequest { + handle: client, + address: stage(3, &path.encode().unwrap()), + user: ROOT, + }, + ))), + LocalSocketResponse::Connect + ); + let LocalSocketResponse::Accept(accepted) = request(BrokerOperation::LocalSocket( + LocalSocketRequest::Accept(AcceptLocalSocketRequest { + handle: listener, + flags: FileOpenFlags::NONE, + }), + )) else { + panic!("expected an accepted connection"); + }; + let get_name = |handle, peer| { + let response = request(BrokerOperation::LocalSocket(LocalSocketRequest::GetName( + GetLocalSocketNameRequest { + handle, + peer, + buffer: single_slot_sequence(10, MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE), + }, + ))); + let LocalSocketResponse::GetName(response) = response else { + panic!("expected a name response, got {response:?}"); + }; + let mut encoded = std::vec![0; response.length as usize]; + shared_buffers + .read(SharedBufferSlotIndex(10), &mut encoded) + .unwrap(); + LocalSocketName::decode(&encoded).unwrap() + }; + assert_eq!(get_name(client, true), path.name()); + assert_eq!(get_name(accepted.handle, false), path.name()); + assert_eq!( + handle_test_request_with_buffers( + &process, + BrokerOperation::File(FileRequest::Unlink(UnlinkFileRequest { + path: stage(3, b"/server.sock"), + user: ROOT, + })), + &shared_buffers, + ), + BrokerResult::File(FileResponse::Unlink) + ); + + // A datagram addressed by name reports the sender's name. + let server = create(SocketType::Datagram); + let sender = create(SocketType::Datagram); + let server_address = LocalSocketAddress::Abstract(b"server".to_vec()); + let sender_address = LocalSocketAddress::Abstract(b"sender".to_vec()); + for (handle, address) in [(server, &server_address), (sender, &sender_address)] { + assert_eq!( + request(BrokerOperation::LocalSocket(LocalSocketRequest::Bind( + BindLocalSocketRequest { + handle, + address: stage(3, &address.encode().unwrap()), + user: ROOT, + mode: FileMode::default(), + }, + ))), + LocalSocketResponse::Bind + ); + } + let mut staged = server_address.encode().unwrap(); + let address_length = u32::try_from(staged.len()).unwrap(); + staged.extend_from_slice(&[7, 8, 9]); + assert_eq!( + request(BrokerOperation::LocalSocket(LocalSocketRequest::Send( + SendLocalSocketRequest { + handle: sender, + buffer: stage(4, &staged), + address_length, + user: ROOT, + }, + ))), + LocalSocketResponse::Send(SendLocalSocketResponse { sent: 3 }) + ); + assert_eq!( + receive(server, 2), + (std::vec![7, 8], 3, sender_address.name()) + ); + + // An undecodable address fails the operation without aborting. + assert_eq!( + request(BrokerOperation::LocalSocket(LocalSocketRequest::Connect( + ConnectLocalSocketRequest { + handle: create(SocketType::Stream), + address: stage(3, &[0xff]), + user: ROOT, + }, + ))), + LocalSocketResponse::Failed(LocalSocketError::InvalidArgument) + ); + + // A receive buffer without room for the source name is malformed. + assert_eq!( + complete_request(handle_request( + &process, + BrokerOperation::LocalSocket(LocalSocketRequest::Receive( + ReceiveLocalSocketRequest { + handle: server, + buffer: single_slot_sequence(9, 8), + capacity: 8, + peek: false, + nonblocking: false, + }, + )), + &shared_buffers, + &test_readiness_sink(), + )), + Err(ErrorCode::MalformedRequest) + ); + } + fn association_shared_buffer_sequences_stage_socket_data(broker: &BrokerCore) { let process = broker .create_process(CallerCredential::Unauthenticated, None) diff --git a/litebox_broker_local/src/lib.rs b/litebox_broker_local/src/lib.rs index e15d7396bf..52750d2397 100644 --- a/litebox_broker_local/src/lib.rs +++ b/litebox_broker_local/src/lib.rs @@ -23,6 +23,7 @@ extern crate std; mod error; mod event; mod fs; +mod local_socket; mod pipe; mod random; mod signal; diff --git a/litebox_broker_local/src/local_socket.rs b/litebox_broker_local/src/local_socket.rs new file mode 100644 index 0000000000..5ae3f9c5bb --- /dev/null +++ b/litebox_broker_local/src/local_socket.rs @@ -0,0 +1,391 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +use alloc::vec::Vec; + +use litebox_broker_protocol::ObjectHandle; +use litebox_broker_protocol::error::ErrorCode; +use litebox_broker_protocol::fs::{FileMode, FileOpenFlags, FileUser}; +use litebox_broker_protocol::local_socket::{ + AcceptLocalSocketRequest, BindLocalSocketRequest, ConnectLocalSocketRequest, + CreateLocalSocketPairResponse, CreateLocalSocketRequest, GetLocalSocketNameRequest, + GetLocalSocketOptionsResponse, ListenLocalSocketRequest, LocalSocketError, LocalSocketName, + LocalSocketOption, MAX_ENCODED_LOCAL_SOCKET_ADDRESS_SIZE, MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE, + MAX_LOCAL_SOCKET_TRANSFER_SIZE, ReceiveLocalSocketRequest, ReceiveLocalSocketResponse, + SendLocalSocketRequest, SetLocalSocketOptionRequest, ShutdownLocalSocketRequest, +}; +use litebox_broker_protocol::message::{ + BrokerOperation, BrokerResult, LocalSocketRequest, LocalSocketResponse, +}; +use litebox_broker_protocol::shared_buffer::SharedBufferSequence; +use litebox_broker_protocol::socket::{ShutdownMode, SocketType}; +use litebox_broker_transport::channel::LocalCallChannel; + +use crate::{BrokerLocal, BrokerLocalError, Result}; + +type LocalSocketResult = core::result::Result; + +impl BrokerLocal { + /// Creates a broker-owned local socket. + /// + /// # Panics + /// + /// Panics if the broker returns a response for a different operation. + pub fn create_local_socket( + &self, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> Result { + match self.request_local_socket(LocalSocketRequest::Create(CreateLocalSocketRequest { + socket_type, + flags, + }))? { + LocalSocketResponse::Create(response) => Ok(response.handle), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Creates a pair of connected broker-owned local sockets. + /// + /// # Panics + /// + /// Panics if the broker returns a response for a different operation. + pub fn create_local_socket_pair( + &self, + socket_type: SocketType, + flags: FileOpenFlags, + ) -> Result { + match self.request_local_socket(LocalSocketRequest::CreatePair( + CreateLocalSocketRequest { socket_type, flags }, + ))? { + LocalSocketResponse::CreatePair(response) => Ok(response), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Binds a local socket to an encoded address staged in `buffer`. + /// + /// # Panics + /// + /// Panics if `buffer` does not match `address` or the broker returns a + /// response for a different operation. + pub fn bind_local_socket( + &self, + handle: ObjectHandle, + buffer: SharedBufferSequence, + address: &[u8], + user: FileUser, + mode: FileMode, + ) -> Result, Channel::Error> { + self.stage_local_socket_buffer(buffer, address, MAX_ENCODED_LOCAL_SOCKET_ADDRESS_SIZE)?; + match self.request_local_socket(LocalSocketRequest::Bind(BindLocalSocketRequest { + handle, + address: buffer, + user, + mode, + }))? { + LocalSocketResponse::Bind => Ok(Ok(())), + LocalSocketResponse::Failed(error) => Ok(Err(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Makes a bound stream socket accept connections. + /// + /// # Panics + /// + /// Panics if the broker returns a response for a different operation. + pub fn listen_local_socket( + &self, + handle: ObjectHandle, + backlog: u32, + ) -> Result, Channel::Error> { + match self.request_local_socket(LocalSocketRequest::Listen(ListenLocalSocketRequest { + handle, + backlog, + }))? { + LocalSocketResponse::Listen => Ok(Ok(())), + LocalSocketResponse::Failed(error) => Ok(Err(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Connects a local socket to an encoded address staged in `buffer`. + /// + /// # Panics + /// + /// Panics if `buffer` does not match `address` or the broker returns a + /// response for a different operation. + pub fn connect_local_socket( + &self, + handle: ObjectHandle, + buffer: SharedBufferSequence, + address: &[u8], + user: FileUser, + ) -> Result, Channel::Error> { + self.stage_local_socket_buffer(buffer, address, MAX_ENCODED_LOCAL_SOCKET_ADDRESS_SIZE)?; + match self.request_local_socket(LocalSocketRequest::Connect(ConnectLocalSocketRequest { + handle, + address: buffer, + user, + }))? { + LocalSocketResponse::Connect => Ok(Ok(())), + LocalSocketResponse::Failed(error) => Ok(Err(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Accepts one pending connection from a listening local socket. + /// + /// # Panics + /// + /// Panics if the broker returns a response for a different operation. + pub fn accept_local_socket( + &self, + handle: ObjectHandle, + flags: FileOpenFlags, + ) -> Result, Channel::Error> { + match self.request_local_socket(LocalSocketRequest::Accept(AcceptLocalSocketRequest { + handle, + flags, + }))? { + LocalSocketResponse::Accept(response) => Ok(Ok(response.handle)), + LocalSocketResponse::Failed(error) => Ok(Err(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Sends `staged`, an encoded destination address of `address_length` + /// bytes followed by the data, from a local socket. + /// + /// Returns the number of data bytes sent. + /// + /// # Panics + /// + /// Panics if `buffer` does not match `staged`, `address_length` exceeds + /// it, or the broker returns an invalid response. + pub fn send_local_socket( + &self, + handle: ObjectHandle, + buffer: SharedBufferSequence, + staged: &[u8], + address_length: usize, + user: FileUser, + ) -> Result, Channel::Error> { + let data_length = staged + .len() + .checked_sub(address_length) + .expect("staged send must hold its address"); + self.stage_local_socket_buffer(buffer, staged, MAX_LOCAL_SOCKET_TRANSFER_SIZE)?; + match self.request_local_socket(LocalSocketRequest::Send(SendLocalSocketRequest { + handle, + buffer, + address_length: u32::try_from(address_length) + .expect("staged address fits the transfer limit"), + user, + }))? { + LocalSocketResponse::Send(response) => { + let sent = response.sent as usize; + assert!( + sent <= data_length && (sent != 0 || data_length == 0), + "broker returned an invalid local socket send length" + ); + Ok(Ok(sent)) + } + LocalSocketResponse::Failed(error) => Ok(Err(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Receives data from a local socket into `destination` through + /// `buffer`, which must hold `destination` plus an encoded name, as for a + /// non-blocking socket if `nonblocking` is set. + /// + /// Returns the received lengths and the sender's name. + /// + /// # Panics + /// + /// Panics if `buffer` is not `destination.len()` plus + /// [`MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE`] bytes, or the broker returns an + /// invalid response. + pub fn receive_local_socket( + &self, + handle: ObjectHandle, + buffer: SharedBufferSequence, + destination: &mut [u8], + peek: bool, + nonblocking: bool, + ) -> Result, Channel::Error> + { + let capacity = destination.len(); + self.validate_local_socket_buffer( + buffer, + capacity + MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE as usize, + MAX_LOCAL_SOCKET_TRANSFER_SIZE, + )?; + match self.request_local_socket(LocalSocketRequest::Receive(ReceiveLocalSocketRequest { + handle, + buffer, + capacity: u32::try_from(capacity).expect("validated capacity fits in u32"), + peek, + nonblocking, + }))? { + LocalSocketResponse::Receive(response) => { + let received = response.received as usize; + let source_length = response.source_length as usize; + assert!( + received <= capacity + && received <= response.length as usize + && source_length <= MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE as usize, + "broker returned inconsistent local socket receive lengths" + ); + let mut staged = Vec::new(); + staged + .try_reserve_exact(received + source_length) + .map_err(|_| BrokerLocalError::Broker(ErrorCode::OutOfMemory))?; + staged.resize(received + source_length, 0); + self.read_shared_buffer(buffer, &mut staged); + let (data, source) = staged.split_at(received); + destination[..received].copy_from_slice(data); + let source = LocalSocketName::decode(source) + .expect("broker returned an invalid local socket source name"); + Ok(Ok((response, source))) + } + LocalSocketResponse::Failed(error) => Ok(Err(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Shuts down one or both directions of a local socket. + /// + /// # Panics + /// + /// Panics if the broker returns a response for a different operation. + pub fn shutdown_local_socket( + &self, + handle: ObjectHandle, + mode: ShutdownMode, + ) -> Result, Channel::Error> { + match self.request_local_socket(LocalSocketRequest::Shutdown( + ShutdownLocalSocketRequest { handle, mode }, + ))? { + LocalSocketResponse::Shutdown => Ok(Ok(())), + LocalSocketResponse::Failed(error) => Ok(Err(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Reads the name of a local socket, or of its peer if `peer` is set, + /// through `buffer` of [`MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE`] bytes. + /// + /// # Panics + /// + /// Panics if `buffer` has the wrong length or the broker returns an + /// invalid response. + pub fn local_socket_name( + &self, + handle: ObjectHandle, + peer: bool, + buffer: SharedBufferSequence, + ) -> Result, Channel::Error> { + self.validate_local_socket_buffer( + buffer, + MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE as usize, + MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE, + )?; + match self.request_local_socket(LocalSocketRequest::GetName(GetLocalSocketNameRequest { + handle, + peer, + buffer, + }))? { + LocalSocketResponse::GetName(response) => { + let length = response.length as usize; + assert!( + length <= MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE as usize, + "broker returned an oversized local socket name" + ); + let mut encoded = [0; MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE as usize]; + self.read_shared_buffer(buffer, &mut encoded[..length]); + Ok(Ok(LocalSocketName::decode(&encoded[..length]) + .expect("broker returned an invalid local socket name"))) + } + LocalSocketResponse::Failed(error) => Ok(Err(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Stores one option of a local socket. + /// + /// # Panics + /// + /// Panics if the broker returns a response for a different operation. + pub fn set_local_socket_option( + &self, + handle: ObjectHandle, + option: LocalSocketOption, + ) -> Result<(), Channel::Error> { + match self.request_local_socket(LocalSocketRequest::SetOption( + SetLocalSocketOptionRequest { handle, option }, + ))? { + LocalSocketResponse::SetOption => Ok(()), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + /// Reads the type and stored options of a local socket. + /// + /// # Panics + /// + /// Panics if the broker returns a response for a different operation. + pub fn local_socket_options( + &self, + handle: ObjectHandle, + ) -> Result { + match self.request_local_socket(LocalSocketRequest::GetOptions(handle))? { + LocalSocketResponse::GetOptions(response) => Ok(response), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } + + fn stage_local_socket_buffer( + &self, + buffer: SharedBufferSequence, + data: &[u8], + max_length: u32, + ) -> Result<(), Channel::Error> { + self.validate_local_socket_buffer(buffer, data.len(), max_length)?; + self.write_shared_buffer(buffer, data); + Ok(()) + } + + fn validate_local_socket_buffer( + &self, + buffer: SharedBufferSequence, + expected_length: usize, + max_length: u32, + ) -> Result<(), Channel::Error> { + if buffer.length() > max_length { + return Err(BrokerLocalError::Broker(ErrorCode::ResourceExhausted)); + } + assert_eq!( + expected_length, + buffer.length() as usize, + "shared data must match its buffer sequence" + ); + let _ = buffer + .descriptors(self.shared_buffers.layout()) + .expect("shared buffer sequence must identify valid slot ranges"); + Ok(()) + } + + fn request_local_socket( + &self, + request: LocalSocketRequest, + ) -> Result { + match self.request(BrokerOperation::LocalSocket(request))? { + BrokerResult::LocalSocket(response) => Ok(response), + BrokerResult::Error(error) => Err(BrokerLocalError::Broker(error)), + response => panic!("broker returned unexpected local socket response: {response:?}"), + } + } +} diff --git a/litebox_broker_protocol/src/lib.rs b/litebox_broker_protocol/src/lib.rs index f110205c65..322ac7dc9b 100644 --- a/litebox_broker_protocol/src/lib.rs +++ b/litebox_broker_protocol/src/lib.rs @@ -19,6 +19,7 @@ extern crate std; pub mod error; pub mod event; pub mod fs; +pub mod local_socket; pub mod message; pub mod pipe; pub mod process; diff --git a/litebox_broker_protocol/src/local_socket.rs b/litebox_broker_protocol/src/local_socket.rs new file mode 100644 index 0000000000..743d953791 --- /dev/null +++ b/litebox_broker_protocol/src/local_socket.rs @@ -0,0 +1,532 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +//! Broker-owned local sockets. +//! +//! A local socket connects processes on the same broker without host +//! networking. Its names live in either a filesystem path, identified by the +//! node the path resolves to, or an abstract namespace of byte strings, as +//! Unix domain sockets name them on POSIX systems and Windows. Data moves +//! through broker-owned queues, so every reference to a socket, including +//! references in other processes, shares its state. + +use alloc::string::String; +use alloc::vec::Vec; +use core::time::Duration; + +use thiserror::Error; + +use crate::ObjectHandle; +use crate::fs::{FileError, FileMode, FileOpenFlags, FileUser}; +use crate::shared_buffer::{SHARED_BUFFER_SLOT_SIZE, SharedBufferSequence}; +use crate::socket::{ShutdownMode, SocketType}; + +/// Maximum length in bytes of the guest-visible part of a local socket name. +pub const MAX_LOCAL_SOCKET_NAME_SIZE: u32 = 108; + +/// Maximum length in bytes of an encoded [`LocalSocketName`]. +pub const MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE: u32 = 1 + MAX_LOCAL_SOCKET_NAME_SIZE; + +/// Maximum length in bytes of the absolute lookup path of a path name. +pub const MAX_LOCAL_SOCKET_PATH_SIZE: u32 = 4096; + +/// Maximum length in bytes of an encoded [`LocalSocketAddress`]. +pub const MAX_ENCODED_LOCAL_SOCKET_ADDRESS_SIZE: u32 = + 3 + MAX_LOCAL_SOCKET_PATH_SIZE + MAX_LOCAL_SOCKET_NAME_SIZE; + +/// Bytes each direction of a connection may queue, which also bounds one +/// datagram. +pub const LOCAL_SOCKET_BUFFER_SIZE: u32 = 212_992; + +/// Maximum bytes one send or receive request transfers, including an encoded +/// address or name staged with the data. +pub const MAX_LOCAL_SOCKET_TRANSFER_SIZE: u32 = 4 * SHARED_BUFFER_SLOT_SIZE; + +/// Largest accepted listen backlog; larger requests are clamped. +pub const MAX_LOCAL_SOCKET_BACKLOG: u32 = 4096; + +const _: () = assert!( + LOCAL_SOCKET_BUFFER_SIZE + MAX_ENCODED_LOCAL_SOCKET_ADDRESS_SIZE + <= MAX_LOCAL_SOCKET_TRANSFER_SIZE +); +const _: () = assert!( + MAX_LOCAL_SOCKET_TRANSFER_SIZE.div_ceil(SHARED_BUFFER_SLOT_SIZE) as usize + <= crate::shared_buffer::MAX_SHARED_BUFFER_SEQUENCE_SLOTS +); +const _: () = assert!(MAX_LOCAL_SOCKET_PATH_SIZE <= u16::MAX as u32); + +const NAME_TAG_UNNAMED: u8 = 0; +const NAME_TAG_PATH: u8 = 1; +const NAME_TAG_ABSTRACT: u8 = 2; + +/// The name a local socket reports for itself or its peer. +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub enum LocalSocketName { + /// The socket has no name. + #[default] + Unnamed, + /// A filesystem name, holding the path bytes the socket was bound with. + Path(Vec), + /// A name in the abstract namespace. + Abstract(Vec), +} + +impl LocalSocketName { + /// Encodes this name for a shared buffer. + /// + /// Returns `None` if the name is longer than [`MAX_LOCAL_SOCKET_NAME_SIZE`]. + #[must_use] + pub fn encode(&self) -> Option> { + let (tag, bytes): (u8, &[u8]) = match self { + Self::Unnamed => (NAME_TAG_UNNAMED, &[]), + Self::Path(bytes) => (NAME_TAG_PATH, bytes), + Self::Abstract(bytes) => (NAME_TAG_ABSTRACT, bytes), + }; + if bytes.len() > MAX_LOCAL_SOCKET_NAME_SIZE as usize { + return None; + } + let mut encoded = Vec::with_capacity(1 + bytes.len()); + encoded.push(tag); + encoded.extend_from_slice(bytes); + Some(encoded) + } + + /// Decodes a name produced by [`Self::encode`]. + pub fn decode(encoded: &[u8]) -> Result { + let (&tag, bytes) = encoded + .split_first() + .ok_or(LocalSocketCodecError::Truncated)?; + if bytes.len() > MAX_LOCAL_SOCKET_NAME_SIZE as usize { + return Err(LocalSocketCodecError::TooLong); + } + match tag { + NAME_TAG_UNNAMED if bytes.is_empty() => Ok(Self::Unnamed), + NAME_TAG_PATH => Ok(Self::Path(bytes.into())), + NAME_TAG_ABSTRACT => Ok(Self::Abstract(bytes.into())), + _ => Err(LocalSocketCodecError::Invalid), + } + } +} + +/// An address that binds, connects, or sends to a local socket. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum LocalSocketAddress { + /// A filesystem name. + Path { + /// Absolute UTF-8 path the broker resolves. + path: String, + /// Guest-visible name that binding records. + name: Vec, + }, + /// A name in the abstract namespace. + Abstract(Vec), +} + +impl LocalSocketAddress { + /// Returns the name a socket bound to this address reports. + #[must_use] + pub fn name(&self) -> LocalSocketName { + match self { + Self::Path { name, .. } => LocalSocketName::Path(name.clone()), + Self::Abstract(bytes) => LocalSocketName::Abstract(bytes.clone()), + } + } + + /// Encodes this address for a shared buffer. + /// + /// Returns `None` if the path or name is too long. + #[must_use] + pub fn encode(&self) -> Option> { + match self { + Self::Path { path, name } => { + if path.len() > MAX_LOCAL_SOCKET_PATH_SIZE as usize + || name.len() > MAX_LOCAL_SOCKET_NAME_SIZE as usize + { + return None; + } + let path_length = u16::try_from(path.len()).ok()?; + let mut encoded = Vec::with_capacity(3 + path.len() + name.len()); + encoded.push(NAME_TAG_PATH); + encoded.extend_from_slice(&path_length.to_le_bytes()); + encoded.extend_from_slice(path.as_bytes()); + encoded.extend_from_slice(name); + Some(encoded) + } + Self::Abstract(bytes) => LocalSocketName::Abstract(bytes.clone()).encode(), + } + } + + /// Decodes an address produced by [`Self::encode`]. + pub fn decode(encoded: &[u8]) -> Result { + match encoded.split_first() { + Some((&NAME_TAG_PATH, rest)) => { + let (length, rest) = rest + .split_first_chunk::<2>() + .ok_or(LocalSocketCodecError::Truncated)?; + let length = usize::from(u16::from_le_bytes(*length)); + if length > MAX_LOCAL_SOCKET_PATH_SIZE as usize { + return Err(LocalSocketCodecError::TooLong); + } + let (path, name) = rest + .split_at_checked(length) + .ok_or(LocalSocketCodecError::Truncated)?; + if name.len() > MAX_LOCAL_SOCKET_NAME_SIZE as usize { + return Err(LocalSocketCodecError::TooLong); + } + let path = + core::str::from_utf8(path).map_err(|_| LocalSocketCodecError::Invalid)?; + if !path.starts_with('/') { + return Err(LocalSocketCodecError::Invalid); + } + Ok(Self::Path { + path: path.into(), + name: name.into(), + }) + } + Some((&NAME_TAG_ABSTRACT, _)) => match LocalSocketName::decode(encoded)? { + LocalSocketName::Abstract(bytes) => Ok(Self::Abstract(bytes)), + _ => Err(LocalSocketCodecError::Invalid), + }, + Some(_) => Err(LocalSocketCodecError::Invalid), + None => Err(LocalSocketCodecError::Truncated), + } + } +} + +/// Failure to decode a local socket name or address from a shared buffer. +#[derive(Clone, Copy, Debug, Error, PartialEq, Eq)] +pub enum LocalSocketCodecError { + #[error("truncated local socket name")] + Truncated, + #[error("local socket name is too long")] + TooLong, + #[error("invalid local socket name")] + Invalid, +} + +/// Local socket operation failure that is meaningful to the guest ABI. +#[derive(Clone, Copy, Debug, Error, PartialEq, Eq)] +#[non_exhaustive] +pub enum LocalSocketError { + #[error("address is already in use")] + AddressInUse, + #[error("no socket accepts connections at the address")] + ConnectionRefused, + #[error("socket at the address has a different type")] + WrongType, + #[error("invalid argument for the socket's state")] + InvalidArgument, + #[error("socket is already connected")] + AlreadyConnected, + #[error("socket is not connected")] + NotConnected, + #[error("operation is not supported by the socket type")] + Unsupported, + #[error("message is too large")] + MessageTooLarge, + #[error("sending side is shut down")] + BrokenPipe, + #[error("receiving socket is connected to another socket")] + NotPermitted, + #[error("address lookup failed: {0}")] + File(FileError), +} + +/// Options stored with a local socket and shared by its references. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub struct LocalSocketOptions { + /// Longest time a blocking receive or accept waits, or `None` to wait + /// indefinitely. + pub receive_timeout: Option, + /// Longest time a blocking send or connect waits, or `None` to wait + /// indefinitely. + pub send_timeout: Option, + /// Linger time on close, or `None` to close in the background. + pub linger: Option, + /// Whether address reuse was requested. + pub reuse_address: bool, + /// Whether keepalive was requested. + pub keep_alive: bool, + /// Whether broadcast was requested. + pub broadcast: bool, +} + +/// One local socket option value. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum LocalSocketOption { + /// [`LocalSocketOptions::receive_timeout`]. + ReceiveTimeout(Option), + /// [`LocalSocketOptions::send_timeout`]. + SendTimeout(Option), + /// [`LocalSocketOptions::linger`]. + Linger(Option), + /// [`LocalSocketOptions::reuse_address`]. + ReuseAddress(bool), + /// [`LocalSocketOptions::keep_alive`]. + KeepAlive(bool), + /// [`LocalSocketOptions::broadcast`]. + Broadcast(bool), +} + +impl LocalSocketOptions { + /// Stores `option`. + pub fn set(&mut self, option: LocalSocketOption) { + match option { + LocalSocketOption::ReceiveTimeout(timeout) => self.receive_timeout = timeout, + LocalSocketOption::SendTimeout(timeout) => self.send_timeout = timeout, + LocalSocketOption::Linger(timeout) => self.linger = timeout, + LocalSocketOption::ReuseAddress(value) => self.reuse_address = value, + LocalSocketOption::KeepAlive(value) => self.keep_alive = value, + LocalSocketOption::Broadcast(value) => self.broadcast = value, + } + } +} + +/// Request to create one local socket or a connected pair. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct CreateLocalSocketRequest { + /// Socket type. + pub socket_type: SocketType, + /// Initial status flags, within [`FileOpenFlags::STATUS`]. + pub flags: FileOpenFlags, +} + +/// Response to a local socket create request. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct CreateLocalSocketResponse { + /// Handle for the new socket. + pub handle: ObjectHandle, +} + +/// Response to a local socket pair create request. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct CreateLocalSocketPairResponse { + /// Handle for the first socket. + pub first: ObjectHandle, + /// Handle for the second socket, connected to the first. + pub second: ObjectHandle, +} + +/// Request to bind a local socket to a name. +/// +/// A path name creates a filesystem node at the path, which fails if one +/// exists. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct BindLocalSocketRequest { + /// Socket handle. + pub handle: ObjectHandle, + /// Shared-buffer region holding one encoded [`LocalSocketAddress`]. + pub address: SharedBufferSequence, + /// Caller identity for filesystem permission checks. + pub user: FileUser, + /// Mode of the filesystem node a path name creates. + pub mode: FileMode, +} + +/// Request to make a bound stream socket accept connections. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ListenLocalSocketRequest { + /// Socket handle. + pub handle: ObjectHandle, + /// Connections that may wait to be accepted beyond the first. + pub backlog: u32, +} + +/// Request to connect a stream socket, or to set a datagram socket's default +/// destination. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ConnectLocalSocketRequest { + /// Socket handle. + pub handle: ObjectHandle, + /// Shared-buffer region holding one encoded [`LocalSocketAddress`]. + pub address: SharedBufferSequence, + /// Caller identity for filesystem permission checks. + pub user: FileUser, +} + +/// Request to accept one pending connection. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct AcceptLocalSocketRequest { + /// Listening socket handle. + pub handle: ObjectHandle, + /// Status flags of the accepted socket, within [`FileOpenFlags::STATUS`]. + pub flags: FileOpenFlags, +} + +/// Response to an accept request. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct AcceptLocalSocketResponse { + /// Handle for the accepted connection. + pub handle: ObjectHandle, +} + +/// Request to send bytes staged in shared memory. +/// +/// A stream socket may send only part of the bytes. A datagram socket sends +/// them as one datagram. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct SendLocalSocketRequest { + /// Socket handle. + pub handle: ObjectHandle, + /// Shared-buffer region holding an encoded destination + /// [`LocalSocketAddress`] of `address_length` bytes followed by the data. + pub buffer: SharedBufferSequence, + /// Length of the destination address, or zero to send to the connected + /// peer. + pub address_length: u32, + /// Caller identity for filesystem permission checks. + pub user: FileUser, +} + +/// Response to a send request. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct SendLocalSocketResponse { + /// Number of data bytes sent. + pub sent: u32, +} + +/// Request to receive bytes into shared memory. +/// +/// A stream socket receives queued bytes. A datagram socket receives one +/// datagram, discarding the bytes beyond `capacity` unless peeking. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ReceiveLocalSocketRequest { + /// Socket handle. + pub handle: ObjectHandle, + /// Shared-buffer region that receives the data followed by the sender's + /// encoded [`LocalSocketName`]. + /// + /// It must hold `capacity` bytes plus + /// [`MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE`]. + pub buffer: SharedBufferSequence, + /// Maximum data bytes to receive. + pub capacity: u32, + /// Whether to leave the data queued. + pub peek: bool, + /// Whether the caller will not wait for data, as if the socket were + /// non-blocking. + /// + /// A datagram socket whose receive direction is shut down then reports + /// that the receive would block instead of the end of data. + pub nonblocking: bool, +} + +/// Response to a receive request. +/// +/// A response with zero `received` and `length` for a nonzero `capacity` +/// means no more data can arrive. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ReceiveLocalSocketResponse { + /// Number of data bytes placed in the buffer. + pub received: u32, + /// Length of the whole datagram, or `received` for a stream socket. + pub length: u32, + /// Length of the sender's encoded name, which follows the data. + pub source_length: u32, +} + +/// Request to shut down one or both directions of a socket. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ShutdownLocalSocketRequest { + /// Socket handle. + pub handle: ObjectHandle, + /// [`ShutdownMode::Read`], [`ShutdownMode::Write`], or [`ShutdownMode::Both`]. + pub mode: ShutdownMode, +} + +/// Request to read the name of a socket or its peer. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct GetLocalSocketNameRequest { + /// Socket handle. + pub handle: ObjectHandle, + /// Whether to read the peer's name instead of the socket's own. + pub peer: bool, + /// Shared-buffer region of [`MAX_ENCODED_LOCAL_SOCKET_NAME_SIZE`] bytes + /// that receives the encoded [`LocalSocketName`]. + pub buffer: SharedBufferSequence, +} + +/// Response to a name request. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct GetLocalSocketNameResponse { + /// Length of the encoded name. + pub length: u32, +} + +/// Request to store one socket option. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct SetLocalSocketOptionRequest { + /// Socket handle. + pub handle: ObjectHandle, + /// Option value. + pub option: LocalSocketOption, +} + +/// Response describing a socket's type and options. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct GetLocalSocketOptionsResponse { + /// Socket type. + pub socket_type: SocketType, + /// Stored options. + pub options: LocalSocketOptions, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn names_round_trip() { + for name in [ + LocalSocketName::Unnamed, + LocalSocketName::Path(b"./sock".into()), + LocalSocketName::Abstract(b"\0x".into()), + LocalSocketName::Abstract(Vec::new()), + ] { + assert_eq!(LocalSocketName::decode(&name.encode().unwrap()), Ok(name)); + } + assert_eq!( + LocalSocketName::Path(alloc::vec![1; MAX_LOCAL_SOCKET_NAME_SIZE as usize + 1]).encode(), + None + ); + assert_eq!( + LocalSocketName::decode(&[NAME_TAG_UNNAMED, 1]), + Err(LocalSocketCodecError::Invalid) + ); + assert_eq!( + LocalSocketName::decode(&[]), + Err(LocalSocketCodecError::Truncated) + ); + } + + #[test] + fn addresses_round_trip() { + for address in [ + LocalSocketAddress::Path { + path: "/tmp/sock".into(), + name: b"sock".into(), + }, + LocalSocketAddress::Abstract(b"name".into()), + ] { + let encoded = address.encode().unwrap(); + assert!(encoded.len() <= MAX_ENCODED_LOCAL_SOCKET_ADDRESS_SIZE as usize); + assert_eq!(LocalSocketAddress::decode(&encoded), Ok(address)); + } + let relative = LocalSocketAddress::Path { + path: "tmp/sock".into(), + name: Vec::new(), + }; + assert_eq!( + LocalSocketAddress::decode(&relative.encode().unwrap()), + Err(LocalSocketCodecError::Invalid) + ); + assert_eq!( + LocalSocketAddress::decode(&[NAME_TAG_UNNAMED]), + Err(LocalSocketCodecError::Invalid) + ); + assert_eq!( + LocalSocketAddress::decode(&[NAME_TAG_PATH, 5, 0, b'/']), + Err(LocalSocketCodecError::Truncated) + ); + } +} diff --git a/litebox_broker_protocol/src/message.rs b/litebox_broker_protocol/src/message.rs index c3a8054aba..c3cd4398e5 100644 --- a/litebox_broker_protocol/src/message.rs +++ b/litebox_broker_protocol/src/message.rs @@ -14,6 +14,14 @@ use crate::fs::{ SetStatusFlagsRequest, TruncateFileRequest, UnlinkFileRequest, WriteFileRequest, WriteFileResponse, }; +use crate::local_socket::{ + AcceptLocalSocketRequest, AcceptLocalSocketResponse, BindLocalSocketRequest, + ConnectLocalSocketRequest, CreateLocalSocketPairResponse, CreateLocalSocketRequest, + CreateLocalSocketResponse, GetLocalSocketNameRequest, GetLocalSocketNameResponse, + GetLocalSocketOptionsResponse, ListenLocalSocketRequest, LocalSocketError, + ReceiveLocalSocketRequest, ReceiveLocalSocketResponse, SendLocalSocketRequest, + SendLocalSocketResponse, SetLocalSocketOptionRequest, ShutdownLocalSocketRequest, +}; use crate::pipe::{ CreatePipeRequest, CreatePipeResponse, ReadPipeRequest, ReadPipeResponse, WritePipeRequest, WritePipeResponse, @@ -101,6 +109,8 @@ pub enum BrokerOperation { Timer(TimerRequest), /// Signal request family. Signal(SignalRequest), + /// Local socket request family. + LocalSocket(LocalSocketRequest), } impl BrokerOperation { @@ -139,7 +149,18 @@ impl BrokerOperation { handles: buffer, .. }) - | Self::WriteChildMemory(WriteChildMemoryRequest { data: buffer, .. }) => Some(*buffer), + | Self::WriteChildMemory(WriteChildMemoryRequest { data: buffer, .. }) + | Self::LocalSocket( + LocalSocketRequest::Bind(BindLocalSocketRequest { + address: buffer, .. + }) + | LocalSocketRequest::Connect(ConnectLocalSocketRequest { + address: buffer, .. + }) + | LocalSocketRequest::Send(SendLocalSocketRequest { buffer, .. }) + | LocalSocketRequest::Receive(ReceiveLocalSocketRequest { buffer, .. }) + | LocalSocketRequest::GetName(GetLocalSocketNameRequest { buffer, .. }), + ) => Some(*buffer), Self::CreateThread(_) | Self::ExitThread(_) | Self::CloseObject(_) @@ -170,6 +191,15 @@ impl BrokerOperation { | FileRequest::Truncate(_) | FileRequest::HandleStatus(_) | FileRequest::IsTerminal(_), + ) + | Self::LocalSocket( + LocalSocketRequest::Create(_) + | LocalSocketRequest::CreatePair(_) + | LocalSocketRequest::Listen(_) + | LocalSocketRequest::Accept(_) + | LocalSocketRequest::Shutdown(_) + | LocalSocketRequest::SetOption(_) + | LocalSocketRequest::GetOptions(_), ) => None, } } @@ -268,6 +298,35 @@ pub enum PipeRequest { Write(WritePipeRequest), } +/// Broker-owned local socket request. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum LocalSocketRequest { + /// Create an unbound, unconnected socket. + Create(CreateLocalSocketRequest), + /// Create a pair of connected unnamed sockets. + CreatePair(CreateLocalSocketRequest), + /// Bind a socket to a name. + Bind(BindLocalSocketRequest), + /// Make a bound stream socket accept connections. + Listen(ListenLocalSocketRequest), + /// Connect a socket to a named socket. + Connect(ConnectLocalSocketRequest), + /// Accept one pending connection. + Accept(AcceptLocalSocketRequest), + /// Send bytes staged in shared memory. + Send(SendLocalSocketRequest), + /// Receive bytes into shared memory. + Receive(ReceiveLocalSocketRequest), + /// Shut down one or both directions. + Shutdown(ShutdownLocalSocketRequest), + /// Read the name of a socket or its peer. + GetName(GetLocalSocketNameRequest), + /// Store one socket option. + SetOption(SetLocalSocketOptionRequest), + /// Read a socket's type and options. + GetOptions(ObjectHandle), +} + /// Broker-owned socket object request. #[derive(Clone, Debug, PartialEq, Eq)] pub enum SocketRequest { @@ -343,6 +402,8 @@ pub enum BrokerResult { Timer(TimerResponse), /// Signal response family. Signal(SignalResponse), + /// Local socket response family. + LocalSocket(LocalSocketResponse), /// Operation failed with an ABI-neutral broker error. Error(ErrorCode), } @@ -442,6 +503,40 @@ pub enum SocketResponse { Failed(SocketError), } +/// Broker-owned local socket response. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum LocalSocketResponse { + /// Create operation response. + Create(CreateLocalSocketResponse), + /// Pair create operation response. + CreatePair(CreateLocalSocketPairResponse), + /// Bind operation completed. + Bind, + /// Listen operation completed. + Listen, + /// Connect operation completed. + Connect, + /// Accept operation response. + Accept(AcceptLocalSocketResponse), + /// Send operation response. + Send(SendLocalSocketResponse), + /// Receive operation response. + Receive(ReceiveLocalSocketResponse), + /// Shutdown operation completed. + Shutdown, + /// Name operation response. + GetName(GetLocalSocketNameResponse), + /// Option was stored. + SetOption, + /// Options operation response. + GetOptions(GetLocalSocketOptionsResponse), + /// The operation failed in a way the guest ABI reports. + /// + /// Waiting, resource, and request-validation failures use + /// [`BrokerResult::Error`] instead. + Failed(LocalSocketError), +} + /// Broker-owned fs request. #[derive(Clone, Debug, PartialEq, Eq)] pub enum FileRequest { diff --git a/litebox_broker_protocol/src/readiness.rs b/litebox_broker_protocol/src/readiness.rs index be00b1b60d..1f399679dc 100644 --- a/litebox_broker_protocol/src/readiness.rs +++ b/litebox_broker_protocol/src/readiness.rs @@ -17,6 +17,8 @@ impl ReadinessFlags { pub const HANGUP: Self = Self(1 << 2); /// The object is in an error state. pub const ERROR: Self = Self(1 << 3); + /// Both directions are shut down. + pub const CLOSED: Self = Self(1 << 4); /// Returns whether every flag in `other` is set. #[must_use] diff --git a/litebox_broker_protocol/src/socket.rs b/litebox_broker_protocol/src/socket.rs index 06422ae196..33eb313b4e 100644 --- a/litebox_broker_protocol/src/socket.rs +++ b/litebox_broker_protocol/src/socket.rs @@ -45,7 +45,7 @@ pub enum AddressFamily { } /// Communication semantics of a broker socket. -#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] #[non_exhaustive] pub enum SocketType { /// Reliable ordered byte stream. diff --git a/litebox_broker_protocol/src/wire.rs b/litebox_broker_protocol/src/wire.rs index ade0745239..721f92cf9d 100644 --- a/litebox_broker_protocol/src/wire.rs +++ b/litebox_broker_protocol/src/wire.rs @@ -33,6 +33,7 @@ use primitive::{Decoder, Encoder}; mod event; mod fs; +mod local_socket; mod pipe; mod primitive; mod signal; @@ -60,6 +61,7 @@ const REQUEST_TAG_GET_STATUS_FLAGS: u8 = 18; const REQUEST_TAG_SET_STATUS_FLAGS: u8 = 19; const REQUEST_TAG_SIGNAL: u8 = 20; const REQUEST_TAG_WRITE_CHILD_MEMORY: u8 = 21; +const REQUEST_TAG_LOCAL_SOCKET: u8 = 22; const CREATE_THREAD_TAG_THREAD: u8 = 0; const CREATE_THREAD_TAG_PROCESS: u8 = 1; @@ -89,6 +91,7 @@ const RESPONSE_TAG_STATUS_FLAGS: u8 = 18; const RESPONSE_TAG_STATUS_FLAGS_SET: u8 = 19; const RESPONSE_TAG_SIGNAL: u8 = 20; const RESPONSE_TAG_CHILD_MEMORY_WRITTEN: u8 = 21; +const RESPONSE_TAG_LOCAL_SOCKET: u8 = 22; // Reserve the top of the tag space for responses without paired requests. const RESPONSE_TAG_ERROR: u8 = 253; @@ -157,7 +160,8 @@ pub fn decode_handshake_request(frame: &[u8]) -> Result { + | REQUEST_TAG_WRITE_CHILD_MEMORY + | REQUEST_TAG_LOCAL_SOCKET => { return Err(WireError::WrongMessagePhase); } _ => return Err(WireError::InvalidTag), @@ -295,6 +299,11 @@ pub fn encode_request(request: BrokerRequest) -> Vec { encoder.request_id(request_id); signal::encode_signal_request(&mut encoder, request); } + BrokerOperation::LocalSocket(request) => { + encoder.u8(REQUEST_TAG_LOCAL_SOCKET); + encoder.request_id(request_id); + local_socket::encode_local_socket_request(&mut encoder, request); + } } encoder.finish() } @@ -324,7 +333,8 @@ pub fn decode_request(frame: &[u8]) -> Result { | REQUEST_TAG_GET_STATUS_FLAGS | REQUEST_TAG_SET_STATUS_FLAGS | REQUEST_TAG_SIGNAL - | REQUEST_TAG_WRITE_CHILD_MEMORY => {} + | REQUEST_TAG_WRITE_CHILD_MEMORY + | REQUEST_TAG_LOCAL_SOCKET => {} _ => return Err(WireError::InvalidTag), } let request_id = decoder.request_id()?; @@ -386,6 +396,9 @@ pub fn decode_request(frame: &[u8]) -> Result { } REQUEST_TAG_TIMER => BrokerOperation::Timer(timer::decode_timer_request(&mut decoder)?), REQUEST_TAG_SIGNAL => BrokerOperation::Signal(signal::decode_signal_request(&mut decoder)?), + REQUEST_TAG_LOCAL_SOCKET => { + BrokerOperation::LocalSocket(local_socket::decode_local_socket_request(&mut decoder)?) + } _ => unreachable!("active request tag was validated"), }; decoder.finish()?; @@ -471,7 +484,8 @@ pub fn decode_handshake_response(frame: &[u8]) -> Result { + | RESPONSE_TAG_CHILD_MEMORY_WRITTEN + | RESPONSE_TAG_LOCAL_SOCKET => { return Err(WireError::WrongMessagePhase); } RESPONSE_TAG_VERSION_MISMATCH => BrokerHandshakeResponse::VersionMismatch { @@ -606,6 +620,11 @@ pub fn encode_response(response: BrokerResponse) -> Vec { encoder.request_id(request_id); signal::encode_signal_response(&mut encoder, response); } + BrokerResult::LocalSocket(response) => { + encoder.u8(RESPONSE_TAG_LOCAL_SOCKET); + encoder.request_id(request_id); + local_socket::encode_local_socket_response(&mut encoder, response); + } BrokerResult::Error(error) => { encoder.u8(RESPONSE_TAG_ERROR); encoder.request_id(request_id); @@ -643,7 +662,8 @@ pub fn decode_response(frame: &[u8]) -> Result { | RESPONSE_TAG_STATUS_FLAGS | RESPONSE_TAG_STATUS_FLAGS_SET | RESPONSE_TAG_SIGNAL - | RESPONSE_TAG_CHILD_MEMORY_WRITTEN => {} + | RESPONSE_TAG_CHILD_MEMORY_WRITTEN + | RESPONSE_TAG_LOCAL_SOCKET => {} _ => return Err(WireError::InvalidTag), } let request_id = decoder.request_id()?; @@ -688,6 +708,9 @@ pub fn decode_response(frame: &[u8]) -> Result { RESPONSE_TAG_CHILD_MEMORY_WRITTEN => BrokerResult::ChildMemoryWritten, RESPONSE_TAG_TIMER => BrokerResult::Timer(timer::decode_timer_response(&mut decoder)?), RESPONSE_TAG_SIGNAL => BrokerResult::Signal(signal::decode_signal_response(&mut decoder)?), + RESPONSE_TAG_LOCAL_SOCKET => { + BrokerResult::LocalSocket(local_socket::decode_local_socket_response(&mut decoder)?) + } _ => unreachable!("active response tag was validated"), }; decoder.finish()?; @@ -805,9 +828,19 @@ mod tests { SetStatusFlagsRequest, TruncateFileRequest, UnlinkFileRequest, WriteFileRequest, WriteFileResponse, }; + use crate::local_socket::{ + AcceptLocalSocketRequest, AcceptLocalSocketResponse, BindLocalSocketRequest, + ConnectLocalSocketRequest, CreateLocalSocketPairResponse, CreateLocalSocketRequest, + CreateLocalSocketResponse, GetLocalSocketNameRequest, GetLocalSocketNameResponse, + GetLocalSocketOptionsResponse, ListenLocalSocketRequest, LocalSocketError, + LocalSocketOption, LocalSocketOptions, ReceiveLocalSocketRequest, + ReceiveLocalSocketResponse, SendLocalSocketRequest, SendLocalSocketResponse, + SetLocalSocketOptionRequest, ShutdownLocalSocketRequest, + }; use crate::message::{ - EventRequest, EventResponse, FileRequest, FileResponse, PipeRequest, PipeResponse, - SignalRequest, SignalResponse, SocketRequest, SocketResponse, TimerRequest, TimerResponse, + EventRequest, EventResponse, FileRequest, FileResponse, LocalSocketRequest, + LocalSocketResponse, PipeRequest, PipeResponse, SignalRequest, SignalResponse, + SocketRequest, SocketResponse, TimerRequest, TimerResponse, }; use crate::pipe::{ CreatePipeRequest, CreatePipeResponse, ReadPipeRequest, ReadPipeResponse, WritePipeRequest, @@ -889,6 +922,7 @@ mod tests { RESPONSE_TAG_STATUS_FLAGS_SET, RESPONSE_TAG_SIGNAL, RESPONSE_TAG_CHILD_MEMORY_WRITTEN, + RESPONSE_TAG_LOCAL_SOCKET, ], [ REQUEST_TAG_NEGOTIATE, @@ -912,6 +946,7 @@ mod tests { REQUEST_TAG_SET_STATUS_FLAGS, REQUEST_TAG_SIGNAL, REQUEST_TAG_WRITE_CHILD_MEMORY, + REQUEST_TAG_LOCAL_SOCKET, ] ); assert_eq!( @@ -1213,6 +1248,84 @@ mod tests { offset: u64::MAX, data: largest_sequence, }), + BrokerOperation::LocalSocket(LocalSocketRequest::Create(CreateLocalSocketRequest { + socket_type: SocketType::Stream, + flags: FileOpenFlags::NONBLOCKING, + })), + BrokerOperation::LocalSocket(LocalSocketRequest::CreatePair( + CreateLocalSocketRequest { + socket_type: SocketType::Datagram, + flags: FileOpenFlags::NONE, + }, + )), + BrokerOperation::LocalSocket(LocalSocketRequest::Bind(BindLocalSocketRequest { + handle, + address: largest_sequence, + user: FileUser { + user: 1000, + group: u16::MAX, + }, + mode: FileMode::from_bits(0o755).unwrap(), + })), + BrokerOperation::LocalSocket(LocalSocketRequest::Listen(ListenLocalSocketRequest { + handle, + backlog: u32::MAX, + })), + BrokerOperation::LocalSocket(LocalSocketRequest::Connect(ConnectLocalSocketRequest { + handle, + address: largest_sequence, + user: FileUser { user: 0, group: 0 }, + })), + BrokerOperation::LocalSocket(LocalSocketRequest::Accept(AcceptLocalSocketRequest { + handle, + flags: FileOpenFlags::NONBLOCKING, + })), + BrokerOperation::LocalSocket(LocalSocketRequest::Send(SendLocalSocketRequest { + handle, + buffer: largest_sequence, + address_length: u32::MAX, + user: FileUser { user: 1, group: 2 }, + })), + BrokerOperation::LocalSocket(LocalSocketRequest::Receive(ReceiveLocalSocketRequest { + handle, + buffer: largest_sequence, + capacity: u32::MAX, + peek: true, + nonblocking: true, + })), + BrokerOperation::LocalSocket(LocalSocketRequest::Shutdown( + ShutdownLocalSocketRequest { + handle, + mode: ShutdownMode::Both, + }, + )), + BrokerOperation::LocalSocket(LocalSocketRequest::GetName(GetLocalSocketNameRequest { + handle, + peer: true, + buffer: largest_sequence, + })), + BrokerOperation::LocalSocket(LocalSocketRequest::SetOption( + SetLocalSocketOptionRequest { + handle, + option: LocalSocketOption::ReceiveTimeout(Some(core::time::Duration::new( + u64::MAX, + 999_999_999, + ))), + }, + )), + BrokerOperation::LocalSocket(LocalSocketRequest::SetOption( + SetLocalSocketOptionRequest { + handle, + option: LocalSocketOption::Linger(None), + }, + )), + BrokerOperation::LocalSocket(LocalSocketRequest::SetOption( + SetLocalSocketOptionRequest { + handle, + option: LocalSocketOption::Broadcast(true), + }, + )), + BrokerOperation::LocalSocket(LocalSocketRequest::GetOptions(handle)), ]; let mut maximum_encoded_size = 0; @@ -1458,6 +1571,52 @@ mod tests { BrokerResult::Timer(TimerResponse::Read(ReadTimerResponse { expirations: u64::MAX, })), + BrokerResult::LocalSocket(LocalSocketResponse::Create(CreateLocalSocketResponse { + handle, + })), + BrokerResult::LocalSocket(LocalSocketResponse::CreatePair( + CreateLocalSocketPairResponse { + first: handle, + second: ObjectHandle(u64::MAX), + }, + )), + BrokerResult::LocalSocket(LocalSocketResponse::Bind), + BrokerResult::LocalSocket(LocalSocketResponse::Listen), + BrokerResult::LocalSocket(LocalSocketResponse::Connect), + BrokerResult::LocalSocket(LocalSocketResponse::Accept(AcceptLocalSocketResponse { + handle, + })), + BrokerResult::LocalSocket(LocalSocketResponse::Send(SendLocalSocketResponse { + sent: u32::MAX, + })), + BrokerResult::LocalSocket(LocalSocketResponse::Receive(ReceiveLocalSocketResponse { + received: 1, + length: u32::MAX, + source_length: 109, + })), + BrokerResult::LocalSocket(LocalSocketResponse::Shutdown), + BrokerResult::LocalSocket(LocalSocketResponse::GetName(GetLocalSocketNameResponse { + length: 3, + })), + BrokerResult::LocalSocket(LocalSocketResponse::SetOption), + BrokerResult::LocalSocket(LocalSocketResponse::GetOptions( + GetLocalSocketOptionsResponse { + socket_type: SocketType::Datagram, + options: LocalSocketOptions { + receive_timeout: Some(core::time::Duration::new(u64::MAX, 1)), + send_timeout: None, + linger: Some(core::time::Duration::from_secs(3)), + reuse_address: true, + keep_alive: false, + broadcast: true, + }, + }, + )), + BrokerResult::LocalSocket(LocalSocketResponse::Failed(LocalSocketError::AddressInUse)), + BrokerResult::LocalSocket(LocalSocketResponse::Failed(LocalSocketError::NotPermitted)), + BrokerResult::LocalSocket(LocalSocketResponse::Failed(LocalSocketError::File( + FileError::NoSuchFileOrDirectory, + ))), BrokerResult::Signal(SignalResponse::Open(OpenSignalsResponse { handle })), BrokerResult::Signal(SignalResponse::Sent), BrokerResult::Signal(SignalResponse::Take(PendingSignal { diff --git a/litebox_broker_protocol/src/wire/fs.rs b/litebox_broker_protocol/src/wire/fs.rs index 3474a9c307..0e469b6475 100644 --- a/litebox_broker_protocol/src/wire/fs.rs +++ b/litebox_broker_protocol/src/wire/fs.rs @@ -152,7 +152,7 @@ pub(super) fn encode_fs_request(encoder: &mut Encoder, request: FileRequest) { } } -fn decode_mode(decoder: &mut Decoder<'_>) -> Result { +pub(super) fn decode_mode(decoder: &mut Decoder<'_>) -> Result { FileMode::from_bits(decoder.u16()?).ok_or(WireError::InvalidTag) } @@ -312,7 +312,7 @@ pub(super) fn decode_fs_response(decoder: &mut Decoder<'_>) -> Result 1, FileError::NoWritePermissions => 2, @@ -338,7 +338,7 @@ fn encode_file_error(encoder: &mut Encoder, error: FileError) { }); } -fn decode_file_error(decoder: &mut Decoder<'_>) -> Result { +pub(super) fn decode_file_error(decoder: &mut Decoder<'_>) -> Result { match decoder.u8()? { 1 => Ok(FileError::AccessNotAllowed), 2 => Ok(FileError::NoWritePermissions), @@ -365,12 +365,12 @@ fn decode_file_error(decoder: &mut Decoder<'_>) -> Result } } -fn encode_user(encoder: &mut Encoder, user: FileUser) { +pub(super) fn encode_user(encoder: &mut Encoder, user: FileUser) { encoder.u16(user.user); encoder.u16(user.group); } -fn decode_user(decoder: &mut Decoder<'_>) -> Result { +pub(super) fn decode_user(decoder: &mut Decoder<'_>) -> Result { Ok(FileUser { user: decoder.u16()?, group: decoder.u16()?, diff --git a/litebox_broker_protocol/src/wire/local_socket.rs b/litebox_broker_protocol/src/wire/local_socket.rs new file mode 100644 index 0000000000..f497cb7b1d --- /dev/null +++ b/litebox_broker_protocol/src/wire/local_socket.rs @@ -0,0 +1,412 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +use core::time::Duration; + +use super::WireError; +use super::fs::{ + decode_file_error, decode_mode, decode_open_flags, decode_user, encode_file_error, encode_user, +}; +use super::primitive::{Decoder, Encoder}; +use super::socket::{ + decode_shutdown_mode, decode_socket_type, encode_shutdown_mode, encode_socket_type, +}; +use crate::local_socket::{ + AcceptLocalSocketRequest, AcceptLocalSocketResponse, BindLocalSocketRequest, + ConnectLocalSocketRequest, CreateLocalSocketPairResponse, CreateLocalSocketRequest, + CreateLocalSocketResponse, GetLocalSocketNameRequest, GetLocalSocketNameResponse, + GetLocalSocketOptionsResponse, ListenLocalSocketRequest, LocalSocketError, LocalSocketOption, + LocalSocketOptions, ReceiveLocalSocketRequest, ReceiveLocalSocketResponse, + SendLocalSocketRequest, SendLocalSocketResponse, SetLocalSocketOptionRequest, + ShutdownLocalSocketRequest, +}; +use crate::message::{LocalSocketRequest, LocalSocketResponse}; + +const TAG_CREATE: u8 = 0; +const TAG_CREATE_PAIR: u8 = 1; +const TAG_BIND: u8 = 2; +const TAG_LISTEN: u8 = 3; +const TAG_CONNECT: u8 = 4; +const TAG_ACCEPT: u8 = 5; +const TAG_SEND: u8 = 6; +const TAG_RECEIVE: u8 = 7; +const TAG_SHUTDOWN: u8 = 8; +const TAG_GET_NAME: u8 = 9; +const TAG_SET_OPTION: u8 = 10; +const TAG_GET_OPTIONS: u8 = 11; +const RESPONSE_TAG_FAILED: u8 = 12; + +const OPTION_TAG_RECEIVE_TIMEOUT: u8 = 0; +const OPTION_TAG_SEND_TIMEOUT: u8 = 1; +const OPTION_TAG_LINGER: u8 = 2; +const OPTION_TAG_REUSE_ADDRESS: u8 = 3; +const OPTION_TAG_KEEP_ALIVE: u8 = 4; +const OPTION_TAG_BROADCAST: u8 = 5; + +const ERROR_TAG_ADDRESS_IN_USE: u8 = 0; +const ERROR_TAG_CONNECTION_REFUSED: u8 = 1; +const ERROR_TAG_WRONG_TYPE: u8 = 2; +const ERROR_TAG_INVALID_ARGUMENT: u8 = 3; +const ERROR_TAG_ALREADY_CONNECTED: u8 = 4; +const ERROR_TAG_NOT_CONNECTED: u8 = 5; +const ERROR_TAG_UNSUPPORTED: u8 = 6; +const ERROR_TAG_MESSAGE_TOO_LARGE: u8 = 7; +const ERROR_TAG_BROKEN_PIPE: u8 = 8; +const ERROR_TAG_FILE: u8 = 9; +const ERROR_TAG_NOT_PERMITTED: u8 = 10; + +pub(super) fn encode_local_socket_request(encoder: &mut Encoder, request: LocalSocketRequest) { + match request { + LocalSocketRequest::Create(request) => { + encoder.u8(TAG_CREATE); + encode_create_request(encoder, request); + } + LocalSocketRequest::CreatePair(request) => { + encoder.u8(TAG_CREATE_PAIR); + encode_create_request(encoder, request); + } + LocalSocketRequest::Bind(request) => { + encoder.u8(TAG_BIND); + encoder.handle(request.handle); + encoder.shared_buffer_sequence(request.address); + encode_user(encoder, request.user); + encoder.u16(request.mode.bits()); + } + LocalSocketRequest::Listen(request) => { + encoder.u8(TAG_LISTEN); + encoder.handle(request.handle); + encoder.u32(request.backlog); + } + LocalSocketRequest::Connect(request) => { + encoder.u8(TAG_CONNECT); + encoder.handle(request.handle); + encoder.shared_buffer_sequence(request.address); + encode_user(encoder, request.user); + } + LocalSocketRequest::Accept(request) => { + encoder.u8(TAG_ACCEPT); + encoder.handle(request.handle); + encoder.u16(request.flags.bits()); + } + LocalSocketRequest::Send(request) => { + encoder.u8(TAG_SEND); + encoder.handle(request.handle); + encoder.shared_buffer_sequence(request.buffer); + encoder.u32(request.address_length); + encode_user(encoder, request.user); + } + LocalSocketRequest::Receive(request) => { + encoder.u8(TAG_RECEIVE); + encoder.handle(request.handle); + encoder.shared_buffer_sequence(request.buffer); + encoder.u32(request.capacity); + encode_bool(encoder, request.peek); + encode_bool(encoder, request.nonblocking); + } + LocalSocketRequest::Shutdown(request) => { + encoder.u8(TAG_SHUTDOWN); + encoder.handle(request.handle); + encode_shutdown_mode(encoder, request.mode); + } + LocalSocketRequest::GetName(request) => { + encoder.u8(TAG_GET_NAME); + encoder.handle(request.handle); + encode_bool(encoder, request.peer); + encoder.shared_buffer_sequence(request.buffer); + } + LocalSocketRequest::SetOption(request) => { + encoder.u8(TAG_SET_OPTION); + encoder.handle(request.handle); + encode_option(encoder, request.option); + } + LocalSocketRequest::GetOptions(handle) => { + encoder.u8(TAG_GET_OPTIONS); + encoder.handle(handle); + } + } +} + +pub(super) fn decode_local_socket_request( + decoder: &mut Decoder<'_>, +) -> Result { + Ok(match decoder.u8()? { + TAG_CREATE => LocalSocketRequest::Create(decode_create_request(decoder)?), + TAG_CREATE_PAIR => LocalSocketRequest::CreatePair(decode_create_request(decoder)?), + TAG_BIND => LocalSocketRequest::Bind(BindLocalSocketRequest { + handle: decoder.handle()?, + address: decoder.shared_buffer_sequence()?, + user: decode_user(decoder)?, + mode: decode_mode(decoder)?, + }), + TAG_LISTEN => LocalSocketRequest::Listen(ListenLocalSocketRequest { + handle: decoder.handle()?, + backlog: decoder.u32()?, + }), + TAG_CONNECT => LocalSocketRequest::Connect(ConnectLocalSocketRequest { + handle: decoder.handle()?, + address: decoder.shared_buffer_sequence()?, + user: decode_user(decoder)?, + }), + TAG_ACCEPT => LocalSocketRequest::Accept(AcceptLocalSocketRequest { + handle: decoder.handle()?, + flags: decode_open_flags(decoder)?, + }), + TAG_SEND => LocalSocketRequest::Send(SendLocalSocketRequest { + handle: decoder.handle()?, + buffer: decoder.shared_buffer_sequence()?, + address_length: decoder.u32()?, + user: decode_user(decoder)?, + }), + TAG_RECEIVE => LocalSocketRequest::Receive(ReceiveLocalSocketRequest { + handle: decoder.handle()?, + buffer: decoder.shared_buffer_sequence()?, + capacity: decoder.u32()?, + peek: decode_bool(decoder)?, + nonblocking: decode_bool(decoder)?, + }), + TAG_SHUTDOWN => LocalSocketRequest::Shutdown(ShutdownLocalSocketRequest { + handle: decoder.handle()?, + mode: decode_shutdown_mode(decoder)?, + }), + TAG_GET_NAME => LocalSocketRequest::GetName(GetLocalSocketNameRequest { + handle: decoder.handle()?, + peer: decode_bool(decoder)?, + buffer: decoder.shared_buffer_sequence()?, + }), + TAG_SET_OPTION => LocalSocketRequest::SetOption(SetLocalSocketOptionRequest { + handle: decoder.handle()?, + option: decode_option(decoder)?, + }), + TAG_GET_OPTIONS => LocalSocketRequest::GetOptions(decoder.handle()?), + _ => return Err(WireError::InvalidTag), + }) +} + +pub(super) fn encode_local_socket_response(encoder: &mut Encoder, response: LocalSocketResponse) { + match response { + LocalSocketResponse::Create(response) => { + encoder.u8(TAG_CREATE); + encoder.handle(response.handle); + } + LocalSocketResponse::CreatePair(response) => { + encoder.u8(TAG_CREATE_PAIR); + encoder.handle(response.first); + encoder.handle(response.second); + } + LocalSocketResponse::Bind => encoder.u8(TAG_BIND), + LocalSocketResponse::Listen => encoder.u8(TAG_LISTEN), + LocalSocketResponse::Connect => encoder.u8(TAG_CONNECT), + LocalSocketResponse::Accept(response) => { + encoder.u8(TAG_ACCEPT); + encoder.handle(response.handle); + } + LocalSocketResponse::Send(response) => { + encoder.u8(TAG_SEND); + encoder.u32(response.sent); + } + LocalSocketResponse::Receive(response) => { + encoder.u8(TAG_RECEIVE); + encoder.u32(response.received); + encoder.u32(response.length); + encoder.u32(response.source_length); + } + LocalSocketResponse::Shutdown => encoder.u8(TAG_SHUTDOWN), + LocalSocketResponse::GetName(response) => { + encoder.u8(TAG_GET_NAME); + encoder.u32(response.length); + } + LocalSocketResponse::SetOption => encoder.u8(TAG_SET_OPTION), + LocalSocketResponse::GetOptions(response) => { + encoder.u8(TAG_GET_OPTIONS); + encode_socket_type(encoder, response.socket_type); + encode_options(encoder, response.options); + } + LocalSocketResponse::Failed(error) => { + encoder.u8(RESPONSE_TAG_FAILED); + encode_error(encoder, error); + } + } +} + +pub(super) fn decode_local_socket_response( + decoder: &mut Decoder<'_>, +) -> Result { + Ok(match decoder.u8()? { + TAG_CREATE => LocalSocketResponse::Create(CreateLocalSocketResponse { + handle: decoder.handle()?, + }), + TAG_CREATE_PAIR => LocalSocketResponse::CreatePair(CreateLocalSocketPairResponse { + first: decoder.handle()?, + second: decoder.handle()?, + }), + TAG_BIND => LocalSocketResponse::Bind, + TAG_LISTEN => LocalSocketResponse::Listen, + TAG_CONNECT => LocalSocketResponse::Connect, + TAG_ACCEPT => LocalSocketResponse::Accept(AcceptLocalSocketResponse { + handle: decoder.handle()?, + }), + TAG_SEND => LocalSocketResponse::Send(SendLocalSocketResponse { + sent: decoder.u32()?, + }), + TAG_RECEIVE => LocalSocketResponse::Receive(ReceiveLocalSocketResponse { + received: decoder.u32()?, + length: decoder.u32()?, + source_length: decoder.u32()?, + }), + TAG_SHUTDOWN => LocalSocketResponse::Shutdown, + TAG_GET_NAME => LocalSocketResponse::GetName(GetLocalSocketNameResponse { + length: decoder.u32()?, + }), + TAG_SET_OPTION => LocalSocketResponse::SetOption, + TAG_GET_OPTIONS => LocalSocketResponse::GetOptions(GetLocalSocketOptionsResponse { + socket_type: decode_socket_type(decoder)?, + options: decode_options(decoder)?, + }), + RESPONSE_TAG_FAILED => LocalSocketResponse::Failed(decode_error(decoder)?), + _ => return Err(WireError::InvalidTag), + }) +} + +fn encode_create_request(encoder: &mut Encoder, request: CreateLocalSocketRequest) { + encode_socket_type(encoder, request.socket_type); + encoder.u16(request.flags.bits()); +} + +fn decode_create_request(decoder: &mut Decoder<'_>) -> Result { + Ok(CreateLocalSocketRequest { + socket_type: decode_socket_type(decoder)?, + flags: decode_open_flags(decoder)?, + }) +} + +fn encode_bool(encoder: &mut Encoder, value: bool) { + encoder.u8(u8::from(value)); +} + +fn decode_bool(decoder: &mut Decoder<'_>) -> Result { + match decoder.u8()? { + 0 => Ok(false), + 1 => Ok(true), + _ => Err(WireError::InvalidTag), + } +} + +fn encode_duration(encoder: &mut Encoder, duration: Option) { + match duration { + None => encoder.u8(0), + Some(duration) => { + encoder.u8(1); + encoder.u64(duration.as_secs()); + encoder.u32(duration.subsec_nanos()); + } + } +} + +fn decode_duration(decoder: &mut Decoder<'_>) -> Result, WireError> { + if !decode_bool(decoder)? { + return Ok(None); + } + let seconds = decoder.u64()?; + let nanoseconds = decoder.u32()?; + if nanoseconds >= 1_000_000_000 { + return Err(WireError::InvalidTag); + } + Ok(Some(Duration::new(seconds, nanoseconds))) +} + +fn encode_option(encoder: &mut Encoder, option: LocalSocketOption) { + match option { + LocalSocketOption::ReceiveTimeout(timeout) => { + encoder.u8(OPTION_TAG_RECEIVE_TIMEOUT); + encode_duration(encoder, timeout); + } + LocalSocketOption::SendTimeout(timeout) => { + encoder.u8(OPTION_TAG_SEND_TIMEOUT); + encode_duration(encoder, timeout); + } + LocalSocketOption::Linger(timeout) => { + encoder.u8(OPTION_TAG_LINGER); + encode_duration(encoder, timeout); + } + LocalSocketOption::ReuseAddress(value) => { + encoder.u8(OPTION_TAG_REUSE_ADDRESS); + encode_bool(encoder, value); + } + LocalSocketOption::KeepAlive(value) => { + encoder.u8(OPTION_TAG_KEEP_ALIVE); + encode_bool(encoder, value); + } + LocalSocketOption::Broadcast(value) => { + encoder.u8(OPTION_TAG_BROADCAST); + encode_bool(encoder, value); + } + } +} + +fn decode_option(decoder: &mut Decoder<'_>) -> Result { + Ok(match decoder.u8()? { + OPTION_TAG_RECEIVE_TIMEOUT => LocalSocketOption::ReceiveTimeout(decode_duration(decoder)?), + OPTION_TAG_SEND_TIMEOUT => LocalSocketOption::SendTimeout(decode_duration(decoder)?), + OPTION_TAG_LINGER => LocalSocketOption::Linger(decode_duration(decoder)?), + OPTION_TAG_REUSE_ADDRESS => LocalSocketOption::ReuseAddress(decode_bool(decoder)?), + OPTION_TAG_KEEP_ALIVE => LocalSocketOption::KeepAlive(decode_bool(decoder)?), + OPTION_TAG_BROADCAST => LocalSocketOption::Broadcast(decode_bool(decoder)?), + _ => return Err(WireError::InvalidTag), + }) +} + +fn encode_options(encoder: &mut Encoder, options: LocalSocketOptions) { + encode_duration(encoder, options.receive_timeout); + encode_duration(encoder, options.send_timeout); + encode_duration(encoder, options.linger); + encode_bool(encoder, options.reuse_address); + encode_bool(encoder, options.keep_alive); + encode_bool(encoder, options.broadcast); +} + +fn decode_options(decoder: &mut Decoder<'_>) -> Result { + Ok(LocalSocketOptions { + receive_timeout: decode_duration(decoder)?, + send_timeout: decode_duration(decoder)?, + linger: decode_duration(decoder)?, + reuse_address: decode_bool(decoder)?, + keep_alive: decode_bool(decoder)?, + broadcast: decode_bool(decoder)?, + }) +} + +fn encode_error(encoder: &mut Encoder, error: LocalSocketError) { + match error { + LocalSocketError::AddressInUse => encoder.u8(ERROR_TAG_ADDRESS_IN_USE), + LocalSocketError::ConnectionRefused => encoder.u8(ERROR_TAG_CONNECTION_REFUSED), + LocalSocketError::WrongType => encoder.u8(ERROR_TAG_WRONG_TYPE), + LocalSocketError::InvalidArgument => encoder.u8(ERROR_TAG_INVALID_ARGUMENT), + LocalSocketError::AlreadyConnected => encoder.u8(ERROR_TAG_ALREADY_CONNECTED), + LocalSocketError::NotConnected => encoder.u8(ERROR_TAG_NOT_CONNECTED), + LocalSocketError::Unsupported => encoder.u8(ERROR_TAG_UNSUPPORTED), + LocalSocketError::MessageTooLarge => encoder.u8(ERROR_TAG_MESSAGE_TOO_LARGE), + LocalSocketError::BrokenPipe => encoder.u8(ERROR_TAG_BROKEN_PIPE), + LocalSocketError::NotPermitted => encoder.u8(ERROR_TAG_NOT_PERMITTED), + LocalSocketError::File(error) => { + encoder.u8(ERROR_TAG_FILE); + encode_file_error(encoder, error); + } + } +} + +fn decode_error(decoder: &mut Decoder<'_>) -> Result { + Ok(match decoder.u8()? { + ERROR_TAG_ADDRESS_IN_USE => LocalSocketError::AddressInUse, + ERROR_TAG_CONNECTION_REFUSED => LocalSocketError::ConnectionRefused, + ERROR_TAG_WRONG_TYPE => LocalSocketError::WrongType, + ERROR_TAG_INVALID_ARGUMENT => LocalSocketError::InvalidArgument, + ERROR_TAG_ALREADY_CONNECTED => LocalSocketError::AlreadyConnected, + ERROR_TAG_NOT_CONNECTED => LocalSocketError::NotConnected, + ERROR_TAG_UNSUPPORTED => LocalSocketError::Unsupported, + ERROR_TAG_MESSAGE_TOO_LARGE => LocalSocketError::MessageTooLarge, + ERROR_TAG_BROKEN_PIPE => LocalSocketError::BrokenPipe, + ERROR_TAG_NOT_PERMITTED => LocalSocketError::NotPermitted, + ERROR_TAG_FILE => LocalSocketError::File(decode_file_error(decoder)?), + _ => return Err(WireError::InvalidTag), + }) +} diff --git a/litebox_broker_protocol/src/wire/socket.rs b/litebox_broker_protocol/src/wire/socket.rs index 1e5731ca4e..acc027a8f4 100644 --- a/litebox_broker_protocol/src/wire/socket.rs +++ b/litebox_broker_protocol/src/wire/socket.rs @@ -66,10 +66,7 @@ pub(super) fn encode_socket_request(encoder: &mut Encoder, request: SocketReques encoder.u8(match request.address_family { AddressFamily::Ipv4 => ADDRESS_FAMILY_TAG_IPV4, }); - encoder.u8(match request.socket_type { - SocketType::Stream => TYPE_TAG_STREAM, - SocketType::Datagram => TYPE_TAG_DATAGRAM, - }); + encode_socket_type(encoder, request.socket_type); encoder.u8(match request.protocol { IpProtocol::Tcp => IP_PROTOCOL_TAG_TCP, IpProtocol::Udp => IP_PROTOCOL_TAG_UDP, @@ -124,13 +121,7 @@ pub(super) fn encode_socket_request(encoder: &mut Encoder, request: SocketReques SocketRequest::Shutdown(request) => { encoder.u8(SOCKET_TAG_SHUTDOWN); encoder.handle(request.handle); - encoder.u8(match request.mode { - ShutdownMode::Read => SHUTDOWN_TAG_READ, - ShutdownMode::Write => SHUTDOWN_TAG_WRITE, - ShutdownMode::Both => SHUTDOWN_TAG_BOTH, - ShutdownMode::Abort => SHUTDOWN_TAG_ABORT, - ShutdownMode::StopListening => SHUTDOWN_TAG_STOP_LISTENING, - }); + encode_shutdown_mode(encoder, request.mode); } SocketRequest::SetTcpOption(request) => { encoder.u8(SOCKET_TAG_SET_TCP_OPTION); @@ -156,11 +147,7 @@ pub(super) fn decode_socket_request(decoder: &mut Decoder<'_>) -> Result AddressFamily::Ipv4, _ => return Err(WireError::InvalidTag), }, - socket_type: match decoder.u8()? { - TYPE_TAG_STREAM => SocketType::Stream, - TYPE_TAG_DATAGRAM => SocketType::Datagram, - _ => return Err(WireError::InvalidTag), - }, + socket_type: decode_socket_type(decoder)?, protocol: match decoder.u8()? { IP_PROTOCOL_TAG_TCP => IpProtocol::Tcp, IP_PROTOCOL_TAG_UDP => IpProtocol::Udp, @@ -207,14 +194,7 @@ pub(super) fn decode_socket_request(decoder: &mut Decoder<'_>) -> Result Ok(SocketRequest::Shutdown(ShutdownSocketRequest { handle: decoder.handle()?, - mode: match decoder.u8()? { - SHUTDOWN_TAG_READ => ShutdownMode::Read, - SHUTDOWN_TAG_WRITE => ShutdownMode::Write, - SHUTDOWN_TAG_BOTH => ShutdownMode::Both, - SHUTDOWN_TAG_ABORT => ShutdownMode::Abort, - SHUTDOWN_TAG_STOP_LISTENING => ShutdownMode::StopListening, - _ => return Err(WireError::InvalidTag), - }, + mode: decode_shutdown_mode(decoder)?, })), SOCKET_TAG_SET_TCP_OPTION => Ok(SocketRequest::SetTcpOption(SetTcpOptionRequest { handle: decoder.handle()?, @@ -352,6 +332,42 @@ pub(super) fn decode_socket_response( } } +pub(super) fn encode_socket_type(encoder: &mut Encoder, socket_type: SocketType) { + encoder.u8(match socket_type { + SocketType::Stream => TYPE_TAG_STREAM, + SocketType::Datagram => TYPE_TAG_DATAGRAM, + }); +} + +pub(super) fn decode_socket_type(decoder: &mut Decoder<'_>) -> Result { + match decoder.u8()? { + TYPE_TAG_STREAM => Ok(SocketType::Stream), + TYPE_TAG_DATAGRAM => Ok(SocketType::Datagram), + _ => Err(WireError::InvalidTag), + } +} + +pub(super) fn encode_shutdown_mode(encoder: &mut Encoder, mode: ShutdownMode) { + encoder.u8(match mode { + ShutdownMode::Read => SHUTDOWN_TAG_READ, + ShutdownMode::Write => SHUTDOWN_TAG_WRITE, + ShutdownMode::Both => SHUTDOWN_TAG_BOTH, + ShutdownMode::Abort => SHUTDOWN_TAG_ABORT, + ShutdownMode::StopListening => SHUTDOWN_TAG_STOP_LISTENING, + }); +} + +pub(super) fn decode_shutdown_mode(decoder: &mut Decoder<'_>) -> Result { + match decoder.u8()? { + SHUTDOWN_TAG_READ => Ok(ShutdownMode::Read), + SHUTDOWN_TAG_WRITE => Ok(ShutdownMode::Write), + SHUTDOWN_TAG_BOTH => Ok(ShutdownMode::Both), + SHUTDOWN_TAG_ABORT => Ok(ShutdownMode::Abort), + SHUTDOWN_TAG_STOP_LISTENING => Ok(ShutdownMode::StopListening), + _ => Err(WireError::InvalidTag), + } +} + fn encode_tcp_option_name(encoder: &mut Encoder, name: TcpOptionName) { encoder.u8(match name { TcpOptionName::NoDelay => TCP_OPTION_TAG_NODELAY, diff --git a/litebox_common_linux/src/errno/mod.rs b/litebox_common_linux/src/errno/mod.rs index c90ce4c090..8f2217e46c 100644 --- a/litebox_common_linux/src/errno/mod.rs +++ b/litebox_common_linux/src/errno/mod.rs @@ -762,6 +762,51 @@ impl From for Errno { } } +impl From for Errno { + fn from(value: litebox::local_sockets::errors::LocalSocketError) -> Self { + use litebox::local_sockets::errors::LocalSocketError; + use litebox_broker_protocol::fs::FileError; + use litebox_broker_protocol::local_socket::LocalSocketError as SocketError; + match value { + LocalSocketError::Socket(error) => match error { + SocketError::AddressInUse => Errno::EADDRINUSE, + SocketError::ConnectionRefused => Errno::ECONNREFUSED, + SocketError::WrongType => Errno::EPROTOTYPE, + SocketError::InvalidArgument => Errno::EINVAL, + SocketError::AlreadyConnected => Errno::EISCONN, + SocketError::NotConnected => Errno::ENOTCONN, + SocketError::Unsupported => Errno::EOPNOTSUPP, + SocketError::MessageTooLarge => Errno::EMSGSIZE, + SocketError::BrokenPipe => Errno::EPIPE, + SocketError::NotPermitted => Errno::EPERM, + SocketError::File(error) => match error { + FileError::NoSuchFileOrDirectory | FileError::MissingComponent => Errno::ENOENT, + FileError::AccessNotAllowed + | FileError::NoWritePermissions + | FileError::NoSearchPermissions => Errno::EACCES, + FileError::ReadOnlyFs => Errno::EROFS, + FileError::AlreadyExists => Errno::EADDRINUSE, + FileError::InvalidPathname => Errno::EINVAL, + FileError::ComponentNotDirectory | FileError::NotDirectory => Errno::ENOTDIR, + _ => Errno::EIO, + }, + _ => Errno::EIO, + }, + LocalSocketError::WouldBlock => Errno::EAGAIN, + LocalSocketError::Interrupted => Errno::EINTR, + LocalSocketError::WaitError(error) => match error { + litebox::event::wait::WaitError::Interrupted => Errno::ERESTARTSYS, + litebox::event::wait::WaitError::TimedOut => Errno::ETIMEDOUT, + }, + LocalSocketError::ResourceExhausted => Errno::ENOBUFS, + LocalSocketError::OutOfMemory => Errno::ENOMEM, + LocalSocketError::PermissionDenied => Errno::EACCES, + LocalSocketError::Unsupported => Errno::EINVAL, + _ => Errno::EIO, + } + } +} + impl From for Errno { fn from(value: litebox::pipes::errors::ClosedError) -> Self { match value { diff --git a/litebox_common_linux/src/program_startup.rs b/litebox_common_linux/src/program_startup.rs index 40a2fa0fdb..3d2bb691ca 100644 --- a/litebox_common_linux/src/program_startup.rs +++ b/litebox_common_linux/src/program_startup.rs @@ -26,6 +26,7 @@ const INHERITED_FD_HEADER_SIZE: usize = size_of::() + size_of::() + si const FORK_REGION_SIZE: usize = size_of::<[u64; 2]>() + size_of::(); const FILE_TAG: u8 = 0; const PIPE_TAG: u8 = 1; +const LOCAL_SOCKET_TAG: u8 = 2; const PROGRAM_STARTUP_TAG: u8 = 0; const FORK_STARTUP_TAG: u8 = 1; @@ -115,6 +116,8 @@ pub enum InheritedFdKind { /// Which end. endpoint: HalfPipeType, }, + /// A local socket, such as a Unix domain socket. + LocalSocket, } /// Linux process state needed to continue a child duplicated by `fork` in a fresh runner. @@ -517,7 +520,7 @@ impl InheritedFd { fn encoded_len(&self) -> usize { INHERITED_FD_HEADER_SIZE + match self.kind { - InheritedFdKind::File => 0, + InheritedFdKind::File | InheritedFdKind::LocalSocket => 0, InheritedFdKind::Pipe { .. } => size_of::(), } } @@ -527,6 +530,7 @@ impl InheritedFd { push_u64(output, self.handle.0); match self.kind { InheritedFdKind::File => output.push(FILE_TAG), + InheritedFdKind::LocalSocket => output.push(LOCAL_SOCKET_TAG), InheritedFdKind::Pipe { endpoint } => { output.push(PIPE_TAG); output.push(match endpoint { @@ -542,6 +546,7 @@ impl InheritedFd { let handle = ObjectHandle(read_u64(input)?); let kind = match read_u8(input)? { FILE_TAG => InheritedFdKind::File, + LOCAL_SOCKET_TAG => InheritedFdKind::LocalSocket, PIPE_TAG => InheritedFdKind::Pipe { endpoint: match read_u8(input)? { 0 => HalfPipeType::ReceiverHalf, @@ -726,6 +731,11 @@ mod tests { endpoint: HalfPipeType::SenderHalf, }, }, + InheritedFd { + fd: 7, + handle: ObjectHandle(10), + kind: InheritedFdKind::LocalSocket, + }, ], }; @@ -827,6 +837,14 @@ mod tests { }, close_on_exec: true, }, + ForkedFd { + inherited: InheritedFd { + fd: 4, + handle: ObjectHandle(9), + kind: InheritedFdKind::LocalSocket, + }, + close_on_exec: false, + }, ], } } diff --git a/litebox_runner_linux_userland/tests/fork_unix_parent.c b/litebox_runner_linux_userland/tests/fork_unix_parent.c new file mode 100644 index 0000000000..9dc6b12b6e --- /dev/null +++ b/litebox_runner_linux_userland/tests/fork_unix_parent.c @@ -0,0 +1,225 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT license. + +// Checks that forked and exec'd children share Unix domain sockets with their parent. + +#define _GNU_SOURCE +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#define SOCKET_PATH "/tmp/fork_unix_parent.sock" +// Bounds every wait, so a failure cannot hang the test. +#define TIMEOUT_MS 30000 + +static int wait_readable(int fd) { + struct pollfd poll_fd = {.fd = fd, .events = POLLIN}; + int ready; + do { + ready = poll(&poll_fd, 1, TIMEOUT_MS); + } while (ready == -1 && errno == EINTR); + return ready == 1; +} + +static int read_exact(int fd, char *buffer, size_t length) { + size_t received = 0; + while (received < length) { + if (!wait_readable(fd)) { + return 0; + } + ssize_t n = read(fd, buffer + received, length - received); + if (n < 0 && errno == EINTR) { + continue; + } + if (n <= 0) { + return 0; + } + received += n; + } + return 1; +} + +static int expect(int fd, const char *message) { + char buffer[16] = {0}; + size_t length = strlen(message); + return read_exact(fd, buffer, length) && memcmp(buffer, message, length) == 0; +} + +static int send_all(int fd, const char *message) { + size_t length = strlen(message); + return write(fd, message, length) == (ssize_t)length; +} + +// `listener` is nonblocking, so a connection taken by the other process sharing it fails the accept +// rather than blocking it. +static int accept_within(int listener) { + return wait_readable(listener) ? accept(listener, NULL, NULL) : -1; +} + +// Returns the exit code of `child`, or -1 if it did not exit normally. +static int wait_code(pid_t child) { + int status = 0; + pid_t waited; + do { + waited = waitpid(child, &status, 0); + } while (waited == -1 && errno == EINTR); + return waited == child && WIFEXITED(status) ? WEXITSTATUS(status) : -1; +} + +// Reaps `child`, first killing it if the parent already failed, since it may be waiting on the +// parent. Returns the failures including one if the child did not exit with `expected`. +static int finish(pid_t child, int failures, int expected) { + if (failures != 0) { + kill(child, SIGKILL); + } + return failures + (wait_code(child) != expected); +} + +// A forked child talks to its parent over an inherited stream socket pair. +static int stream_pair(void) { + int pair[2]; + if (socketpair(AF_UNIX, SOCK_STREAM, 0, pair) != 0) { + return 1; + } + pid_t child = fork(); + if (child == 0) { + close(pair[0]); + _exit(send_all(pair[1], "ping") && expect(pair[1], "pong") ? 0 : 1); + } + close(pair[1]); + if (child < 0) { + close(pair[0]); + return 1; + } + int failures = !expect(pair[0], "ping"); + failures += !send_all(pair[0], "pong"); + failures = finish(child, failures, 0); + // The child's exit closed the only other reference to the peer. + char byte; + failures += !wait_readable(pair[0]) || read(pair[0], &byte, 1) != 0; + close(pair[0]); + return failures; +} + +// A forked child sends a datagram over an inherited datagram socket pair. +static int datagram_pair(void) { + int pair[2]; + if (socketpair(AF_UNIX, SOCK_DGRAM, 0, pair) != 0) { + return 1; + } + pid_t child = fork(); + if (child == 0) { + close(pair[0]); + _exit(send(pair[1], "dgram", 5, 0) == 5 ? 0 : 1); + } + close(pair[1]); + if (child < 0) { + close(pair[0]); + return 1; + } + char buffer[16] = {0}; + // Closing a datagram peer does not end receives, so the child exits before a nonblocking + // receive. + int failures = finish(child, 0, 0); + failures += + recv(pair[0], buffer, sizeof buffer, MSG_DONTWAIT) != 5 || memcmp(buffer, "dgram", 5) != 0; + close(pair[0]); + return failures; +} + +// A forked child connects to its parent's listener by path, then accepts the parent's connection +// on the listener it inherited. +static int path_listener(void) { + struct sockaddr_un address = {.sun_family = AF_UNIX}; + strncpy(address.sun_path, SOCKET_PATH, sizeof address.sun_path - 1); + unlink(SOCKET_PATH); + int listener = socket(AF_UNIX, SOCK_STREAM | SOCK_NONBLOCK, 0); + if (listener < 0 || bind(listener, (struct sockaddr *)&address, sizeof address) != 0 || + listen(listener, 4) != 0) { + return 1; + } + pid_t child = fork(); + if (child == 0) { + int ok = 1; + struct sockaddr_un name = {0}; + socklen_t length = sizeof name; + ok &= getsockname(listener, (struct sockaddr *)&name, &length) == 0 && + strcmp(name.sun_path, SOCKET_PATH) == 0; + int client = socket(AF_UNIX, SOCK_STREAM, 0); + ok &= client >= 0 && connect(client, (struct sockaddr *)&address, sizeof address) == 0; + ok &= send_all(client, "hello") && expect(client, "world"); + int accepted = accept_within(listener); + ok &= accepted >= 0 && expect(accepted, "parent"); + _exit(ok ? 0 : 1); + } + if (child < 0) { + close(listener); + unlink(SOCKET_PATH); + return 1; + } + int accepted = accept_within(listener); + int failures = accepted < 0 || !expect(accepted, "hello") || !send_all(accepted, "world"); + int client = socket(AF_UNIX, SOCK_STREAM, 0); + failures += client < 0 || connect(client, (struct sockaddr *)&address, sizeof address) != 0 || + !send_all(client, "parent"); + failures = finish(child, failures, 0); + close(client); + close(accepted); + close(listener); + unlink(SOCKET_PATH); + return failures; +} + +// An exec'd child keeps an inherited socket but not a close-on-exec one. +static int exec_child(const char *child_path) { + int pair[2]; + int hidden[2]; + if (socketpair(AF_UNIX, SOCK_STREAM, 0, pair) != 0 || + socketpair(AF_UNIX, SOCK_STREAM | SOCK_CLOEXEC, 0, hidden) != 0) { + return 1; + } + fflush(stdout); + pid_t child = fork(); + if (child == 0) { + close(pair[0]); + char socket_fd[16]; + char hidden_fd[16]; + snprintf(socket_fd, sizeof socket_fd, "%d", pair[1]); + snprintf(hidden_fd, sizeof hidden_fd, "%d", hidden[0]); + char *const argv[] = {(char *)child_path, "unix-fd", socket_fd, hidden_fd, NULL}; + execv(child_path, argv); + _exit(99); + } + close(pair[1]); + close(hidden[0]); + close(hidden[1]); + if (child < 0) { + close(pair[0]); + return 1; + } + int failures = !expect(pair[0], "exec"); + failures += !send_all(pair[0], "reply"); + failures = finish(child, failures, 42); + close(pair[0]); + return failures; +} + +int main(int argc, char **argv) { + if (argc < 2) { + return 2; + } + fflush(stdout); + int stream = stream_pair(); + int datagram = datagram_pair(); + int path = path_listener(); + int exec = exec_child(argv[1]); + printf("unix-fork stream=%d dgram=%d path=%d exec=%d\n", stream, datagram, path, exec); + return stream + datagram + path + exec; +} diff --git a/litebox_runner_linux_userland/tests/run.rs b/litebox_runner_linux_userland/tests/run.rs index c7d2524e7e..ef62c69237 100644 --- a/litebox_runner_linux_userland/tests/run.rs +++ b/litebox_runner_linux_userland/tests/run.rs @@ -63,6 +63,7 @@ const DEDICATED_C_TESTS: &[&str] = &[ "async_x16.c", "fork_parent.c", "fork_threads_parent.c", + "fork_unix_parent.c", "gate_signals.c", "sigreturn.c", "sigreturn_simd.c", @@ -1013,6 +1014,47 @@ fn fork_pauses_sibling_threads() { assert_eq!(numeric_field(line, "slept="), 1); } +#[cfg(all(target_arch = "x86_64", target_os = "linux"))] +#[test] +fn forked_and_exec_children_share_unix_sockets() { + let parent = common::compile( + "./tests/fork_unix_parent.c", + "fork_unix_parent", + true, + false, + ); + let child = common::compile("./tests/vfork_exec_child.c", "fork_unix_child", true, false); + let child_guest_path = std::path::absolute(&child).unwrap(); + let mut runner = Runner::new(&parent, "fork_unix_parent"); + runner + .allow_process_duplication() + .arg(&child_guest_path) + .with_fs_path(|root| { + let destination = root.join(child_guest_path.strip_prefix("/").unwrap()); + assert!(common::rewrite_with_cache(&child, &destination, &[])); + }); + + let output = String::from_utf8(runner.output()).unwrap(); + let line = |prefix: &str| { + output + .lines() + .find(|line| line.starts_with(prefix)) + .unwrap_or_else(|| panic!("missing {prefix:?} output in {output:?}")) + }; + let parent_line = line("unix-fork "); + for field in ["stream=", "dgram=", "path=", "exec="] { + assert_eq!( + numeric_field(parent_line, field), + 0, + "{field} in {output:?}" + ); + } + let child_line = line("child-unix "); + for field in ["stream=", "sent=", "replied=", "hidden_closed="] { + assert_eq!(numeric_field(child_line, field), 1, "{field} in {output:?}"); + } +} + #[cfg(all(target_arch = "x86_64", target_os = "linux"))] #[test] fn fork_requires_process_duplication() { diff --git a/litebox_runner_linux_userland/tests/vfork_exec_child.c b/litebox_runner_linux_userland/tests/vfork_exec_child.c index 1b98e0eb40..2391eb54f0 100644 --- a/litebox_runner_linux_userland/tests/vfork_exec_child.c +++ b/litebox_runner_linux_userland/tests/vfork_exec_child.c @@ -9,6 +9,7 @@ #include #include #include +#include #include #include #include @@ -123,6 +124,21 @@ int main(int argc, char **argv) { printf("child-fds stdin_setfl=%d hidden_closed=%d\n", stdin_setfl, hidden_closed); return 42; } + if (strcmp(marker, "unix-fd") == 0 && argc > 3) { + int socket_fd = atoi(argv[2]); + int hidden = atoi(argv[3]); + int type = 0; + socklen_t length = sizeof type; + int is_stream = getsockopt(socket_fd, SOL_SOCKET, SO_TYPE, &type, &length) == 0 && + type == SOCK_STREAM; + int sent = write(socket_fd, "exec", 4) == 4; + char reply[8] = {0}; + int replied = read(socket_fd, reply, sizeof reply - 1) == 5 && strcmp(reply, "reply") == 0; + int hidden_closed = fcntl(hidden, F_GETFD) == -1 && errno == EBADF; + printf("child-unix stream=%d sent=%d replied=%d hidden_closed=%d\n", is_stream, sent, + replied, hidden_closed); + return 42; + } if (strcmp(marker, "stdout-mode") == 0) { printf("child-stdout wronly=%d\n", (fcntl(1, F_GETFL) & O_ACCMODE) == O_WRONLY); marker = "open-fds"; diff --git a/litebox_shim_linux/src/channel.rs b/litebox_shim_linux/src/channel.rs deleted file mode 100644 index a370eef0ef..0000000000 --- a/litebox_shim_linux/src/channel.rs +++ /dev/null @@ -1,334 +0,0 @@ -// Copyright (c) Microsoft Corporation. -// Licensed under the MIT license. - -use core::sync::atomic::{AtomicBool, Ordering}; - -use alloc::sync::{Arc, Weak}; -use litebox::{ - event::{Events, observer::Observer, polling::Pollee}, - sync::{Mutex, RawSyncPrimitivesProvider}, -}; -use litebox_common_linux::errno::Errno; -use litebox_platform::time::TimeProvider; -use ringbuf::traits::{Consumer as _, Observer as _, Producer as _}; - -use crate::ShimPlatform; - -macro_rules! common_functions_for_channel { - () => { - pub(crate) fn is_shutdown(&self) -> bool { - self.endpoint.is_shutdown() - } - - /// Shuts the endpoint down. Returns `true` only on the call that - /// effected the transition (idempotent thereafter — not a fallibility - /// signal). The first transition also wakes the peer's pollee so a - /// peer blocked in send/recv unblocks immediately. - pub(crate) fn shutdown(&self) -> bool { - if self.endpoint.shutdown() { - if let Some(peer) = self.peer.upgrade() { - peer.pollee.notify_observers(litebox::event::Events::HUP); - } - true - } else { - false - } - } - - /// Has the peer (i.e., other end) been shut down? - pub(crate) fn is_peer_shutdown(&self) -> bool { - if let Some(peer) = self.peer.upgrade() { - peer.is_shutdown() - } else { - true - } - } - }; -} - -struct EndPointer { - rb: Mutex, - pollee: Arc>, - is_shutdown: AtomicBool, -} - -impl EndPointer { - fn new(rb: T, pollee: Arc>) -> Self { - Self { - rb: Mutex::new(rb), - pollee, - is_shutdown: AtomicBool::new(false), - } - } - - fn is_shutdown(&self) -> bool { - self.is_shutdown.load(Ordering::Acquire) - } - - /// Returns `true` on the call that affected the transition so callers can - /// gate one-shot side-effects (e.g. peer wake-ups); idempotent thereafter. - /// The boolean reports newness, not fallibility — the state is always shut - /// down after this call. - fn shutdown(&self) -> bool { - !self.is_shutdown.swap(true, Ordering::Release) - } -} - -pub(crate) struct ReadEnd { - endpoint: alloc::sync::Arc>>, - peer: alloc::sync::Weak>>, -} - -impl ReadEnd { - fn update_pollee(&self) { - if let Some(peer) = self.peer.upgrade() { - peer.pollee.notify_observers(litebox::event::Events::OUT); - } - } - - pub(crate) fn is_empty(&self) -> bool { - self.endpoint.rb.lock().is_empty() - } - - /// Peeks at the first item in the channel and conditionally consumes it. - /// - /// This method allows examining and potentially modifying the first item in the - /// channel through a closure. The closure decides whether to consume the item - /// by returning a boolean in its result tuple. - pub(crate) fn peek_and_consume_one( - &self, - mut f: impl FnMut(&mut T) -> Result<(bool, R), Errno>, - ) -> Result { - // Linux preserves bytes already queued when the read side is shut down - // (via shutdown(SHUT_RD) or peer close), so consult the buffer before - // returning ESHUTDOWN; the caller observes EOF only once the queue drains. - let is_shutdown = self.is_shutdown() || self.is_peer_shutdown(); - let mut guard = self.endpoint.rb.lock(); - if let Some(item) = guard.first_mut() { - let (should_consume, ret) = f(item)?; - if should_consume { - guard - .try_pop() - .expect("Guaranteed to have an element to consume"); - self.update_pollee(); - } - return Ok(ret); - } - if is_shutdown { - return Err(Errno::ESHUTDOWN); - } - - Err(Errno::EAGAIN) - } - - pub(crate) fn for_each_queued(&self, mut f: impl FnMut(&T) -> bool) -> Result<(), Errno> { - let is_shutdown = self.is_shutdown() || self.is_peer_shutdown(); - let guard = self.endpoint.rb.lock(); - if guard.is_empty() { - return if is_shutdown { - Err(Errno::ESHUTDOWN) - } else { - Err(Errno::EAGAIN) - }; - } - for item in guard.iter() { - if !f(item) { - break; - } - } - Ok(()) - } - - common_functions_for_channel!(); -} - -pub(crate) struct WriteEnd { - endpoint: alloc::sync::Arc>>, - peer: alloc::sync::Weak>>, -} - -impl Clone for WriteEnd { - fn clone(&self) -> Self { - Self { - endpoint: self.endpoint.clone(), - peer: self.peer.clone(), - } - } -} - -impl WriteEnd { - pub(crate) fn try_write_one(&self, elem: T) -> Result<(), (T, Errno)> { - if self.is_shutdown() || self.is_peer_shutdown() { - return Err((elem, Errno::EPIPE)); - } - - let ret = self.endpoint.rb.lock().try_push(elem); - match ret { - Ok(()) => { - if let Some(peer) = self.peer.upgrade() { - peer.pollee.notify_observers(litebox::event::Events::IN); - } - Ok(()) - } - Err(e) => Err((e, Errno::EAGAIN)), - } - } - - pub(crate) fn is_full(&self) -> bool { - self.endpoint.rb.lock().is_full() - } - - pub(crate) fn is_pair(&self, reader: &ReadEnd) -> bool { - if let Some(peer) = self.peer.upgrade() { - Arc::ptr_eq(&peer, &reader.endpoint) - } else { - false - } - } - - pub(crate) fn register_observer(&self, observer: Weak>, filter: Events) { - self.endpoint.pollee.register_observer(observer, filter); - } - - common_functions_for_channel!(); -} - -pub(crate) struct Channel { - writer: WriteEnd, - reader: ReadEnd, -} - -impl Channel { - pub(crate) fn new( - capacity: usize, - writer_pollee: Arc>, - reader_pollee: Arc>, - ) -> Self { - use ringbuf::traits::Split as _; - let rb: ringbuf::HeapRb = ringbuf::HeapRb::new(capacity); - let (rb_prod, rb_cons) = rb.split(); - - let mut writer = WriteEnd { - endpoint: Arc::new(EndPointer::new(rb_prod, writer_pollee)), - peer: alloc::sync::Weak::new(), - }; - let mut reader = ReadEnd { - endpoint: Arc::new(EndPointer::new(rb_cons, reader_pollee)), - peer: alloc::sync::Weak::new(), - }; - - writer.peer = Arc::downgrade(&reader.endpoint); - reader.peer = Arc::downgrade(&writer.endpoint); - - Self { writer, reader } - } - - /// Turn the channel into a pair of its read and write ends. - pub(crate) fn split(self) -> (WriteEnd, ReadEnd) { - let Channel { writer, reader } = self; - (writer, reader) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::syscalls::tests::TestPlatform; - use core::sync::atomic::{AtomicBool, Ordering}; - use litebox::event::observer::Observer; - - fn split_pair() -> (WriteEnd, ReadEnd) { - Channel::::new(4, Arc::new(Pollee::new()), Arc::new(Pollee::new())).split() - } - - /// Test observer that flips a flag the first time it is notified. - struct FlagOnNotify(Arc); - impl Observer for FlagOnNotify { - fn on_events(&self, _events: &Events) { - self.0.store(true, Ordering::Release); - } - } - - #[test] - fn peek_and_consume_one_drains_queue_after_self_shutdown() { - let (writer, reader) = split_pair::(); - writer.try_write_one(42).unwrap(); - reader.shutdown(); - // Queued bytes must remain readable after shutdown(SHUT_RD): we should - // get the 42 first, ESHUTDOWN only once the buffer is empty. - let got = reader - .peek_and_consume_one(|x| Ok((true, *x))) - .expect("queued item must be returned even after self shutdown"); - assert_eq!(got, 42); - let err = reader - .peek_and_consume_one(|x: &mut u32| Ok((true, *x))) - .unwrap_err(); - assert_eq!(err, Errno::ESHUTDOWN); - } - - #[test] - fn peek_and_consume_one_drains_queue_after_peer_shutdown() { - let (writer, reader) = split_pair::(); - writer.try_write_one(7).unwrap(); - writer.shutdown(); - let got = reader - .peek_and_consume_one(|x| Ok((true, *x))) - .expect("queued item must be returned even after peer shutdown"); - assert_eq!(got, 7); - let err = reader - .peek_and_consume_one(|x: &mut u32| Ok((true, *x))) - .unwrap_err(); - assert_eq!(err, Errno::ESHUTDOWN); - } - - #[test] - fn peek_and_consume_one_returns_eagain_when_empty_and_alive() { - let (_writer, reader) = split_pair::(); - let err = reader - .peek_and_consume_one(|x: &mut u32| Ok((true, *x))) - .unwrap_err(); - assert_eq!(err, Errno::EAGAIN); - } - - #[test] - fn try_write_one_returns_epipe_after_self_shutdown() { - let (writer, _reader) = split_pair::(); - writer.shutdown(); - let (_val, err) = writer.try_write_one(1).unwrap_err(); - assert_eq!(err, Errno::EPIPE); - } - - #[test] - fn try_write_one_returns_epipe_after_peer_shutdown() { - let (writer, reader) = split_pair::(); - reader.shutdown(); - let (_val, err) = writer.try_write_one(1).unwrap_err(); - assert_eq!(err, Errno::EPIPE); - } - - /// Regression: `shutdown()` must wake observers on the peer's pollee so a peer blocked - /// in send/recv notices the new state without waiting for an unrelated event. HUP is in - /// `Events::ALWAYS_POLLED`, so any observer (even one registered with a different mask) - /// must be notified. - #[test] - fn shutdown_notifies_peer_pollee_hup() { - let writer_pollee = Arc::new(Pollee::new()); - let reader_pollee = Arc::new(Pollee::new()); - let (_writer, reader) = - Channel::::new(4, writer_pollee.clone(), reader_pollee).split(); - let flag = Arc::new(AtomicBool::new(false)); - let observer: Arc = Arc::new(FlagOnNotify(flag.clone())); - // The peer of `reader` is the writer's endpoint, whose pollee is `writer_pollee`; - // register the observer there to detect that `reader.shutdown()` reaches it. - writer_pollee.register_observer( - Arc::downgrade(&observer) as Weak>, - Events::OUT, - ); - assert!(!flag.load(Ordering::Acquire), "observer must start cleared"); - reader.shutdown(); - assert!( - flag.load(Ordering::Acquire), - "shutdown(ReadEnd) must wake peer pollee observers" - ); - } -} diff --git a/litebox_shim_linux/src/lib.rs b/litebox_shim_linux/src/lib.rs index 21442b8f98..efac8449cc 100644 --- a/litebox_shim_linux/src/lib.rs +++ b/litebox_shim_linux/src/lib.rs @@ -62,7 +62,6 @@ macro_rules! log_unsupported { }; } -pub(crate) mod channel; pub mod loader; pub(crate) mod stdio; pub mod syscalls; @@ -409,7 +408,6 @@ impl LinuxShimBuilder { boot_time: self.platform.now(), process_id: self.process_id, litebox, - unix_addr_table: litebox::sync::RwLock::new(syscalls::unix::UnixAddrTable::new()), elf_patch_cache: litebox::sync::Mutex::new(alloc::collections::BTreeMap::new()), }); LinuxShim(global) @@ -1538,8 +1536,6 @@ struct GlobalState { boot_time: ::Instant, /// Process ID assigned to this shim. process_id: i32, - /// UNIX domain socket address table - unix_addr_table: litebox::sync::RwLock>, /// Per-process collection of ELF patching state for runtime syscall rewriting. elf_patch_cache: litebox::sync::Mutex, } diff --git a/litebox_shim_linux/src/syscalls/file.rs b/litebox_shim_linux/src/syscalls/file.rs index 26024d0999..f56b048fd0 100644 --- a/litebox_shim_linux/src/syscalls/file.rs +++ b/litebox_shim_linux/src/syscalls/file.rs @@ -171,7 +171,7 @@ fn set_file_status_flags( /// Returns the `F_GETFL` flags for the access mode and status flags `status` of an open file /// description. -fn linux_status_flags(status: FileStatusFlags) -> Result { +pub(super) fn linux_status_flags(status: FileStatusFlags) -> Result { let mut flags = match status.access { FileAccessMode::ReadOnly => OFlags::RDONLY, FileAccessMode::WriteOnly => OFlags::WRONLY, @@ -321,6 +321,13 @@ impl FilesState { .map_err(|_| crate::loader::elf::ElfLoaderError::InvalidInheritedFds) })?; } + InheritedFdKind::LocalSocket => { + self.install_inherited_fd(litebox, first_fd, raw_fd, || { + global + .adopt_inherited_unix_socket(handle) + .map_err(|_| crate::loader::elf::ElfLoaderError::InvalidInheritedFds) + })?; + } } descriptions.entry(handle).or_insert(raw_fd); next_fd = raw_fd + 1; @@ -600,7 +607,7 @@ impl Task { /// them is: /// /// - a network socket: the broker does not duplicate socket objects; - /// - an eventfd, epoll, or Unix socket descriptor: the object is local to this runner; + /// - an eventfd or epoll descriptor: the object is local to this runner; /// - a timerfd: the broker does not duplicate timer objects; /// - a directory with a position, set by `getdents64` or `lseek`: the position is local to this /// runner; or @@ -640,6 +647,14 @@ impl Task { self.global.inherited_linux_pipe_kind(&pipe)?, InheritableFd::Pipe(pipe), ), + AnyTypedFd::Unix(socket) => ( + InheritedFdKind::LocalSocket, + InheritableFd::LocalSocket( + self.global.with_unix_socket(&socket, |socket| { + Ok(alloc::sync::Arc::clone(socket.local_socket())) + })?, + ), + ), _ => return Err(Errno::EAGAIN), }; next_fd = raw_fd + 1; @@ -931,13 +946,15 @@ impl Task { .entry_handle(fd) .ok_or(Errno::EBADF)?; handle.with_entry(|file| { + let buf = &mut buf.borrow_mut(); + // A datagram longer than `buf` reports its full length. file.recvfrom( &self.wait_cx(), - &mut buf.borrow_mut(), + buf, litebox_common_linux::ReceiveFlags::empty(), None, - None, ) + .map(|size| size.min(buf.len())) }) }, |fd| { @@ -1478,117 +1495,154 @@ impl Task { }) } + /// Reads one datagram into `iovs` if `fd` is a datagram socket, which a + /// series of reads would split. fn try_read_datagram_from_iovec( &self, fd: &AnyTypedFd, iovs: &[IoReadVec], ) -> Result, Errno> { - let AnyTypedFd::Network(fd) = fd else { - return Ok(None); - }; - let socket = self.global.pin_socket(fd)?; - if socket.is_broker_datagram() { - self.read_datagram_from_iovec(&socket, iovs).map(Some) - } else { - Ok(None) - } - } - - fn read_datagram_from_iovec( - &self, - socket: &super::net::InetSocketPin<'_, Platform>, - iovs: &[IoReadVec], - ) -> Result { - check_iov_lens(iovs.iter().map(|iov| iov.iov_len))?; - let capacity = iovs.iter().map(|iov| iov.iov_len).sum::(); - if capacity == 0 { - return Ok(0); - } - let mut buffer = alloc::vec::Vec::new(); - let capacity = capacity.min(litebox::net::SOCKET_RECEIVE_OPERATION_SIZE); - buffer - .try_reserve_exact(capacity) - .map_err(|_| Errno::ENOMEM)?; - buffer.resize(capacity, 0); - let received = self - .global - .receive_from_socket( - &self.wait_cx(), - socket, - &mut buffer, - litebox_common_linux::ReceiveFlags::empty(), - super::net::ReceiveContext::new(None, false), - None, - )? - .min(buffer.len()); - let mut copied = 0; - for iov in iovs { - if copied == received { - break; - } - let length = (received - copied).min(iov.iov_len); - if iov - .iov_base - .copy_from_slice::(0, &buffer[copied..copied + length]) - .is_none() - { - return Err(Errno::EFAULT); + let receive_flags = litebox_common_linux::ReceiveFlags::empty(); + match fd { + AnyTypedFd::Network(fd) => { + let socket = self.global.pin_socket(fd)?; + if !socket.is_broker_datagram() { + return Ok(None); + } + read_datagram_from_iovec::(iovs, |buffer| { + self.global.receive_from_socket( + &self.wait_cx(), + &socket, + buffer, + receive_flags, + super::net::ReceiveContext::new(None, false), + None, + ) + }) + .map(Some) } - copied += length; + AnyTypedFd::Unix(fd) => self.global.with_unix_socket(fd, |socket| { + if socket.is_stream() { + return Ok(None); + } + read_datagram_from_iovec::(iovs, |buffer| { + socket.recvfrom(&self.wait_cx(), buffer, receive_flags, None) + }) + .map(Some) + }), + _ => Ok(None), } - Ok(received) } + /// Writes `iovs` as one datagram if `fd` is a datagram socket, which a + /// series of writes would split. fn try_write_datagram_to_iovec( &self, fd: &AnyTypedFd, iovs: &[IoWriteVec], ) -> Result, Errno> { - let AnyTypedFd::Network(fd) = fd else { - return Ok(None); - }; - let socket = self.global.pin_socket(fd)?; - if socket.is_broker_datagram() { - self.write_datagram_to_iovec(&socket, iovs).map(Some) - } else { - Ok(None) + let send_flags = litebox_common_linux::SendFlags::empty(); + match fd { + AnyTypedFd::Network(fd) => { + let socket = self.global.pin_socket(fd)?; + if !socket.is_broker_datagram() { + return Ok(None); + } + write_datagram_to_iovec::( + iovs, + litebox::net::MAX_UDP_DATAGRAM_SIZE, + |buffer| { + self.global.send_to_pinned_socket( + &self.wait_cx(), + &socket, + buffer, + send_flags, + ) + }, + ) + .map(Some) + } + AnyTypedFd::Unix(fd) => self.global.with_unix_socket(fd, |socket| { + if socket.is_stream() { + return Ok(None); + } + write_datagram_to_iovec::( + iovs, + litebox_broker_protocol::local_socket::LOCAL_SOCKET_BUFFER_SIZE as usize, + |buffer| socket.sendto(self, buffer, send_flags, None), + ) + .map(Some) + }), + _ => Ok(None), } } +} - fn write_datagram_to_iovec( - &self, - socket: &super::net::InetSocketPin<'_, Platform>, - iovs: &[IoWriteVec], - ) -> Result { - check_iov_lens(iovs.iter().map(|iov| iov.iov_len))?; - let length = iovs.iter().map(|iov| iov.iov_len).sum::(); - if length == 0 { - return Ok(0); - } - if length > litebox::net::MAX_UDP_DATAGRAM_SIZE { - return Err(Errno::EMSGSIZE); - } - let mut buffer = alloc::vec::Vec::new(); - buffer - .try_reserve_exact(length) - .map_err(|_| Errno::ENOMEM)?; - for iov in iovs { - for offset in 0..iov.iov_len { - let offset = isize::try_from(offset).map_err(|_| Errno::EFAULT)?; - buffer.push( - iov.iov_base - .read_at_offset::(offset) - .ok_or(Errno::EFAULT)?, - ); - } - } - self.global.send_to_pinned_socket( - &self.wait_cx(), - socket, - &buffer, - litebox_common_linux::SendFlags::empty(), - ) +/// Receives one datagram with `receive` and scatters it into `iovs`. +fn read_datagram_from_iovec( + iovs: &[IoReadVec], + receive: impl FnOnce(&mut [u8]) -> Result, +) -> Result { + check_iov_lens(iovs.iter().map(|iov| iov.iov_len))?; + let capacity = iovs.iter().map(|iov| iov.iov_len).sum::(); + if capacity == 0 { + return Ok(0); + } + let mut buffer = alloc::vec::Vec::new(); + let capacity = capacity.min(litebox::net::SOCKET_RECEIVE_OPERATION_SIZE); + buffer + .try_reserve_exact(capacity) + .map_err(|_| Errno::ENOMEM)?; + buffer.resize(capacity, 0); + let received = receive(&mut buffer)?.min(buffer.len()); + let mut copied = 0; + for iov in iovs { + if copied == received { + break; + } + let length = (received - copied).min(iov.iov_len); + if iov + .iov_base + .copy_from_slice::(0, &buffer[copied..copied + length]) + .is_none() + { + return Err(Errno::EFAULT); + } + copied += length; + } + Ok(received) +} + +/// Gathers `iovs` into one datagram of at most `max_length` bytes and sends it +/// with `send`. +fn write_datagram_to_iovec( + iovs: &[IoWriteVec], + max_length: usize, + send: impl FnOnce(&[u8]) -> Result, +) -> Result { + check_iov_lens(iovs.iter().map(|iov| iov.iov_len))?; + let length = iovs.iter().map(|iov| iov.iov_len).sum::(); + if length == 0 { + return Ok(0); + } + if length > max_length { + return Err(Errno::EMSGSIZE); + } + let mut buffer = alloc::vec::Vec::new(); + buffer + .try_reserve_exact(length) + .map_err(|_| Errno::ENOMEM)?; + for iov in iovs { + for offset in 0..iov.iov_len { + let offset = isize::try_from(offset).map_err(|_| Errno::EFAULT)?; + buffer.push( + iov.iov_base + .read_at_offset::(offset) + .ok_or(Errno::EFAULT)?, + ); + } } + send(&buffer) } /// Linux's `IOV_MAX` / `UIO_MAXIOV`: the kernel rejects iovec counts above this @@ -2155,7 +2209,10 @@ impl Task { |fd| linux_status_flags(self.global.pipes.get_status_flags(fd)?), |fd| getfl_from_handle!(fd), |fd| getfl_from_handle!(fd), - |fd| getfl_from_handle!(fd), + |fd| { + self.global + .with_unix_socket(fd, super::unix::UnixSocket::get_status) + }, |fd| getfl_from_handle!(fd), )? .bits()) @@ -2236,8 +2293,12 @@ impl Task { }, |_fd| todo!("epoll"), |fd| { - toggle_flags!(fd); - Ok(()) + if flags.intersects(OFlags::DIRECT | OFlags::NOATIME) { + log_unsupported!("unsupported flags"); + } + self.global.with_unix_socket(fd, |socket| { + socket.set_status_flags(setfl_mask, flags) + }) }, |fd| { toggle_flags!(fd); @@ -2531,7 +2592,16 @@ impl Task { }, |fd| set_nonblock_on_entry!(fd), |fd| set_nonblock_on_entry!(fd), - |fd| set_nonblock_on_entry!(fd), + |fd| { + let flags = if val != 0 { + OFlags::NONBLOCK + } else { + OFlags::empty() + }; + self.global.with_unix_socket(fd, |socket| { + socket.set_status_flags(OFlags::NONBLOCK, flags) + }) + }, |fd| set_nonblock_on_entry!(fd), )?; Ok(0) diff --git a/litebox_shim_linux/src/syscalls/net.rs b/litebox_shim_linux/src/syscalls/net.rs index f791505957..6992d19005 100644 --- a/litebox_shim_linux/src/syscalls/net.rs +++ b/litebox_shim_linux/src/syscalls/net.rs @@ -1422,7 +1422,7 @@ impl Task { } AddressFamily::UNIX => { let _ = UnixProtocol::try_from(protocol).map_err(|_| Errno::EPROTONOSUPPORT)?; - let socket = UnixSocket::new(ty, flags).ok_or(Errno::ESOCKTNOSUPPORT)?; + let socket = UnixSocket::new(&self.global.litebox, ty, flags)?; let typed = self .global .litebox @@ -1481,7 +1481,7 @@ impl Task { AddressFamily::UNIX => { let _ = UnixProtocol::try_from(protocol).map_err(|_| Errno::EPROTONOSUPPORT)?; let (sock1, sock2) = - UnixSocket::new_connected_pair(ty, flags).ok_or(Errno::ESOCKTNOSUPPORT)?; + UnixSocket::new_connected_pair(&self.global.litebox, ty, flags)?; let files = self.files.borrow(); let mut dt = self.global.litebox.descriptor_table_mut(); let typed1 = @@ -1859,7 +1859,7 @@ impl Task { &self.global, socket, |fd| self.global.listen(fd, backlog), - |file| file.listen(backlog, &self.global), + |file| file.listen(backlog), ) } @@ -2168,19 +2168,10 @@ impl Task { .transpose()?; Ok((chunk_waitall, is_stream, deadline)) } - ReceiveSocket::Unix(handle) => handle.with_entry(|entry| { - let deadline = entry - .recv_timeout() - .map(|timeout| { - self.global - .platform - .now() - .checked_add(timeout) - .ok_or(Errno::EINVAL) - }) - .transpose()?; - Ok((false, entry.is_stream(), deadline)) - }), + // The local socket applies its own receive timeout to each receive. + ReceiveSocket::Unix(handle) => { + Ok((false, handle.with_entry(UnixSocket::is_stream), None)) + } } } @@ -2263,16 +2254,10 @@ impl Task { } ReceiveSocket::Unix(handle) => handle.with_entry(|entry| { let mut addr = None; - let timeout = deadline.map(|deadline| { - deadline - .checked_duration_since(&self.global.platform.now()) - .unwrap_or(core::time::Duration::ZERO) - }); let size = entry.recvfrom( &self.wait_cx(), buf, flags, - timeout, if want_source { Some(&mut addr) } else { None }, )?; let src_addr = addr.map(SocketAddress::Unix); @@ -2707,7 +2692,7 @@ impl Task { .map(SocketAddress::Inet) .map_err(Errno::from) }, - |unix| Ok(SocketAddress::Unix(unix.get_local_addr())), + |unix| unix.get_local_addr().map(SocketAddress::Unix), ) } @@ -2734,11 +2719,7 @@ impl Task { .map(SocketAddress::Inet) .map_err(Errno::from) }, - |file| { - file.get_peer_addr() - .ok_or(Errno::ENOTCONN) - .map(SocketAddress::Unix) - }, + |file| file.get_peer_addr().map(SocketAddress::Unix), ) } @@ -2774,8 +2755,7 @@ impl Task { }, |file| { let how = ShutdownHow::try_from(how).map_err(|_| Errno::EINVAL)?; - file.shutdown(how); - Ok(()) + file.shutdown(how) }, ) } @@ -3963,7 +3943,7 @@ mod unix_tests { &client_fd, SocketAddress::Unix(UnixSocketAddr::Path(addr.to_string())), ); - assert_eq!(result.unwrap_err(), Errno::ECONNREFUSED); + assert_eq!(result.unwrap_err(), Errno::ENOENT); close_socket(&task, raw_client_fd); let raw_server_fd = create_unix_server_socket( @@ -4047,11 +4027,15 @@ mod unix_tests { }, ); let client_fd = typed_socket(&task, raw_client_fd); - if is_nonblocking { - ppoll(&task, raw_server_fd, Events::OUT); + // A listener never polls writable, so a non-blocking connect + // retries while the backlog is full. + loop { + match task.do_connect(&client_fd, SocketAddress::Unix(client_addr.clone())) { + Ok(()) => break, + Err(Errno::EAGAIN) if is_nonblocking => std::thread::yield_now(), + Err(error) => panic!("connect failed: {error:?}"), + } } - task.do_connect(&client_fd, SocketAddress::Unix(client_addr.clone())) - .unwrap(); client_fds.push((raw_client_fd, client_fd)); } @@ -4304,6 +4288,66 @@ mod unix_tests { unix_socketpair_bidirectional(SockType::Datagram, true); } + #[test] + fn unix_datagram_reads_and_vectors_keep_boundaries() { + use litebox_common_linux::{IoReadVec, IoWriteVec}; + + let task = init_platform(); + let (sender, receiver) = task + .do_socketpair( + AddressFamily::UNIX, + SockType::Datagram, + SockFlags::empty(), + 0, + ) + .unwrap(); + let sender_fd = i32::try_from(sender).unwrap(); + let receiver_fd = i32::try_from(receiver).unwrap(); + + let (first, second) = (b"ab", b"cdef"); + let write_iovs = [ + IoWriteVec { + iov_base: UserPtr::from_usize(first.as_ptr().expose_provenance()), + iov_len: first.len(), + }, + IoWriteVec { + iov_base: UserPtr::from_usize(second.as_ptr().expose_provenance()), + iov_len: second.len(), + }, + ]; + let write_iovs_ptr = UserPtr::from_usize(write_iovs.as_ptr().expose_provenance()); + assert_eq!(task.sys_writev(sender_fd, write_iovs_ptr, 2), Ok(6)); + assert_eq!(task.sys_writev(sender_fd, write_iovs_ptr, 2), Ok(6)); + + let mut short = [0; 3]; + assert_eq!(task.sys_read(receiver_fd, &mut short, None), Ok(3)); + assert_eq!(&short, b"abc"); + + let (mut head, mut tail) = ([0; 1], [0; 8]); + let read_iovs = [ + IoReadVec { + iov_base: UserPtrMut::from_usize(head.as_mut_ptr().expose_provenance()), + iov_len: head.len(), + }, + IoReadVec { + iov_base: UserPtrMut::from_usize(tail.as_mut_ptr().expose_provenance()), + iov_len: tail.len(), + }, + ]; + let read_iovs_ptr = UserPtr::from_usize(read_iovs.as_ptr().expose_provenance()); + assert_eq!(task.sys_readv(receiver_fd, read_iovs_ptr, 2), Ok(6)); + assert_eq!(&head, b"a"); + assert_eq!(&tail[..5], b"bcdef"); + + let receiver_typed = typed_socket(&task, receiver); + assert_eq!( + task.do_recvfrom(&receiver_typed, &mut short, ReceiveFlags::DONTWAIT, None), + Err(Errno::EAGAIN) + ); + close_socket(&task, sender); + close_socket(&task, receiver); + } + #[test] fn pinned_receive_does_not_follow_dup2_replacement() { let task = init_platform(); diff --git a/litebox_shim_linux/src/syscalls/test_broker.rs b/litebox_shim_linux/src/syscalls/test_broker.rs index 612da8828d..bf9f387f59 100644 --- a/litebox_shim_linux/src/syscalls/test_broker.rs +++ b/litebox_shim_linux/src/syscalls/test_broker.rs @@ -28,7 +28,7 @@ use litebox_broker_transport::channel::LocalCallChannel; use crate::syscalls::tests::TestPlatform; -pub(crate) const MAX_TEST_BROKER_REFERENCES: usize = 16; +pub(crate) const MAX_TEST_BROKER_REFERENCES: usize = 64; static BROKER: OnceLock = OnceLock::new(); diff --git a/litebox_shim_linux/src/syscalls/unix.rs b/litebox_shim_linux/src/syscalls/unix.rs index 6636959599..c46fab4907 100644 --- a/litebox_shim_linux/src/syscalls/unix.rs +++ b/litebox_shim_linux/src/syscalls/unix.rs @@ -1,41 +1,42 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT license. -//! Unix domain socket implementation for the Linux shim layer. - -use core::{ - sync::atomic::{AtomicBool, AtomicU32, Ordering}, - time::Duration, -}; +//! Unix domain sockets for the Linux shim layer. +//! +//! Each socket is a broker-owned [`LocalSocket`], so processes that share a socket through +//! inheritance see one socket. This module translates between Linux socket calls and the +//! broker's guest-neutral local socket operations. use alloc::{ - collections::{btree_map::BTreeMap, vec_deque::VecDeque}, - string::String, + string::{String, ToString as _}, sync::{Arc, Weak}, vec::Vec, }; use litebox::{ - event::{ - Events, IOPollable, - polling::{Pollee, TryOpError}, - wait::WaitContext, + event::{Events, IOPollable, observer::Observer, wait::WaitContext}, + fd::{FdEnabledSubsystem, FdEnabledSubsystemEntry, TypedFd}, + local_sockets::LocalSocket, + process::ProcessError, +}; +use litebox_broker_protocol::{ + ObjectHandle, + fs::{FileMode as Mode, FileOpenFlags, FileUser}, + local_socket::{ + LOCAL_SOCKET_BUFFER_SIZE, LocalSocketAddress, LocalSocketName, LocalSocketOption, }, - fd::{FdEnabledSubsystem, FdEnabledSubsystemEntry}, - fs::errors::OpenError, - sync::{Mutex, RwLock}, - utils::TruncateExt as _, + socket::{ShutdownMode, SocketType}, }; -use litebox_broker_protocol::fs::{FileAccessMode, FileMode as Mode, FileOpenFlags}; use litebox_common_linux::{ IpOption, OFlags, ReceiveFlags, SendFlags, ShutdownHow, SockFlags, SockType, SocketOption, SocketOptionName, errno::Errno, }; use crate::{ - FileFd, GlobalState, ShimPlatform, Task, UserPtr, UserPtrMut, - channel::{Channel, ReadEnd, WriteEnd}, - syscalls::net::{SocketOptionValue, SocketOptions}, - wait::wait_errno, + GlobalState, ShimPlatform, Task, UserPtr, UserPtrMut, + syscalls::{ + file::{linux_status_flags, status_flags_change}, + net::SocketOptionValue, + }, }; pub(crate) struct UnixSocketSubsystem(core::marker::PhantomData); @@ -66,1363 +67,118 @@ pub(crate) enum UnixSocketAddr { Abstract(Vec), } -/// A bound Unix socket address with associated resources. -/// -/// For path-based sockets, this includes a file descriptor to ensure -/// the socket file remains accessible. The file is automatically closed -/// when this structure is dropped. -enum UnixBoundSocketAddr { - Path((String, FileFd, Arc>)), - Abstract(Vec), -} - -/// Key type for indexing Unix socket addresses in the global address table. -/// -/// This is used internally to track which addresses are currently bound -/// by listening sockets. -#[derive(PartialEq, Eq, Hash, Debug, Ord, PartialOrd)] -pub(crate) enum UnixSocketAddrKey { - // TODO: add inode reference once the file system supports it. - Path(String), - Abstract(Vec), -} - impl UnixSocketAddr { - /// Returns true if this is an unnamed socket address. - fn is_unnamed(&self) -> bool { - matches!(self, UnixSocketAddr::Unnamed) - } - - /// Binds this address to the filesystem or abstract namespace. - /// - /// # Arguments + /// Returns the broker address naming `self`, resolving a relative path against the current + /// working directory of `task`. /// - /// * `task` - The current task context - /// * `is_server` - Whether this is a server socket (creates the file if true) - /// - /// # Errors - /// - /// Returns an error if the address cannot be bound (e.g., file doesn't exist, - /// permission denied). - fn bind( - self, + /// Fails with `EINVAL` for an unnamed address, which names no socket. + fn to_local( + &self, task: &Task, - is_server: bool, - ) -> Result, Errno> { + ) -> Result { match self { + UnixSocketAddr::Unnamed => Err(Errno::EINVAL), + UnixSocketAddr::Abstract(name) => Ok(LocalSocketAddress::Abstract(name.clone())), UnixSocketAddr::Path(path) => { - let flags = if is_server { - // create the socket file if not exists; - // use O_EXCL to ensure exclusive creation - FileOpenFlags::CREATE | FileOpenFlags::EXCLUSIVE - } else { - FileOpenFlags::NONE - }; - // TODO: extend fs to support creating sock file (i.e., with type `InodeType::Socket`) - let file = { - let fs = task.fs.borrow(); - let context = fs.context.read(); - task.global - .litebox - .open_file( - &context, - path.as_str(), - FileAccessMode::ReadWrite, - flags, - Mode::RWXU | Mode::RGRP | Mode::XGRP | Mode::ROTH | Mode::XOTH, - ) - .map_err(|err| match err { - OpenError::AlreadyExists => Errno::EADDRINUSE, - other => Errno::from(other), - })? - }; - Ok(UnixBoundSocketAddr::Path(( - path, - file, - Arc::clone(&task.global.litebox), - ))) - } - UnixSocketAddr::Abstract(data) => { - // TODO: check if the abstract address is already in use - Ok(UnixBoundSocketAddr::Abstract(data)) - } - UnixSocketAddr::Unnamed => todo!("autobind for unnamed unix socket"), - } - } - - /// Converts this address to a key for the global address table. - /// - /// Returns `None` for unnamed addresses, which cannot be looked up. - fn to_key(&self) -> Option { - match self { - Self::Unnamed => None, - Self::Path(path) => Some(UnixSocketAddrKey::Path(path.clone())), - Self::Abstract(addr) => Some(UnixSocketAddrKey::Abstract(addr.clone())), - } - } -} - -impl UnixBoundSocketAddr { - /// Converts this bound address to a key for the global address table. - fn to_key(&self) -> UnixSocketAddrKey { - match self { - Self::Path((path, ..)) => UnixSocketAddrKey::Path(path.clone()), - Self::Abstract(addr) => UnixSocketAddrKey::Abstract(addr.clone()), - } - } -} - -impl Drop for UnixBoundSocketAddr { - fn drop(&mut self) { - match self { - Self::Path((_, file, fs)) => { - let _ = fs.close_file(file); - } - Self::Abstract(_) => {} - } - } -} - -impl From<&UnixBoundSocketAddr> for UnixSocketAddr { - fn from(addr: &UnixBoundSocketAddr) -> Self { - match addr { - UnixBoundSocketAddr::Path((path, ..)) => UnixSocketAddr::Path(path.clone()), - UnixBoundSocketAddr::Abstract(data) => UnixSocketAddr::Abstract(data.clone()), - } - } -} - -/// Represents a Unix stream socket in its initial state. -/// -/// This is the state immediately after socket creation, before the socket -/// has been connected, or put into listening mode. -struct UnixInitStream { - /// Optional bound address for this socket - addr: Option>, - pollee: Pollee, - read_shutdown: AtomicBool, - write_shutdown: AtomicBool, -} - -impl UnixInitStream { - fn new() -> Self { - Self { - addr: None, - pollee: Pollee::new(), - read_shutdown: AtomicBool::new(false), - write_shutdown: AtomicBool::new(false), - } - } - - fn shutdown(&self, how: ShutdownHow) { - if how.is_shutdown_read() && !self.read_shutdown.swap(true, Ordering::Release) { - self.pollee.notify_observers(Events::IN); - } - if how.is_shutdown_write() { - self.write_shutdown.store(true, Ordering::Release); - } - } - - /// Binds this socket to the given address. - fn bind(&mut self, task: &Task, addr: UnixSocketAddr) -> Result<(), Errno> { - if self.addr.is_some() && !addr.is_unnamed() { - return Err(Errno::EINVAL); - } - if self.addr.is_none() { - let bound_addr = addr.bind(task, true)?; - self.addr = Some(bound_addr); - } - Ok(()) - } - - /// Transitions this socket to listening state. - /// - /// # Arguments - /// - /// * `backlog` - Maximum number of pending connections to queue - fn listen( - self, - backlog: u16, - global: &Arc>, - ) -> Result, (Self, Errno)> { - let Some(addr) = self.addr else { - return Err((self, Errno::EINVAL)); - }; - let key = addr.to_key(); - let backlog = Arc::new(Backlog::new(addr, backlog, self.pollee)); - global - .unix_addr_table - .write() - .insert(key, UnixEntry(UnixEntryInner::Stream(backlog.clone()))); - Ok(UnixListenStream { - backlog, - global: global.clone(), - }) - } - - /// Converts this initial socket into a connected stream pair. - fn into_connected( - self, - peer_addr: Arc>, - ) -> (UnixConnectedStream, UnixConnectedStream) { - let UnixInitStream { - addr, - pollee, - read_shutdown, - write_shutdown, - } = self; - UnixConnectedStream::new_pair( - addr.map(Arc::new), - Some(Arc::new(pollee)), - Some(peer_addr), - read_shutdown.load(Ordering::Acquire), - write_shutdown.load(Ordering::Acquire), - ) - } -} - -/// Connection backlog for a listening Unix socket. -/// -/// Manages the queue of pending connections and the maximum backlog limit. -struct Backlog { - /// The address this socket is listening on - addr: Arc>, - state: Mutex>, - pollee: Pollee, -} - -struct BacklogState { - sockets: VecDeque>, - /// Maximum number of pending connections - limit: u16, - is_shutdown: bool, -} - -impl Backlog { - fn new(addr: UnixBoundSocketAddr, backlog: u16, pollee: Pollee) -> Self { - Self { - addr: Arc::new(addr), - state: litebox::sync::Mutex::new(BacklogState { - sockets: VecDeque::new(), - limit: backlog, - is_shutdown: false, - }), - pollee, - } - } - - /// Updates the maximum backlog size. - fn set_backlog(&self, backlog: u16) { - self.state.lock().limit = backlog; - } - - /// Attempts to establish a connection without blocking. - fn try_connect( - &self, - init: UnixInitStream, - ) -> Result, (UnixInitStream, Errno)> { - let mut state = self.state.lock(); - if state.is_shutdown { - return Err((init, Errno::ECONNREFUSED)); - } - - if state.sockets.len() >= state.limit as usize { - return Err((init, Errno::EAGAIN)); - } - - let (client, server) = init.into_connected(self.addr.clone()); - state.sockets.push_back(server); - - self.pollee.notify_observers(Events::IN); - Ok(client) - } - - /// Attempts to accept a pending connection without blocking. - fn try_accept(&self) -> Result, TryOpError> { - let mut state = self.state.lock(); - match state.sockets.pop_front() { - Some(stream) => { - if !state.is_shutdown { - self.pollee.notify_observers(Events::OUT); - } - Ok(stream) - } - None if state.is_shutdown => Err(TryOpError::Other(Errno::ESHUTDOWN)), - None => Err(TryOpError::TryAgain), - } - } - - fn check_io_events(&self) -> Events { - let state = self.state.lock(); - let mut events = Events::empty(); - if !state.sockets.is_empty() { - events |= Events::IN; - } - if state.is_shutdown { - events |= Events::IN | Events::HUP; - } else if state.sockets.len() < state.limit as usize { - events |= Events::OUT; - } - events - } - - /// Shuts down this backlog, preventing new connections. - fn shutdown(&self) { - let mut state = self.state.lock(); - if !state.is_shutdown { - state.is_shutdown = true; - self.pollee.notify_observers(Events::HUP); - } - } -} - -/// Represents a Unix stream socket in listening state. -struct UnixListenStream { - backlog: Arc>, - global: Arc>, -} - -impl UnixListenStream { - /// Updates the maximum backlog size for pending connections. - fn listen(&self, backlog: u16) { - self.backlog.set_backlog(backlog); - } - - fn register_observer( - &self, - observer: Weak>, - mask: litebox::event::Events, - ) { - self.backlog.pollee.register_observer(observer, mask); - } - - /// Returns the local address this socket is bound to. - fn get_local_addr(&self) -> &UnixBoundSocketAddr { - self.backlog.addr.as_ref() - } -} - -impl Drop for UnixListenStream { - fn drop(&mut self) { - self.backlog.shutdown(); - - let key = self.backlog.addr.to_key(); - let mut table = self.global.unix_addr_table.write(); - // Only remove the entry if it still points to our backlog - if let Some(UnixEntry(UnixEntryInner::Stream(backlog))) = table.get(&key) - && Arc::ptr_eq(backlog, &self.backlog) - { - table.remove(&key); - } - } -} - -/// Tracks the local and peer addresses for a connected socket. -struct AddrView { - addr: Option>>, - peer: Option>>, -} - -impl AddrView { - /// Creates a pair of address views for two connected sockets. - /// - /// The local address of one becomes the peer address of the other. - fn new_pair( - addr: Option>>, - peer: Option>>, - ) -> (Self, Self) { - let first = Self { - addr: addr.clone(), - peer: peer.clone(), - }; - let second = Self { - addr: peer, - peer: addr, - }; - (first, second) - } - - /// Returns the local address, if available. - fn get_local_addr(&self) -> Option<&UnixBoundSocketAddr> { - self.addr.as_deref() - } - - /// Returns the peer address, if available. - fn get_peer_addr(&self) -> Option<&UnixBoundSocketAddr> { - self.peer.as_deref() - } -} - -/// A message sent over a Unix socket. -struct Message { - data: Vec, - // TODO: add control messages - // cmsgs: Option>, -} - -/// Represents a connected Unix stream socket. -struct UnixConnectedStream { - addr: AddrView, - /// The read end of the local socket's channel for receiving messages. - recv_channel: crate::channel::ReadEnd, - /// The write end of the connected peer socket for sending messages. - connected_send_channel: crate::channel::WriteEnd, - pollee: Arc>, -} - -const UNIX_BUF_SIZE: usize = 65536; -impl UnixConnectedStream { - /// Creates a pair of connected Unix stream sockets. - /// - /// `read_shutdown` and `write_shutdown` half-close the corresponding sides of the - /// *first* returned socket only (used to carry pre-connect shutdown flags from - /// `UnixInitStream` across `connect(2)` into the connected state). - fn new_pair( - addr: Option>>, - pollee: Option>>, - peer: Option>>, - read_shutdown: bool, - write_shutdown: bool, - ) -> (Self, Self) { - let (addr1, addr2) = AddrView::new_pair(addr, peer); - let pollee1 = pollee.unwrap_or(Arc::new(Pollee::new())); - let pollee2 = Arc::new(Pollee::new()); - let (send_channel, recv_channel) = - crate::channel::Channel::new(UNIX_BUF_SIZE, pollee2.clone(), pollee1.clone()).split(); - let (send_channel_peer, recv_channel_peer) = - crate::channel::Channel::new(UNIX_BUF_SIZE, pollee1.clone(), pollee2.clone()).split(); - let first = UnixConnectedStream { - addr: addr1, - recv_channel, - connected_send_channel: send_channel_peer, - pollee: pollee1, - }; - let second = UnixConnectedStream { - addr: addr2, - recv_channel: recv_channel_peer, - connected_send_channel: send_channel, - pollee: pollee2, - }; - if read_shutdown { - first.recv_channel.shutdown(); - } - if write_shutdown { - first.connected_send_channel.shutdown(); - } - (first, second) - } - - fn get_local_addr(&self) -> UnixSocketAddr { - match self.addr.get_local_addr() { - Some(addr) => UnixSocketAddr::from(addr), - None => UnixSocketAddr::Unnamed, - } - } - - fn get_peer_addr(&self) -> UnixSocketAddr { - match self.addr.get_peer_addr() { - Some(addr) => UnixSocketAddr::from(addr), - None => UnixSocketAddr::Unnamed, - } - } - - fn try_sendto(&self, msg: Message) -> Result<(), (Message, Errno)> { - // TODO: write partial data? - if msg.data.is_empty() { - return if self.connected_send_channel.is_shutdown() - || self.connected_send_channel.is_peer_shutdown() - { - Err((msg, Errno::EPIPE)) - } else { - Ok(()) - }; - } - self.connected_send_channel.try_write_one(msg) - } - - fn try_recvfrom(&self, mut buf: &mut [u8], peek: bool) -> Result> { - if buf.is_empty() { - return Ok(0); - } - if peek { - let mut total_read = 0; - self.recv_channel - .for_each_queued(|msg| { - let copy_len = (buf.len() - total_read).min(msg.data.len()); - buf[total_read..total_read + copy_len].copy_from_slice(&msg.data[..copy_len]); - total_read += copy_len; - total_read != buf.len() + let resolved = task.fs.borrow().context.read().resolve(path.as_str())?; + Ok(LocalSocketAddress::Path { + path: resolved.to_string(), + name: path.as_bytes().to_vec(), }) - .map_err(|error| match error { - Errno::EAGAIN => TryOpError::TryAgain, - other => TryOpError::Other(other), - })?; - return Ok(total_read); - } - let mut total_read = 0; - while !buf.is_empty() { - let n = match self.recv_channel.peek_and_consume_one(|msg| { - if buf.len() >= msg.data.len() { - buf[..msg.data.len()].copy_from_slice(&msg.data); - Ok((true, msg.data.len())) - } else { - buf.copy_from_slice(&msg.data[..buf.len()]); - msg.data = msg.data.split_off(buf.len()); - Ok((false, buf.len())) - } - }) { - Ok(n) => n, - Err(e) => { - if total_read > 0 { - break; - } - return match e { - Errno::EAGAIN => Err(TryOpError::TryAgain), - other => Err(TryOpError::Other(other)), - }; - } - }; - total_read += n; - buf = &mut buf[n..]; - } - Ok(total_read) - } - - fn check_io_events(&self) -> Events { - let mut events = Events::empty(); - let is_read_shutdown = self.recv_channel.is_shutdown(); - let is_peer_write_shutdown = self.recv_channel.is_peer_shutdown(); - let is_write_shutdown = self.connected_send_channel.is_shutdown(); - if is_read_shutdown || is_peer_write_shutdown { - events |= Events::RDHUP | Events::IN; - if is_write_shutdown { - events |= Events::HUP; } } - if !self.recv_channel.is_empty() { - events |= Events::IN; - } - if !self.connected_send_channel.is_full() { - events |= Events::OUT; - } - events - } - - fn shutdown(&self, how: ShutdownHow) { - let mut events = Events::empty(); - if how.is_shutdown_read() && self.recv_channel.shutdown() { - events |= Events::IN | Events::RDHUP; - } - if how.is_shutdown_write() && self.connected_send_channel.shutdown() { - events |= Events::OUT | Events::HUP; - } - self.pollee.notify_observers(events); - } -} - -enum UnixStreamState { - Init(UnixInitStream), - Listen(UnixListenStream), - Connected(UnixConnectedStream), -} - -impl UnixStreamState { - fn connected(&self) -> Option<&UnixConnectedStream> { - match self { - UnixStreamState::Connected(conn) => Some(conn), - _ => None, - } - } - fn listen(&self) -> Option<&UnixListenStream> { - match self { - UnixStreamState::Listen(listen) => Some(listen), - _ => None, - } - } -} - -struct UnixStream { - state: RwLock>>, -} - -impl UnixStream { - fn new(state: UnixStreamState) -> Self { - Self { - state: litebox::sync::RwLock::new(Some(state)), - } - } - - fn with_state_ref(&self, f: F) -> R - where - F: FnOnce(&UnixStreamState) -> R, - { - let old = self.state.read(); - f(old.as_ref().expect("state should never be None")) - } - - fn with_state_mut_ref(&self, f: F) -> R - where - F: FnOnce(&mut UnixStreamState) -> R, - { - let mut old = self.state.write(); - f(old.as_mut().expect("state should never be None")) - } - - fn with_state(&self, f: F) -> R - where - F: FnOnce(UnixStreamState) -> (UnixStreamState, R), - { - let mut old = self.state.write(); - let (new, result) = f(old.take().expect("state should never be None")); - *old = Some(new); - result - } - - fn bind(&self, task: &Task, addr: UnixSocketAddr) -> Result<(), Errno> { - self.with_state_mut_ref(|state| { - match state { - UnixStreamState::Init(init) => init.bind(task, addr), - UnixStreamState::Listen(_) => { - // Note Linux checks the given address and thus may return - // a different error code (e.g., EADDRINUSE). - Err(Errno::EINVAL) - } - UnixStreamState::Connected(_) => Err(Errno::EISCONN), - } - }) - } - - fn listen(&self, backlog: u16, global: &Arc>) -> Result<(), Errno> { - self.with_state(|state| { - let ret = match state { - UnixStreamState::Init(init) => { - return match init.listen(backlog, global) { - Ok(listen) => (UnixStreamState::Listen(listen), Ok(())), - Err((init, err)) => (UnixStreamState::Init(init), Err(err)), - }; - } - UnixStreamState::Listen(ref listen) => { - listen.listen(backlog); - Ok(()) - } - UnixStreamState::Connected(_) => Err(Errno::EISCONN), - }; - (state, ret) - }) - } - - fn lookup( - &self, - task: &Task, - addr: &UnixSocketAddr, - ) -> Result>, Errno> { - let guard = task.global.unix_addr_table.read(); - let Some(key) = addr.to_key() else { - return Err(Errno::EINVAL); - }; - let Some(entry) = guard.get(&key) else { - return Err(Errno::ECONNREFUSED); - }; - match &entry.0 { - UnixEntryInner::Stream(backlog) => Ok(backlog.clone()), - UnixEntryInner::Datagram(_) => Err(Errno::EPROTOTYPE), - } - } - fn try_connect(&self, backlog: &Backlog) -> Result<(), TryOpError> { - self.with_state(|state| match state { - UnixStreamState::Init(init) => match backlog.try_connect(init) { - Ok(connected) => (UnixStreamState::Connected(connected), Ok(())), - Err((init, err)) => (UnixStreamState::Init(init), Err(err)), - }, - UnixStreamState::Listen(s) => (UnixStreamState::Listen(s), Err(Errno::EINVAL)), - UnixStreamState::Connected(s) => (UnixStreamState::Connected(s), Err(Errno::EISCONN)), - }) - .map_err(|err| match err { - Errno::EAGAIN => TryOpError::TryAgain, - other => TryOpError::Other(other), - }) - } - fn connect( - &self, - task: &Task, - addr: UnixSocketAddr, - is_nonblocking: bool, - ) -> Result<(), Errno> { - let backlog = self.lookup(task, &addr)?; - // check if we can bind to the address - let _ = addr.bind(task, false)?; - task.wait_cx() - .wait_on_events( - is_nonblocking, - Events::OUT, - |observer, mask| { - backlog.pollee.register_observer(observer, mask); - Ok(()) - }, - || self.try_connect(&backlog), - ) - .map_err(Errno::from) - } - - fn accept( - &self, - cx: &WaitContext<'_, Platform>, - mut peer: Option<&mut UnixSocketAddr>, - is_nonblocking: bool, - ) -> Result, Errno> { - let backlog = self.with_state_ref(|state| -> Result>, Errno> { - let listen = state.listen().ok_or(Errno::EINVAL)?; - Ok(listen.backlog.clone()) - })?; - let res = cx - .wait_on_events( - is_nonblocking, - Events::IN, - |observer, mask| { - backlog.pollee.register_observer(observer, mask); - Ok(()) - }, - || { - let accepted = backlog.try_accept()?; - if let Some(peer) = peer.as_deref_mut() { - *peer = accepted.get_peer_addr(); - } - Ok(UnixSocketInner::Stream(UnixStream::new( - UnixStreamState::Connected(accepted), - ))) - }, - ) - .map_err(Errno::from); - // accept on a shut-down listen: Linux returns EAGAIN for non-blocking, EINVAL - // for blocking. try_accept signals shutdown via ESHUTDOWN; translate here. - match res { - Err(Errno::ESHUTDOWN) if is_nonblocking => Err(Errno::EAGAIN), - Err(Errno::ESHUTDOWN) => Err(Errno::EINVAL), - other => other, - } - } - - fn sendto( - &self, - cx: &WaitContext<'_, Platform>, - timeout: Option, - buf: &[u8], - is_nonblocking: bool, - addr: Option, - ) -> Result { - let mut msg = Some(Message { data: buf.to_vec() }); - cx.with_timeout(timeout) - .wait_on_events( - is_nonblocking, - Events::OUT, - |observer, mask| { - self.with_state_ref(|state| { - let conn = state.connected().ok_or(Errno::ENOTCONN)?; - conn.pollee.register_observer(observer, mask); - Ok(()) - }) - }, - || { - self.with_state_ref(|state| { - let conn = state - .connected() - .ok_or(TryOpError::Other(Errno::ENOTCONN))?; - if addr.is_some() { - return Err(TryOpError::Other(Errno::EISCONN)); - } - match conn.try_sendto(msg.take().unwrap()) { - Ok(()) => Ok(buf.len()), - Err((m, Errno::EAGAIN)) => { - let _ = msg.replace(m); - Err(TryOpError::TryAgain) - } - Err((_, err)) => Err(TryOpError::Other(err)), - } - }) - }, - ) - .map_err(|error| wait_errno(timeout, error)) - } - - fn recvfrom( - &self, - cx: &WaitContext<'_, Platform>, - timeout: Option, - buf: &mut [u8], - is_nonblocking: bool, - peek: bool, - mut source_addr: Option<&mut Option>, - ) -> Result { - let res = cx - .with_timeout(timeout) - .wait_on_events( - is_nonblocking, - Events::IN, - |observer, mask| { - self.with_state_ref(|state| { - let conn = state.connected().ok_or(Errno::ENOTCONN)?; - conn.pollee.register_observer(observer, mask); - Ok(()) - }) - }, - || { - self.with_state_ref(|state| { - let conn = state - .connected() - .ok_or(TryOpError::Other(Errno::ENOTCONN))?; - let n = conn.try_recvfrom(buf, peek)?; - // For connected stream sockets, no need to return the source address - if let Some(source_addr) = source_addr.as_deref_mut() { - *source_addr = None; - } - Ok(n) - }) - }, - ) - .map_err(|error| wait_errno(timeout, error)); - match res { - // Linux SO_RCVTIMEO expiry surfaces as `EAGAIN`, not `ETIMEDOUT` - Err(Errno::ETIMEDOUT) => Err(Errno::EAGAIN), - other => other, - } - } - - fn get_local_addr(&self) -> UnixSocketAddr { - self.with_state_ref(|state| match state { - UnixStreamState::Init(init) => init - .addr - .as_ref() - .map_or(UnixSocketAddr::Unnamed, UnixSocketAddr::from), - UnixStreamState::Listen(listen) => UnixSocketAddr::from(listen.get_local_addr()), - UnixStreamState::Connected(connect) => connect.get_local_addr(), - }) - } - fn get_peer_addr(&self) -> Option { - self.with_state_ref(|state| match state { - UnixStreamState::Init(_) | UnixStreamState::Listen(_) => None, - UnixStreamState::Connected(connect) => Some(connect.get_peer_addr()), - }) - } - - fn register_observer( - &self, - observer: Weak>, - mask: Events, - ) { - self.with_state_ref(|state| match state { - UnixStreamState::Init(init) => init.pollee.register_observer(observer, mask), - UnixStreamState::Listen(listen) => listen.register_observer(observer, mask), - UnixStreamState::Connected(connect) => { - connect.pollee.register_observer(observer, mask); - } - }); - } - fn check_io_events(&self) -> Events { - self.with_state_ref(|state| match state { - UnixStreamState::Init(init) => { - // Fresh Init reports OUT|HUP (HUP because not connected). After a - // shutdown(SHUT_RD) on an Init socket, Linux additionally reports IN - // (a recv would return EOF immediately). SHUT_WR has no observable - // effect on Init's poll output. - let mut events = Events::OUT | Events::HUP; - if init.read_shutdown.load(Ordering::Acquire) { - events |= Events::IN; - } - events - } - UnixStreamState::Listen(listen) => listen.backlog.check_io_events(), - UnixStreamState::Connected(conn) => conn.check_io_events(), - }) - } - - fn shutdown(&self, how: ShutdownHow) { - self.with_state_ref(|state| match state { - UnixStreamState::Init(init) => init.shutdown(how), - UnixStreamState::Listen(listen) => { - if how.is_shutdown_read() { - listen.backlog.shutdown(); - } - } - UnixStreamState::Connected(conn) => conn.shutdown(how), - }); - } -} - -/// A datagram message with source address information -#[derive(Clone)] -struct DatagramMessage { - data: Vec, - // TODO: add control messages - // cmsgs: Option>, - source: UnixSocketAddr, -} - -impl WriteEnd { - fn try_write(&self, msg: DatagramMessage) -> Result<(), (DatagramMessage, Errno)> { - self.try_write_one(msg) - } - fn write( - &self, - cx: &WaitContext<'_, Platform>, - timeout: Option, - msg: DatagramMessage, - is_nonblocking: bool, - ) -> Result<(), Errno> { - let mut msg = Some(msg); - cx.with_timeout(timeout) - .wait_on_events( - is_nonblocking, - Events::OUT, - |observer, mask| { - self.register_observer(observer, mask); - Ok(()) - }, - || match self.try_write(msg.take().unwrap()) { - Ok(()) => Ok(()), - Err((m, Errno::EAGAIN)) => { - let _ = msg.replace(m); - Err(TryOpError::TryAgain) - } - Err((_, err)) => Err(TryOpError::Other(err)), - }, - ) - .map_err(|error| wait_errno(timeout, error)) - } -} -impl ReadEnd { - /// Attempts to read a single datagram message without blocking. - /// - /// Reads exactly one message, preserving message boundaries. If the buffer - /// is smaller than the message, the excess data is discarded (truncated). - /// Returns the original message size (which may exceed `buf.len()`). - fn try_read( - &self, - buf: &mut [u8], - peek: bool, - mut source_addr: Option<&mut Option>, - ) -> Result> { - let is_self_shutdown = self.is_shutdown(); - self.peek_and_consume_one(|msg| { - let copy_len = buf.len().min(msg.data.len()); - buf[..copy_len].copy_from_slice(&msg.data[..copy_len]); - if let Some(source_addr) = source_addr.as_deref_mut() { - *source_addr = Some(msg.source.clone()); - } - Ok((!peek, msg.data.len())) - }) - .map_err(|e| match e { - Errno::EAGAIN => TryOpError::TryAgain, - // ESHUTDOWN from the channel layer collapses two distinct conditions: our own - // SHUT_RD (caller wants EOF) and peer SHUT_WR (Linux keeps the socket - // receivable in principle, since other senders could still target it). For - // datagram, only the self case synthesizes EOF; peer-shutdown looks like - // "empty queue, try again". - Errno::ESHUTDOWN if !is_self_shutdown => TryOpError::TryAgain, - other => TryOpError::Other(other), - }) } } -/// The local address of a bound datagram socket together with the global state -/// it was registered in (used to deregister the address on drop). -type BoundDatagramAddr = (UnixBoundSocketAddr, Arc>); - -struct UnixDatagramInner { - /// The local address this socket is bound to, if any. - addr: Option>, - /// The read end of the local socket's channel for receiving messages. - /// Set when the socket is bound via `bind` or `new_pair`. - recv_channel: Option>, - /// The write end of the connected peer socket for sending messages. - /// Set when the socket is connected via `connect` or `new_pair`. - connected_send_channel: Option<(WriteEnd, UnixSocketAddr)>, - read_shutdown: bool, - write_shutdown: bool, - pollee: Arc>, -} -/// Represents a Unix datagram socket. -struct UnixDatagram { - inner: RwLock>, -} - -impl Drop for UnixDatagramInner { - fn drop(&mut self) { - if let Some((addr, global)) = self.addr.take() { - let key = addr.to_key(); - let mut table = global.unix_addr_table.write(); - // Only remove the entry if it matches the current socket - if let Some(UnixEntry(UnixEntryInner::Datagram(send_channel))) = table.get(&key) - && let Some(recv_channel) = &self.recv_channel - && send_channel.is_pair(recv_channel) - { - table.remove(&key); +impl From for UnixSocketAddr { + fn from(name: LocalSocketName) -> Self { + match name { + LocalSocketName::Unnamed => UnixSocketAddr::Unnamed, + LocalSocketName::Path(name) => { + UnixSocketAddr::Path(String::from_utf8_lossy(&name).into_owned()) } + LocalSocketName::Abstract(name) => UnixSocketAddr::Abstract(name), } } } -impl UnixDatagramInner { - /// Binds this socket to the given address. - fn bind(&mut self, task: &Task, addr: UnixSocketAddr) -> Result<(), Errno> { - if self.addr.is_some() { - if addr.is_unnamed() { - return Ok(()); - } - return Err(Errno::EINVAL); - } - - let bound_addr = addr.bind(task, true)?; - let key = bound_addr.to_key(); - // Registers the write end of the socket in the global address table so it - // can receive messages sent to this address. - let (send_channel, recv_channel) = - Channel::new(UNIX_BUF_SIZE, Arc::new(Pollee::new()), self.pollee.clone()).split(); - let _ = task - .global - .unix_addr_table - .write() - .insert(key, UnixEntry(UnixEntryInner::Datagram(send_channel))); - self.addr = Some((bound_addr, task.global.clone())); - if self.read_shutdown { - recv_channel.shutdown(); - } - self.recv_channel = Some(recv_channel); - Ok(()) - } - - fn shutdown(&mut self, how: ShutdownHow) { - let mut events = Events::empty(); - if how.is_shutdown_read() { - self.read_shutdown = true; - if let Some(recv_channel) = &self.recv_channel { - recv_channel.shutdown(); - } - events |= Events::IN | Events::RDHUP; - } - if how.is_shutdown_write() { - self.write_shutdown = true; - if let Some((connected_send_channel, _)) = &self.connected_send_channel { - connected_send_channel.shutdown(); - } - events |= Events::OUT | Events::HUP; - } - self.pollee.notify_observers(events); - } +/// A Unix domain socket descriptor's open file description. +pub(crate) struct UnixSocket { + socket: Arc>, } -impl UnixDatagram { - fn new() -> Self { +impl From> for UnixSocket { + fn from(socket: LocalSocket) -> Self { Self { - inner: RwLock::new(UnixDatagramInner { - addr: None, - recv_channel: None, - connected_send_channel: None, - read_shutdown: false, - write_shutdown: false, - pollee: Arc::new(Pollee::new()), - }), + socket: Arc::new(socket), } } - - fn new_pair() -> (UnixDatagram, UnixDatagram) { - let pollee1 = Arc::new(Pollee::new()); - let pollee2 = Arc::new(Pollee::new()); - let (send_channel, recv_channel) = - crate::channel::Channel::new(UNIX_BUF_SIZE, pollee2.clone(), pollee1.clone()).split(); - let (send_channel_peer, recv_channel_peer) = - crate::channel::Channel::new(UNIX_BUF_SIZE, pollee1.clone(), pollee2.clone()).split(); - ( - // Cross-wire: each socket keeps the other side's send channel. - UnixDatagram { - inner: RwLock::new(UnixDatagramInner { - addr: None, - recv_channel: Some(recv_channel), - connected_send_channel: Some((send_channel_peer, UnixSocketAddr::Unnamed)), - read_shutdown: false, - write_shutdown: false, - pollee: pollee1, - }), - }, - UnixDatagram { - inner: RwLock::new(UnixDatagramInner { - addr: None, - recv_channel: Some(recv_channel_peer), - connected_send_channel: Some((send_channel, UnixSocketAddr::Unnamed)), - read_shutdown: false, - write_shutdown: false, - pollee: pollee2, - }), - }, - ) - } - - /// Binds this socket to the given address. - fn bind(&self, task: &Task, addr: UnixSocketAddr) -> Result<(), Errno> { - self.inner.write().bind(task, addr) - } - - /// Looks up a socket address and returns its write endpoint. - fn lookup( - &self, - task: &Task, - addr: UnixSocketAddr, - ) -> Result, Errno> { - let guard = task.global.unix_addr_table.read(); - let Some(key) = addr.to_key() else { - return Err(Errno::EINVAL); - }; - let Some(entry) = guard.get(&key) else { - return Err(Errno::ECONNREFUSED); - }; - // check if we can bind to the address - let _ = addr.bind(task, false)?; - match &entry.0 { - UnixEntryInner::Stream(_) => Err(Errno::EPROTOTYPE), - UnixEntryInner::Datagram(send_channel) => Ok(send_channel.clone()), - } - } - - /// Connects this socket to a default peer address. - /// - /// Subsequent sends without an address will use this peer. - fn connect(&self, task: &Task, addr: UnixSocketAddr) -> Result<(), Errno> { - let send_channel = self.lookup(task, addr.clone())?; - let mut inner = self.inner.write(); - if inner.write_shutdown { - send_channel.shutdown(); - } - inner.connected_send_channel = Some((send_channel, addr)); - Ok(()) - } - - fn recvfrom( - &self, - cx: &WaitContext<'_, Platform>, - timeout: Option, - buf: &mut [u8], - is_nonblocking: bool, - peek: bool, - mut source_addr: Option<&mut Option>, - ) -> Result { - let res = cx - .with_timeout(timeout) - .wait_on_events( - is_nonblocking, - Events::IN, - |observer, mask| { - self.inner.read().pollee.register_observer(observer, mask); - Ok(()) - }, - || { - let guard = self.inner.read(); - let Some(recv_channel) = &guard.recv_channel else { - return Err(TryOpError::Other(Errno::ENOTCONN)); - }; - recv_channel.try_read(buf, peek, source_addr.as_deref_mut()) - }, - ) - .map_err(|error| wait_errno(timeout, error)); - // - Non-blocking + self-shutdown(SHUT_RD) with empty queue: Linux returns EAGAIN - // instead of EOF (datagram boundaries; no message synthesized for the absent peer). - // - SO_RCVTIMEO expiry on a blocking recv: Linux returns EAGAIN, not ETIMEDOUT - // (the latter is reserved for connect-style timeouts). - match res { - Err(Errno::ESHUTDOWN) if is_nonblocking => Err(Errno::EAGAIN), - Err(Errno::ETIMEDOUT) => Err(Errno::EAGAIN), - other => other, - } - } - - // Sends data to the specified or connected peer. - /// - /// If `addr` is provided, sends to that address. Otherwise, uses the - /// connected peer (set via `connect()`). - fn sendto( - &self, - task: &Task, - timeout: Option, - buf: &[u8], - is_nonblocking: bool, - addr: Option, - ) -> Result { - let source = self.get_local_addr(); - let connected_send_channel = { - let inner = self.inner.read(); - if inner.write_shutdown { - return Err(Errno::EPIPE); - } - inner - .connected_send_channel - .as_ref() - .map(|(send_channel, _)| send_channel.clone()) - }; - - let send_channel = if let Some(addr) = addr { - self.lookup(task, addr)? - } else if let Some(connected_send_channel) = connected_send_channel { - connected_send_channel - } else { - return Err(Errno::ENOTCONN); - }; - send_channel.write( - &task.wait_cx(), - timeout, - DatagramMessage { - data: buf.to_vec(), - source, - }, - is_nonblocking, - )?; - Ok(buf.len()) - } - - fn get_local_addr(&self) -> UnixSocketAddr { - self.inner - .read() - .addr - .as_ref() - .map_or(UnixSocketAddr::Unnamed, |(addr, _)| { - UnixSocketAddr::from(addr) - }) - } - fn get_peer_addr(&self) -> Option { - self.inner - .read() - .connected_send_channel - .as_ref() - .map(|(_, addr)| addr.clone()) - } - - fn check_io_events(&self) -> Events { - let mut events = Events::empty(); - let inner = self.inner.read(); - let recv_shutdown = inner.read_shutdown; - let send_shutdown = inner.write_shutdown; - - if recv_shutdown { - events |= Events::IN | Events::RDHUP; - } else if let Some(recv_channel) = &inner.recv_channel - && !recv_channel.is_empty() - { - events |= Events::IN; - } - - if let Some((connected_send_channel, _)) = &inner.connected_send_channel { - if !connected_send_channel.is_full() { - events |= Events::OUT; - } - } else if !send_shutdown { - // If not connected, allow to sendto any address? - events |= Events::OUT; - } - // Linux reports POLLHUP on a dgram fd only when *both* local directions are - // shut down (peer-side shutdown is invisible since dgrams are connectionless). - if recv_shutdown && send_shutdown { - events |= Events::HUP; - } - events - } - - fn shutdown(&self, how: ShutdownHow) { - let mut inner = self.inner.write(); - inner.shutdown(how); - } -} - -enum UnixSocketInner { - Stream(UnixStream), - Datagram(UnixDatagram), -} -pub(crate) struct UnixSocket { - inner: UnixSocketInner, - status: AtomicU32, - options: Mutex, } impl UnixSocket { - fn new_with_inner(inner: UnixSocketInner, flags: SockFlags) -> Self { - let mut status = OFlags::RDWR; - status.set(OFlags::NONBLOCK, flags.contains(SockFlags::NONBLOCK)); - Self { - inner, - status: AtomicU32::new(status.bits()), - options: litebox::sync::Mutex::new(SocketOptions::default()), - } + pub(super) fn new( + litebox: &litebox::LiteBox, + sock_type: SockType, + flags: SockFlags, + ) -> Result { + let socket = litebox.create_local_socket(socket_type(sock_type)?, open_flags(flags))?; + Ok(socket.into()) } - pub(super) fn recv_timeout(&self) -> Option { - self.options.lock().recv_timeout + pub(super) fn new_connected_pair( + litebox: &litebox::LiteBox, + sock_type: SockType, + flags: SockFlags, + ) -> Result<(Self, Self), Errno> { + let (first, second) = + litebox.create_local_socket_pair(socket_type(sock_type)?, open_flags(flags))?; + Ok((first.into(), second.into())) } - pub(super) fn is_stream(&self) -> bool { - matches!(self.inner, UnixSocketInner::Stream(_)) + /// The broker-owned socket, which a child process inherits. + pub(crate) fn local_socket(&self) -> &Arc> { + &self.socket } - pub(super) fn new(sock_type: SockType, flags: SockFlags) -> Option { - let inner = match sock_type { - SockType::Stream => UnixSocketInner::Stream(UnixStream::new(UnixStreamState::Init( - UnixInitStream::new(), - ))), - SockType::Datagram => UnixSocketInner::Datagram(UnixDatagram::new()), - e => { - log_unsupported!("Unsupported unix socket type: {:?}", e); - return None; - } - }; - Some(Self::new_with_inner(inner, flags)) + pub(super) fn is_stream(&self) -> bool { + self.socket.socket_type() == SocketType::Stream } + /// Binds the socket to `addr`, creating a filesystem node for a path with the permissions + /// that the umask of `task` allows. pub(super) fn bind(&self, task: &Task, addr: UnixSocketAddr) -> Result<(), Errno> { - match &self.inner { - UnixSocketInner::Stream(stream) => stream.bind(task, addr), - UnixSocketInner::Datagram(datagram) => datagram.bind(task, addr), - } + let address = addr.to_local(task)?; + let mode = (Mode::RWXU | Mode::RWXG | Mode::RWXO) & !task.fs.borrow().umask(); + self.socket.bind(&address, acting_user(task), mode)?; + Ok(()) } - pub(super) fn listen( - &self, - backlog: u16, - global: &Arc>, - ) -> Result<(), Errno> { - match &self.inner { - UnixSocketInner::Stream(stream) => stream.listen(backlog, global), - UnixSocketInner::Datagram(_) => Err(Errno::EOPNOTSUPP), - } + pub(super) fn listen(&self, backlog: u16) -> Result<(), Errno> { + self.socket.listen(u32::from(backlog))?; + Ok(()) } pub(super) fn connect(&self, task: &Task, addr: UnixSocketAddr) -> Result<(), Errno> { - match &self.inner { - UnixSocketInner::Stream(stream) => { - let send_timeout = self.options.lock().send_timeout; - stream - .connect(task, addr, self.get_status().contains(OFlags::NONBLOCK)) - .map_err(|error| wait_errno(send_timeout, error)) - } - UnixSocketInner::Datagram(datagram) => datagram.connect(task, addr), - } + let address = addr.to_local(task)?; + self.socket + .connect(&task.wait_cx(), &address, acting_user(task), false)?; + Ok(()) } + /// Accepts a connection, storing the connecting socket's name in `peer` if requested. + /// + /// `flags` sets the status flags of the new socket only. pub(super) fn accept( &self, cx: &WaitContext<'_, Platform>, flags: SockFlags, peer: Option<&mut UnixSocketAddr>, ) -> Result, Errno> { - match &self.inner { - UnixSocketInner::Stream(stream) => { - let recv_timeout = self.recv_timeout(); - let accepted = stream - .accept( - cx, - peer, - self.get_status().contains(OFlags::NONBLOCK) - | flags.contains(SockFlags::NONBLOCK), - ) - .map_err(|error| wait_errno(recv_timeout, error))?; - Ok(UnixSocket::new_with_inner(accepted, flags)) - } - UnixSocketInner::Datagram(_) => Err(Errno::EOPNOTSUPP), + let accepted = self.socket.accept(cx, open_flags(flags), false)?; + if let Some(peer) = peer { + *peer = accepted.name(true)?.into(); } + Ok(accepted.into()) } pub(super) fn sendto( @@ -1437,25 +193,25 @@ impl UnixSocket { log_unsupported!("Unsupported sendto flags: {:?}", flags); return Err(Errno::EINVAL); } - let is_nonblocking = - flags.contains(SendFlags::DONTWAIT) || self.get_status().contains(OFlags::NONBLOCK); - let timeout = self.options.lock().send_timeout; - match &self.inner { - UnixSocketInner::Stream(stream) => { - stream.sendto(&task.wait_cx(), timeout, buf, is_nonblocking, addr) - } - UnixSocketInner::Datagram(datagram) => { - datagram.sendto(task, timeout, buf, is_nonblocking, addr) - } - } + let address = addr.map(|addr| addr.to_local(task)).transpose()?; + Ok(self.socket.send( + &task.wait_cx(), + address.as_ref(), + buf, + acting_user(task), + flags.contains(SendFlags::DONTWAIT), + )?) } + /// Receives into `buf`, returning the length of the received data, which for a datagram + /// socket may exceed `buf`. + /// + /// A datagram socket stores its sender's name in `source_addr` if requested. pub(super) fn recvfrom( &self, cx: &WaitContext<'_, Platform>, buf: &mut [u8], flags: ReceiveFlags, - timeout_override: Option, source_addr: Option<&mut Option>, ) -> Result { let supported_flags = ReceiveFlags::DONTWAIT | ReceiveFlags::PEEK | ReceiveFlags::TRUNC; @@ -1463,64 +219,26 @@ impl UnixSocket { log_unsupported!("Unsupported recvfrom flags: {:?}", flags); return Err(Errno::EINVAL); } - let is_nonblocking = - flags.contains(ReceiveFlags::DONTWAIT) || self.get_status().contains(OFlags::NONBLOCK); - let peek = flags.contains(ReceiveFlags::PEEK); - let timeout = timeout_override.or_else(|| self.options.lock().recv_timeout); - let ret = match &self.inner { - UnixSocketInner::Stream(stream) => { - stream.recvfrom(cx, timeout, buf, is_nonblocking, peek, source_addr) - } - UnixSocketInner::Datagram(datagram) => { - datagram.recvfrom(cx, timeout, buf, is_nonblocking, peek, source_addr) - } - }; - match ret { - Err(Errno::ESHUTDOWN) => Ok(0), - other => other, + let received = self.socket.receive( + cx, + buf, + flags.contains(ReceiveFlags::PEEK), + flags.contains(ReceiveFlags::DONTWAIT), + )?; + if let Some(source_addr) = source_addr + && !self.is_stream() + { + *source_addr = Some(received.source.into()); } + Ok(received.length) } - pub(super) fn get_local_addr(&self) -> UnixSocketAddr { - match &self.inner { - UnixSocketInner::Stream(stream) => stream.get_local_addr(), - UnixSocketInner::Datagram(datagram) => datagram.get_local_addr(), - } - } - pub(super) fn get_peer_addr(&self) -> Option { - match &self.inner { - UnixSocketInner::Stream(stream) => stream.get_peer_addr(), - UnixSocketInner::Datagram(datagram) => datagram.get_peer_addr(), - } + pub(super) fn get_local_addr(&self) -> Result { + Ok(self.socket.name(false)?.into()) } - pub(super) fn new_connected_pair( - ty: SockType, - flags: SockFlags, - ) -> Option<(UnixSocket, UnixSocket)> { - match ty { - SockType::Stream => { - let (conn1, conn2) = UnixConnectedStream::new_pair(None, None, None, false, false); - Some(( - UnixSocket::new_with_inner( - UnixSocketInner::Stream(UnixStream::new(UnixStreamState::Connected(conn1))), - flags, - ), - UnixSocket::new_with_inner( - UnixSocketInner::Stream(UnixStream::new(UnixStreamState::Connected(conn2))), - flags, - ), - )) - } - SockType::Datagram => { - let (datagram1, datagram2) = UnixDatagram::new_pair(); - Some(( - UnixSocket::new_with_inner(UnixSocketInner::Datagram(datagram1), flags), - UnixSocket::new_with_inner(UnixSocketInner::Datagram(datagram2), flags), - )) - } - _ => None, - } + pub(super) fn get_peer_addr(&self) -> Result { + Ok(self.socket.name(true)?.into()) } pub(super) fn setsockopt( @@ -1531,28 +249,28 @@ impl UnixSocket { optlen: usize, ) -> Result<(), Errno> { match global.setsockopt_common(optname, optval, optlen, |so, value| { - match (so, value) { + let option = match (so, value) { (SocketOption::RCVTIMEO, SocketOptionValue::Timeout(timeout)) => { - self.options.lock().recv_timeout = timeout; + LocalSocketOption::ReceiveTimeout(timeout) } (SocketOption::SNDTIMEO, SocketOptionValue::Timeout(timeout)) => { - self.options.lock().send_timeout = timeout; + LocalSocketOption::SendTimeout(timeout) } (SocketOption::LINGER, SocketOptionValue::Timeout(timeout)) => { - self.options.lock().linger_timeout = timeout; + LocalSocketOption::Linger(timeout) } (SocketOption::REUSEADDR, SocketOptionValue::U32(val)) => { - self.options.lock().reuse_address = val != 0; + LocalSocketOption::ReuseAddress(val != 0) } (SocketOption::KEEPALIVE, SocketOptionValue::U32(val)) => { - self.options.lock().keep_alive = val != 0; + LocalSocketOption::KeepAlive(val != 0) } (SocketOption::BROADCAST, SocketOptionValue::U32(val)) => { - self.options.lock().broadcast = val != 0; + LocalSocketOption::Broadcast(val != 0) } _ => unreachable!(), - } - Ok(()) + }; + Ok(self.socket.set_option(option)?) }) { Err(Errno::ENOPROTOOPT) => {} // continue to handle unix other => return other, @@ -1589,6 +307,7 @@ impl UnixSocket { SocketOptionName::TCP(_) => Err(Errno::EOPNOTSUPP), } } + pub(super) fn getsockopt( &self, global: &GlobalState, @@ -1596,23 +315,25 @@ impl UnixSocket { optval: UserPtrMut, len: u32, ) -> Result { - match global.getsockopt_common(optname, optval, len, |sopt| match sopt { - SocketOption::RCVTIMEO => SocketOptionValue::Timeout(self.options.lock().recv_timeout), - SocketOption::SNDTIMEO => SocketOptionValue::Timeout(self.options.lock().send_timeout), - SocketOption::LINGER => SocketOptionValue::Timeout(self.options.lock().linger_timeout), - SocketOption::REUSEADDR => { - SocketOptionValue::U32(u32::from(self.options.lock().reuse_address)) - } - SocketOption::KEEPALIVE => { - SocketOptionValue::U32(u32::from(self.options.lock().keep_alive)) - } - SocketOption::BROADCAST => { - SocketOptionValue::U32(u32::from(self.options.lock().broadcast)) - } - _ => unreachable!(), - }) { - Err(Errno::ENOPROTOOPT) => {} // continue to handle unix - other => return other, + if let SocketOptionName::Socket( + SocketOption::RCVTIMEO + | SocketOption::SNDTIMEO + | SocketOption::LINGER + | SocketOption::REUSEADDR + | SocketOption::KEEPALIVE + | SocketOption::BROADCAST, + ) = optname + { + let options = self.socket.options()?; + return global.getsockopt_common(optname, optval, len, |sopt| match sopt { + SocketOption::RCVTIMEO => SocketOptionValue::Timeout(options.receive_timeout), + SocketOption::SNDTIMEO => SocketOptionValue::Timeout(options.send_timeout), + SocketOption::LINGER => SocketOptionValue::Timeout(options.linger), + SocketOption::REUSEADDR => SocketOptionValue::U32(u32::from(options.reuse_address)), + SocketOption::KEEPALIVE => SocketOptionValue::U32(u32::from(options.keep_alive)), + SocketOption::BROADCAST => SocketOptionValue::U32(u32::from(options.broadcast)), + _ => unreachable!(), + }); } let val: u32 = match optname { @@ -1620,7 +341,7 @@ impl UnixSocket { IpOption::TOS => return Err(Errno::EOPNOTSUPP), }, SocketOptionName::Socket(so) => match so { - // handled by `getsockopt_common` + // handled above SocketOption::RCVTIMEO | SocketOption::SNDTIMEO | SocketOption::LINGER @@ -1631,80 +352,117 @@ impl UnixSocket { } // Unix sockets don't track async errors SocketOption::ERROR => 0, - SocketOption::TYPE => match self.inner { - UnixSocketInner::Stream(_) => SockType::Stream as u32, - UnixSocketInner::Datagram(_) => SockType::Datagram as u32, - }, - SocketOption::RCVBUF | SocketOption::SNDBUF => UNIX_BUF_SIZE.trunc(), - SocketOption::PEERCRED => match &self.inner { - UnixSocketInner::Stream(stream) => { - let ucred = stream.with_state_ref(|state| match state { - UnixStreamState::Connected(_) => { - log_unsupported!("get PEERCRED for unix socket"); - Err(Errno::EOPNOTSUPP) - } - _ => Ok(litebox_common_linux::Ucred { - pid: 0, - uid: u32::MAX, - gid: u32::MAX, - }), - })?; - return super::write_to_user::<_, Platform>(ucred, optval, len); + SocketOption::TYPE => { + if self.is_stream() { + SockType::Stream as u32 + } else { + SockType::Datagram as u32 } - UnixSocketInner::Datagram(_) => { + } + SocketOption::RCVBUF | SocketOption::SNDBUF => LOCAL_SOCKET_BUFFER_SIZE, + SocketOption::PEERCRED => { + if !self.is_stream() { log_unsupported!("get PEERCRED for unix datagram socket"); return Err(Errno::EOPNOTSUPP); } - }, + if self.socket.name(true).is_ok() { + log_unsupported!("get PEERCRED for unix socket"); + return Err(Errno::EOPNOTSUPP); + } + let ucred = litebox_common_linux::Ucred { + pid: 0, + uid: u32::MAX, + gid: u32::MAX, + }; + return super::write_to_user::<_, Platform>(ucred, optval, len); + } }, SocketOptionName::TCP(_) => return Err(Errno::EOPNOTSUPP), }; super::write_to_user::<_, Platform>(val, optval, len) } - pub(super) fn shutdown(&self, how: ShutdownHow) { - match &self.inner { - UnixSocketInner::Stream(stream) => stream.shutdown(how), - UnixSocketInner::Datagram(datagram) => datagram.shutdown(how), - } + pub(super) fn shutdown(&self, how: ShutdownHow) -> Result<(), Errno> { + let mode = match how { + ShutdownHow::Read => ShutdownMode::Read, + ShutdownHow::Write => ShutdownMode::Write, + ShutdownHow::Both => ShutdownMode::Both, + }; + self.socket.shutdown(mode)?; + Ok(()) + } + + /// Returns the `F_GETFL` flags of the socket, which every descriptor sharing it sees. + pub(super) fn get_status(&self) -> Result { + linux_status_flags(self.socket.get_status_flags()?) } - super::common_functions_for_file_status!(); + /// Changes the status flags in `mask` of the socket to their values in `flags`, ignoring + /// flags other than `O_NONBLOCK` and `O_APPEND`. + pub(super) fn set_status_flags(&self, mask: OFlags, flags: OFlags) -> Result<(), Errno> { + let (mask, flags) = status_flags_change(mask, flags); + self.socket.set_status_flags(mask, flags)?; + Ok(()) + } } -impl IOPollable for UnixSocket { - fn register_observer( +impl GlobalState { + /// Runs `f` on the Unix socket at `fd`, without holding the descriptor table. + pub(super) fn with_unix_socket( &self, - observer: Weak>, - mask: Events, - ) { - match &self.inner { - UnixSocketInner::Stream(stream) => { - stream.register_observer(observer, mask); - } - UnixSocketInner::Datagram(datagram) => { - datagram - .inner - .read() - .pollee - .register_observer(observer, mask); - } - } + fd: &TypedFd>, + f: impl FnOnce(&UnixSocket) -> Result, + ) -> Result { + let handle = self + .litebox + .descriptor_table() + .entry_handle(fd) + .ok_or(Errno::EBADF)?; + handle.with_entry(f) + } + + /// Adopts a Unix socket this process inherited from its parent as `handle`. + pub(super) fn adopt_inherited_unix_socket( + &self, + handle: ObjectHandle, + ) -> Result>, ProcessError> { + let socket = self.litebox.adopt_inherited_local_socket(handle)?; + Ok(self + .litebox + .descriptor_table_mut() + .insert::>(UnixSocket::from(socket))) + } +} + +impl IOPollable for UnixSocket { + fn register_observer(&self, observer: Weak>, mask: Events) { + self.socket.register_observer(observer, mask); } fn check_io_events(&self) -> Events { - match &self.inner { - UnixSocketInner::Stream(stream) => stream.check_io_events(), - UnixSocketInner::Datagram(datagram) => datagram.check_io_events(), + self.socket.check_io_events() + } +} + +fn socket_type(sock_type: SockType) -> Result { + match sock_type { + SockType::Stream => Ok(SocketType::Stream), + SockType::Datagram => Ok(SocketType::Datagram), + e => { + log_unsupported!("Unsupported unix socket type: {:?}", e); + Err(Errno::ESOCKTNOSUPPORT) } } } -pub(crate) struct UnixEntry(UnixEntryInner); -enum UnixEntryInner { - Stream(Arc>), - Datagram(WriteEnd), +fn open_flags(flags: SockFlags) -> FileOpenFlags { + if flags.contains(SockFlags::NONBLOCK) { + FileOpenFlags::NONBLOCKING + } else { + FileOpenFlags::NONE + } } -/// Type alias for the global Unix socket address table. -pub(crate) type UnixAddrTable = BTreeMap>; +fn acting_user(task: &Task) -> FileUser { + task.fs.borrow().context.read().acting_user() +} From dc7e63a85914aba4d013dc974ec788c1b14eb5e8 Mon Sep 17 00:00:00 2001 From: Weidong Cui Date: Sat, 3 Oct 2026 21:48:06 -0700 Subject: [PATCH 2/2] Fix Unix sockaddr truncation and datagram SIGPIPE Copy Unix socket addresses truncated to the caller's buffer like Linux, which previously underflowed when addrlen was below 2, and raise SIGPIPE only for Unix stream sockets. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: ceea8bb3-229a-4415-a471-bb3b41d4fed0 --- litebox_shim_linux/src/syscalls/file.rs | 7 +- litebox_shim_linux/src/syscalls/net.rs | 130 +++++++++++++++++------- 2 files changed, 96 insertions(+), 41 deletions(-) diff --git a/litebox_shim_linux/src/syscalls/file.rs b/litebox_shim_linux/src/syscalls/file.rs index f56b048fd0..b2a07d2e58 100644 --- a/litebox_shim_linux/src/syscalls/file.rs +++ b/litebox_shim_linux/src/syscalls/file.rs @@ -996,7 +996,7 @@ impl Task { buf: &[u8], offset: Option, ) -> Result { - let is_inet_datagram = core::cell::Cell::new(false); + let is_datagram = core::cell::Cell::new(false); let result = fd.dispatch( |fd| { self.global @@ -1006,7 +1006,7 @@ impl Task { }, |fd| { espipe_for_non_seekable_offset(offset)?; - is_inet_datagram.set(matches!( + is_datagram.set(matches!( self.global.get_socket_type(fd)?, litebox_common_linux::SockType::Datagram )); @@ -1052,13 +1052,14 @@ impl Task { .entry_handle(fd) .ok_or(Errno::EBADF)?; handle.with_entry(|file| { + is_datagram.set(!file.is_stream()); file.sendto(self, buf, litebox_common_linux::SendFlags::empty(), None) }) }, |_fd| Err(Errno::EINVAL), ); if let Err(Errno::EPIPE) = result - && !is_inet_datagram.get() + && !is_datagram.get() { self.send_signal(Signal::SIGPIPE, signal::siginfo_kill(Signal::SIGPIPE)); } diff --git a/litebox_shim_linux/src/syscalls/net.rs b/litebox_shim_linux/src/syscalls/net.rs index 6992d19005..ba8064a381 100644 --- a/litebox_shim_linux/src/syscalls/net.rs +++ b/litebox_shim_linux/src/syscalls/net.rs @@ -1599,44 +1599,27 @@ pub(crate) fn write_sockaddr_to_user( size_of::() } SocketAddress::Unix(v) => { - let family_ptr = UserPtrMut::::from_usize(addr.as_usize()); - family_ptr - .write_at_offset::(0, AddressFamily::UNIX as u16) - .ok_or(Errno::EFAULT)?; + // Like Linux, copy the full address truncated to the caller's buffer, which may be + // too small even for the family. + let mut bytes = alloc::vec::Vec::new(); + bytes.extend_from_slice(&(AddressFamily::UNIX as u16).to_ne_bytes()); match v { - UnixSocketAddr::Unnamed => { - // only write family - size_of::() - } + UnixSocketAddr::Unnamed => {} UnixSocketAddr::Abstract(name) => { - let offset = offset_of!(CSockUnixAddr, path); - if addrlen_val as usize > offset { - addr.write_at_offset::(isize::try_from(offset).unwrap(), 0) - .ok_or(Errno::EFAULT)?; - let max_len = addrlen_val as usize - offset - 1; - addr.write_slice_at_offset::( - isize::try_from(offset + 1).unwrap(), - &name[..name.len().min(max_len)], - ) - .ok_or(Errno::EFAULT)?; - } - offset + 1 + name.len() + bytes.push(0); + bytes.extend_from_slice(&name); } UnixSocketAddr::Path(path) => { - let offset = offset_of!(CSockUnixAddr, path); - let max_len = addrlen_val as usize - offset; - let name = &path.as_bytes()[..path.len().min(max_len)]; - addr.write_slice_at_offset::(isize::try_from(offset).unwrap(), name) - .ok_or(Errno::EFAULT)?; - let null_offset = offset + name.len(); - // write null terminator if there is space - if addrlen_val as usize > null_offset { - addr.write_at_offset::(isize::try_from(null_offset).unwrap(), 0) - .ok_or(Errno::EFAULT)?; - } - offset + path.len() + 1 + bytes.extend_from_slice(path.as_bytes()); + bytes.push(0); } } + let copied = bytes.len().min(addrlen_val as usize); + if copied != 0 { + addr.write_slice_at_offset::(0, &bytes[..copied]) + .ok_or(Errno::EFAULT)?; + } + bytes.len() } SocketAddress::Inet(SocketAddr::V6(_)) => todo!("copy_sockaddr_to_user for IPv6"), } @@ -1887,12 +1870,12 @@ impl Task { flags: SendFlags, sockaddr: Option, ) -> Result { - let is_inet_datagram = core::cell::Cell::new(false); + let is_datagram = core::cell::Cell::new(false); let res = self.files.borrow().with_typed_socket( &self.global, socket, |fd| { - is_inet_datagram.set(matches!( + is_datagram.set(matches!( self.global.get_socket_type(fd)?, SockType::Datagram )); @@ -1904,6 +1887,7 @@ impl Task { .sendto(&self.wait_cx(), fd, buf, flags, sockaddr) }, |file| { + is_datagram.set(!file.is_stream()); let addr = sockaddr .clone() .map(|addr| addr.unix().ok_or(Errno::EAFNOSUPPORT)) @@ -1913,7 +1897,7 @@ impl Task { ); if let Err(Errno::EPIPE) = res && !flags.contains(SendFlags::NOSIGNAL) - && !is_inet_datagram.get() + && !is_datagram.get() { self.send_signal(Signal::SIGPIPE, signal::siginfo_kill(Signal::SIGPIPE)); } @@ -1963,12 +1947,12 @@ impl Task { .ok_or(Errno::EFAULT)?, ) }; - let is_inet_datagram = core::cell::Cell::new(false); + let is_datagram = core::cell::Cell::new(false); let res = self.files.borrow().with_typed_socket( &self.global, socket, |fd| { - is_inet_datagram.set(matches!( + is_datagram.set(matches!( self.global.get_socket_type(fd)?, SockType::Datagram )); @@ -1981,6 +1965,7 @@ impl Task { .sendto(&self.wait_cx(), fd, &data, flags, sock_addr) }, |file| { + is_datagram.set(!file.is_stream()); let unix_addr = sock_addr .clone() .map(|addr| addr.unix().ok_or(Errno::EAFNOSUPPORT)) @@ -1991,7 +1976,7 @@ impl Task { ); if let Err(Errno::EPIPE) = res && !flags.contains(SendFlags::NOSIGNAL) - && !is_inet_datagram.get() + && !is_datagram.get() { self.send_signal(Signal::SIGPIPE, signal::siginfo_kill(Signal::SIGPIPE)); } @@ -4464,6 +4449,75 @@ mod unix_tests { unix_socket_recv_timeout(SockType::Datagram); } + #[test] + fn unix_addresses_are_truncated_to_the_callers_buffer() { + let task = init_platform(); + let path = "/unix_truncated_name.sock"; + let raw_fd = create_unix_server_socket( + &task, + UnixSocketAddr::Path(path.to_string()), + SockFlags::empty(), + ) + .unwrap(); + let fd = i32::try_from(raw_fd).unwrap(); + let mut expected = (AddressFamily::UNIX as u16).to_ne_bytes().to_vec(); + expected.extend_from_slice(path.as_bytes()); + expected.push(0); + let full_length = u32::try_from(expected.len()).unwrap(); + for capacity in [0, 1, 3, full_length] { + let mut buffer = [0xff_u8; 64]; + let mut length = capacity; + task.sys_getsockname( + fd, + UserPtrMut::from_ptr(buffer.as_mut_ptr()), + UserPtrMut::from_ptr(&raw mut length), + ) + .unwrap(); + assert_eq!(length, full_length); + let copied = capacity as usize; + assert_eq!(&buffer[..copied], &expected[..copied]); + assert!(buffer[copied..].iter().all(|&byte| byte == 0xff)); + } + close_socket(&task, raw_fd); + task.sys_unlinkat(-1, path, AtFlags::empty()).unwrap(); + } + + #[test] + fn only_unix_stream_sockets_raise_sigpipe() { + use litebox_common_linux::signal::Signal; + + let task = init_platform(); + // A raised SIGPIPE stays pending, so the stream socket comes last. + for (ty, raises) in [(SockType::Datagram, false), (SockType::Stream, true)] { + let (sender, receiver) = task + .do_socketpair(AddressFamily::UNIX, ty, SockFlags::empty(), 0) + .unwrap(); + let sender_fd = i32::try_from(sender).unwrap(); + task.sys_shutdown(sender_fd, litebox_common_linux::ShutdownHow::Write as i32) + .unwrap(); + assert_eq!(task.sys_write(sender_fd, b"x", None), Err(Errno::EPIPE)); + let data = b"y"; + assert_eq!( + task.sys_sendto( + sender_fd, + UserPtr::from_usize(data.as_ptr().expose_provenance()), + data.len(), + SendFlags::empty(), + None, + 0, + ), + Err(Errno::EPIPE) + ); + assert_eq!( + task.pending_signal_set().contains(Signal::SIGPIPE), + raises, + "{ty:?}" + ); + close_socket(&task, sender); + close_socket(&task, receiver); + } + } + #[test] fn test_unix_stream_addr() { let task = init_platform();