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..b2a07d2e58 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| { @@ -979,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 @@ -989,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 )); @@ -1035,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)); } @@ -1478,117 +1496,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 +2210,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 +2294,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 +2593,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..ba8064a381 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 = @@ -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"), } @@ -1859,7 +1842,7 @@ impl Task { &self.global, socket, |fd| self.global.listen(fd, backlog), - |file| file.listen(backlog, &self.global), + |file| file.listen(backlog), ) } @@ -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)); } @@ -2168,19 +2153,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 +2239,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 +2677,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 +2704,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 +2740,7 @@ impl Task { }, |file| { let how = ShutdownHow::try_from(how).map_err(|_| Errno::EINVAL)?; - file.shutdown(how); - Ok(()) + file.shutdown(how) }, ) } @@ -3963,7 +3928,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 +4012,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 +4273,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(); @@ -4420,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(); 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() +}