From 8264a31b14f433c3d50dfd91ea93616755ad91b0 Mon Sep 17 00:00:00 2001 From: Weidong Cui Date: Fri, 7 Aug 2026 09:56:28 -0700 Subject: [PATCH] Remove smoltcp networking Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: b5a1a347-37a8-4246-8bbc-306590921475 --- .config/nextest.toml | 12 - Cargo.lock | 154 +- dev_bench/src/main.rs | 9 +- litebox/Cargo.toml | 1 - litebox/src/net/errors.rs | 8 - litebox/src/net/local_ports.rs | 146 - litebox/src/net/mod.rs | 1866 ++-------- litebox/src/net/phy.rs | 103 - litebox/src/net/socket_channel.rs | 1597 +-------- litebox/src/net/tests.rs | 230 -- litebox/src/platform/mock.rs | 22 - litebox/src/platform/mod.rs | 32 +- litebox_broker_core/src/policy.rs | 39 + litebox_broker_core/src/session.rs | 5 +- litebox_broker_core/src/socket.rs | 1053 +++++- litebox_broker_host/src/lib.rs | 114 +- .../src/lib.rs | 2 +- .../src/socket.rs | 3174 +++++++++++++++-- litebox_broker_protocol/src/socket.rs | 8 +- litebox_broker_userland/src/main.rs | 89 +- litebox_common_linux/src/errno/mod.rs | 16 - .../src/host/mock.rs | 8 - .../src/host/snp/snp-sandbox.h | 2 - .../src/host/snp/snp_impl.rs | 20 - litebox_platform_linux_kernel/src/lib.rs | 51 +- litebox_platform_linux_userland/README.md | 9 - .../scripts/_common.sh | 50 - .../scripts/tun-setup.sh | 94 - litebox_platform_linux_userland/src/lib.rs | 188 +- litebox_platform_lvbs/src/host/lvbs_impl.rs | 8 - litebox_platform_lvbs/src/host/mock.rs | 8 - litebox_platform_lvbs/src/lib.rs | 40 +- litebox_platform_windows_userland/src/lib.rs | 19 - litebox_runner_linux_userland/src/lib.rs | 58 +- litebox_runner_linux_userland/tests/loader.rs | 13 +- litebox_runner_linux_userland/tests/run.rs | 36 +- .../tests/tcp_broker_server.c | 10 +- .../src/lib.rs | 2 +- litebox_runner_snp/src/entry.S | 7 - litebox_runner_snp/src/main.rs | 25 - .../src/lib.rs | 2 +- litebox_shim_linux/src/lib.rs | 14 +- litebox_shim_linux/src/loader/elf.rs | 2 +- litebox_shim_linux/src/stdio.rs | 4 +- litebox_shim_linux/src/syscalls/epoll.rs | 12 +- litebox_shim_linux/src/syscalls/eventfd.rs | 2 +- litebox_shim_linux/src/syscalls/file.rs | 10 +- litebox_shim_linux/src/syscalls/misc.rs | 4 +- litebox_shim_linux/src/syscalls/mm.rs | 20 +- litebox_shim_linux/src/syscalls/net.rs | 132 +- litebox_shim_linux/src/syscalls/process.rs | 18 +- litebox_shim_linux/src/syscalls/tests.rs | 46 +- litebox_shim_linux/src/transport.rs | 4 +- litebox_shim_optee/src/syscalls/tests.rs | 6 +- litebox_shim_windows/src/tests.rs | 2 +- 55 files changed, 4400 insertions(+), 5206 deletions(-) delete mode 100644 litebox/src/net/local_ports.rs delete mode 100644 litebox/src/net/phy.rs delete mode 100644 litebox/src/net/tests.rs delete mode 100644 litebox_platform_linux_userland/README.md delete mode 100644 litebox_platform_linux_userland/scripts/_common.sh delete mode 100755 litebox_platform_linux_userland/scripts/tun-setup.sh diff --git a/.config/nextest.toml b/.config/nextest.toml index 9713158962..47f511c0a8 100644 --- a/.config/nextest.toml +++ b/.config/nextest.toml @@ -8,18 +8,6 @@ status-level = "all" failure-output = "immediate-final" # Any tests that take longer than 10 minutes should be killed as a failure slow-timeout = { period = "60s", terminate-after = 10 } -[test-groups] -# ensure only one test accessing tun at a time -tun-access = { max-threads = 1 } - -[[profile.ci.overrides]] -filter = 'test(test_tun)' -test-group = "tun-access" - -[[profile.default.overrides]] -filter = 'test(test_tun)' -test-group = "tun-access" - # For flaky tests on the CI, we attempt to re-run a couple of times before # failing. We explicitly don't do this on the default non-CI profile, so that # flakiness shows up on dev machines accurately, but it is annoying on the CI diff --git a/Cargo.lock b/Cargo.lock index 234b23548b..0970dcf20b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -163,7 +163,7 @@ version = "0.71.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5f58bf3d7db68cfbac37cfc485a8d711e87e064c3d0fe0435b92f7a407f9d6b3" dependencies = [ - "bitflags 2.13.1", + "bitflags", "cexpr", "clang-sys", "itertools", @@ -183,12 +183,6 @@ version = "0.10.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1e4b40c7323adcfc0a41c4b88143ed58346ff65a288fc144329c5c45e05d70c6" -[[package]] -name = "bitflags" -version = "1.3.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" - [[package]] name = "bitflags" version = "2.13.1" @@ -241,12 +235,6 @@ version = "3.20.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5d20789868f4b01b2f2caec9f5c4e0213b41e3e5702a50157d699ae31ced2fcb" -[[package]] -name = "byteorder" -version = "1.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" - [[package]] name = "bytes" version = "1.11.1" @@ -521,47 +509,6 @@ dependencies = [ "syn", ] -[[package]] -name = "defmt" -version = "0.3.100" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f0963443817029b2024136fc4dd07a5107eb8f977eaf18fcd1fdeb11306b64ad" -dependencies = [ - "defmt 1.0.1", -] - -[[package]] -name = "defmt" -version = "1.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "548d977b6da32fa1d1fda2876453da1e7df63ad0304c8b3dae4dbe7b96f39b78" -dependencies = [ - "bitflags 1.3.2", - "defmt-macros", -] - -[[package]] -name = "defmt-macros" -version = "1.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3d4fc12a85bcf441cfe44344c4b72d58493178ce635338a3f3b78943aceb258e" -dependencies = [ - "defmt-parser", - "proc-macro-error2", - "proc-macro2", - "quote", - "syn", -] - -[[package]] -name = "defmt-parser" -version = "1.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "10d60334b3b2e7c9d91ef8150abfb6fa4c1c39ebbcf4a81c2e346aad939fee3e" -dependencies = [ - "thiserror", -] - [[package]] name = "der" version = "0.7.10" @@ -963,15 +910,6 @@ dependencies = [ "regex-syntax", ] -[[package]] -name = "hash32" -version = "0.3.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "47d60b12902ba28e2730cd37e95b8c9223af2808df9e902d4df49588d1470606" -dependencies = [ - "byteorder", -] - [[package]] name = "hashbrown" version = "0.15.5" @@ -983,16 +921,6 @@ dependencies = [ "foldhash", ] -[[package]] -name = "heapless" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0bfb9eb618601c89945a70e254898da93b13be0388091d42117462b265bb3fad" -dependencies = [ - "hash32", - "stable_deref_trait", -] - [[package]] name = "heck" version = "0.5.0" @@ -1418,7 +1346,7 @@ version = "0.1.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1744e39d1d6a9948f4f388969627434e31128196de472883b39f148769bfe30a" dependencies = [ - "bitflags 2.13.1", + "bitflags", "libc", "plain", "redox_syscall", @@ -1435,7 +1363,7 @@ name = "litebox" version = "0.1.0" dependencies = [ "arrayvec", - "bitflags 2.13.1", + "bitflags", "buddy_system_allocator", "either", "hashbrown", @@ -1447,7 +1375,6 @@ dependencies = [ "ringbuf", "slabmalloc", "smallvec", - "smoltcp", "spin 0.9.8", "tar-no-std", "tempfile", @@ -1460,7 +1387,7 @@ dependencies = [ name = "litebox_broker_core" version = "0.1.0" dependencies = [ - "bitflags 2.13.1", + "bitflags", "hashbrown", "litebox_broker_protocol", "spin 0.9.8", @@ -1541,7 +1468,7 @@ dependencies = [ name = "litebox_common_linux" version = "0.1.0" dependencies = [ - "bitflags 2.13.1", + "bitflags", "cfg-if", "elf", "int-enum", @@ -1555,7 +1482,7 @@ dependencies = [ name = "litebox_common_lvbs" version = "0.1.0" dependencies = [ - "bitflags 2.13.1", + "bitflags", "litebox", "litebox_common_linux", "num_enum", @@ -1568,7 +1495,7 @@ dependencies = [ name = "litebox_common_optee" version = "0.1.0" dependencies = [ - "bitflags 2.13.1", + "bitflags", "elf", "litebox", "litebox_common_linux", @@ -1609,7 +1536,7 @@ version = "0.1.0" dependencies = [ "arrayvec", "bindgen", - "bitflags 2.13.1", + "bitflags", "litebox", "litebox_common_linux", "litebox_util_log", @@ -1646,7 +1573,7 @@ dependencies = [ "aligned-vec", "arrayvec", "authenticode", - "bitflags 2.13.1", + "bitflags", "cms", "const-oid", "digest", @@ -1818,7 +1745,7 @@ name = "litebox_shim_linux" version = "0.1.0" dependencies = [ "arrayvec", - "bitflags 2.13.1", + "bitflags", "bitvec", "libc", "litebox", @@ -1866,7 +1793,7 @@ dependencies = [ name = "litebox_shim_windows" version = "0.1.0" dependencies = [ - "bitflags 2.13.1", + "bitflags", "int-enum", "litebox", "litebox_common_linux", @@ -1943,12 +1870,6 @@ dependencies = [ "value-bag", ] -[[package]] -name = "managed" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0ca88d725a0a943b096803bd34e73a4437208b6077654cc4ecb2947a5f91618d" - [[package]] name = "matchers" version = "0.2.0" @@ -2211,7 +2132,7 @@ version = "0.10.80" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a45fa2aa886c42762255da344f0a0d313e254066c46aad76f300c3d3da62d967" dependencies = [ - "bitflags 2.13.1", + "bitflags", "cfg-if", "foreign-types", "libc", @@ -2371,28 +2292,6 @@ dependencies = [ "syn", ] -[[package]] -name = "proc-macro-error-attr2" -version = "2.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "96de42df36bb9bba5542fe9f1a054b8cc87e172759a1868aa05c1f3acc89dfc5" -dependencies = [ - "proc-macro2", - "quote", -] - -[[package]] -name = "proc-macro-error2" -version = "2.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "11ec05c52be0a07b08061f7dd003e7d7092e0472bc731b4af7bb1ef876109802" -dependencies = [ - "proc-macro-error-attr2", - "proc-macro2", - "quote", - "syn", -] - [[package]] name = "proc-macro2" version = "1.0.101" @@ -2477,7 +2376,7 @@ version = "11.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "498cd0dc59d73224351ee52a95fee0f1a617a2eae0e7d9d720cc622c73a54186" dependencies = [ - "bitflags 2.13.1", + "bitflags", ] [[package]] @@ -2506,7 +2405,7 @@ version = "0.7.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6ce70a74e890531977d37e532c34d45e9055d2409ed08ddba14529471ed0be16" dependencies = [ - "bitflags 2.13.1", + "bitflags", ] [[package]] @@ -2619,7 +2518,7 @@ version = "1.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cd15f8a2c5551a84d56efdc1cd049089e409ac19a3072d5037a17fd70719ff3e" dependencies = [ - "bitflags 2.13.1", + "bitflags", "errno", "libc", "linux-raw-sys", @@ -2686,7 +2585,7 @@ version = "2.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "897b2245f0b511c87893af39b033e5ca9cce68824c4d7e7630b5a1d339658d02" dependencies = [ - "bitflags 2.13.1", + "bitflags", "core-foundation", "core-foundation-sys", "libc", @@ -2852,21 +2751,6 @@ version = "1.15.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" -[[package]] -name = "smoltcp" -version = "0.12.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dad095989c1533c1c266d9b1e8d70a1329dd3723c3edac6d03bbd67e7bf6f4bb" -dependencies = [ - "bitflags 1.3.2", - "byteorder", - "cfg-if", - "defmt 0.3.100", - "heapless", - "log", - "managed", -] - [[package]] name = "socket2" version = "0.6.3" @@ -3082,7 +2966,7 @@ version = "0.3.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ac9ee8b664c9f1740cd813fea422116f8ba29997bb7c878d1940424889802897" dependencies = [ - "bitflags 2.13.1", + "bitflags", "log", "num-traits", ] @@ -3224,7 +3108,7 @@ version = "0.6.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d4e6559d53cc268e5031cd8429d05415bc4cb4aefc4aa5d6cc35fbf5b924a1f8" dependencies = [ - "bitflags 2.13.1", + "bitflags", "bytes", "futures-util", "http", @@ -3818,7 +3702,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0f042214de98141e9c8706e8192b73f56494087cc55ebec28ce10f26c5c364ae" dependencies = [ "bit_field", - "bitflags 2.13.1", + "bitflags", "rustversion", "volatile", ] diff --git a/dev_bench/src/main.rs b/dev_bench/src/main.rs index a769af28dc..63524fc168 100644 --- a/dev_bench/src/main.rs +++ b/dev_bench/src/main.rs @@ -712,15 +712,18 @@ fn run_rewritten_iperf3(ctx: BenchCtx<'_>) -> Result<()> { let client_sh = xshell::Shell::new()?; let max_server_attempts = 5; let max_client_attempts = 50; + let guest_port = "5201"; let mut last_failure = "sandboxed server did not become ready".to_owned(); for server_attempt in 1..=max_server_attempts { - let port = TcpListener::bind((Ipv4Addr::LOCALHOST, 0))? + let host_port = TcpListener::bind((Ipv4Addr::LOCALHOST, 0))? .local_addr()? .port() .to_string(); let mut server_command = std::process::Command::new(&broker); server_command + .arg("--publish-tcp") + .arg(format!("127.0.0.1:{host_port}:{guest_port}")) .arg("--runner") .arg(&runner) .arg("--") @@ -733,7 +736,7 @@ fn run_rewritten_iperf3(ctx: BenchCtx<'_>) -> Result<()> { ]) .arg(&tar_file) .arg(&iperf3_rewritten) - .args(["-s", "-1", "-B", "127.0.0.1", "-p", &port]); + .args(["-s", "-1", "-B", "127.0.0.1", "-p", guest_port]); if COMMAND_EXECUTION_IS_QUIET.load(Relaxed) { server_command .stdout(std::process::Stdio::null()) @@ -749,7 +752,7 @@ fn run_rewritten_iperf3(ctx: BenchCtx<'_>) -> Result<()> { for client_attempt in 1..=max_client_attempts { let result = cmd!( client_sh, - "{iperf3_host} -c 127.0.0.1 -p {port} --bytes 1G --connect-timeout 50" + "{iperf3_host} -c 127.0.0.1 -p {host_port} --bytes 1G --connect-timeout 50" ) .quiet() .ignore_stdout() diff --git a/litebox/Cargo.toml b/litebox/Cargo.toml index 5898dc9ce5..3ee51c485b 100644 --- a/litebox/Cargo.toml +++ b/litebox/Cargo.toml @@ -9,7 +9,6 @@ bitflags = "2.13.1" either = { version = "1.13.0", default-features = false } hashbrown = "0.15.2" smallvec = "1.13.2" -smoltcp = { version = "0.12.0", default-features = false, features = ["log", "proto-ipv4", "medium-ip", "socket", "socket-tcp", "socket-udp", "socket-icmp", "socket-raw", "alloc"] } spin = { version = "0.9.8", default-features = false, features = ["spin_mutex"] } tar-no-std = { version = "0.3.3", default-features = false, features = ["alloc"] } thiserror = { version = "2.0.6", default-features = false } diff --git a/litebox/src/net/errors.rs b/litebox/src/net/errors.rs index 0b74ab1f22..a6f7f0b5d6 100644 --- a/litebox/src/net/errors.rs +++ b/litebox/src/net/errors.rs @@ -5,8 +5,6 @@ use core::net::SocketAddr; -use super::local_ports::LocalPortAllocationError; - #[expect( unused_imports, reason = "used for doc string links to work out, but not for code" @@ -63,8 +61,6 @@ pub enum ConnectError { InvalidFd, #[error("Unsupported address {0}")] UnsupportedAddress(SocketAddr), - #[error("Port allocation failed: {0}")] - PortAllocationFailure(#[from] LocalPortAllocationError), #[error("Invalid address")] Unaddressable, #[error("Connection is still in progress")] @@ -123,8 +119,6 @@ pub enum ListenError { InvalidAddress, #[error("Socket is in invalid state")] InvalidState, - #[error("No available free ephemeral ports")] - NoAvailableFreeEphemeralPorts, #[error("Listening is unsupported for this socket")] UnsupportedOperation, #[error("Socket operation failed: {0:?}")] @@ -161,8 +155,6 @@ pub enum SendError { BufferFull, #[error("Datagram is too large")] MessageTooLong, - #[error("port allocation failed: {0}")] - PortAllocationFailure(#[from] LocalPortAllocationError), #[error("unnecessary destination address provided")] UnnecessaryDestinationAddress, #[error("destination address required but not provided")] diff --git a/litebox/src/net/local_ports.rs b/litebox/src/net/local_ports.rs deleted file mode 100644 index f74b5bf5cc..0000000000 --- a/litebox/src/net/local_ports.rs +++ /dev/null @@ -1,146 +0,0 @@ -// Copyright (c) Microsoft Corporation. -// Licensed under the MIT license. - -//! Handling the allocation of local ports - -use core::num::{NonZeroU16, NonZeroU64}; - -use hashbrown::HashMap; -use thiserror::Error; - -use crate::utils::rng::FastRng; - -/// An allocator for local ports, making sure that no already-allocated ports are given out either -/// in case of ephemeral port allocation, or in the case of asking for a specific port. -pub(crate) struct LocalPortAllocator { - // map from port number -> reference count - // - // using a non-zero u16 for the reference count is a memory optimization; if this is ever an - // issue, it can trivially be bumped up to a larger size. - refcount: HashMap, - rng: FastRng, -} - -impl Default for LocalPortAllocator { - fn default() -> Self { - Self::new() - } -} - -impl LocalPortAllocator { - /// Sets up a new local port allocator - pub(crate) fn new() -> Self { - Self { - refcount: HashMap::new(), - rng: FastRng::new_from_seed(NonZeroU64::new(0x13374a4159421337).unwrap()), - } - } - - /// Allocate a new ephemeral local port (i.e., port in the range 49152 and 65535) - pub(crate) fn ephemeral_port(&mut self) -> Result { - for _ in 0..100 { - let port = - NonZeroU16::new(u16::try_from(self.rng.next_in_range_u32(49152..65536)).unwrap()) - .unwrap(); - if let Ok(local_port) = self.specific_port(port) { - return Ok(local_port); - } - } - // If we haven't yet found a port after 100 tries, it is highly likely lots of ports are - // already in use, so we should start looking over them one by one - for port in 49152..=65535 { - let port = NonZeroU16::new(port).unwrap(); - if let Ok(local_port) = self.specific_port(port) { - return Ok(local_port); - } - } - // If we _still_ haven't found any, then we have run out of ports to give out - Err(LocalPortAllocationError::NoAvailableFreePorts) - } - - /// Allocate a specific local port, if available - pub(crate) fn specific_port( - &mut self, - port: NonZeroU16, - ) -> Result { - if self.refcount.contains_key(&port) { - Err(LocalPortAllocationError::AlreadyInUse(port.get())) - } else { - self.refcount.insert(port, NonZeroU16::new(1).unwrap()); - Ok(LocalPort { port }) - } - } - - /// Allocate a local port, either ephemeral (if `port` is 0) or specific (if `port` is non-zero) - pub(crate) fn allocate_local_port( - &mut self, - port: u16, - ) -> Result { - let Some(port) = NonZeroU16::new(port) else { - return self.ephemeral_port(); - }; - self.specific_port(port) - } - - /// Increments the ref-count for a local port, producing a new [`LocalPort`] token to be used - #[must_use] - pub(crate) fn allocate_same_local_port(&mut self, port: &LocalPort) -> LocalPort { - let Some(refcount) = self.refcount.get_mut(&port.port) else { - // Because we have a `LocalPort`, it is (as an invariant) impossible to have the value - // be missing from the refcount. - unreachable!() - }; - // We just bump the refcount, making sure there is no overflow, and then produce the new - // `LocalPort` token. - *refcount = refcount.checked_add(1).unwrap(); - LocalPort { port: port.port } - } - - /// Consumes a [`LocalPort`], possibly marking it as available again. - pub(crate) fn deallocate(&mut self, port: LocalPort) { - let Some(refcount) = self.refcount.get_mut(&port.port) else { - // Because we have a `LocalPort`, it is (as an invariant) impossible to have the value - // be missing from the refcount. - unreachable!() - }; - match refcount.get() { - 0 => unreachable!(), - 1 => { - // Need to drop - self.refcount.remove(&port.port); - } - _ => { - *refcount = NonZeroU16::new(refcount.get() - 1).unwrap(); - } - } - } - - /// Deallocate a port number tracked by this allocator. - pub(crate) fn deallocate_port(&mut self, port: u16) { - if let Some(port) = NonZeroU16::new(port) { - self.deallocate(LocalPort { port }); - } - } -} - -/// A token expressing ownership over a specific local port. -/// -/// Explicitly not cloneable/copyable. -pub(crate) struct LocalPort { - port: NonZeroU16, -} - -impl LocalPort { - pub(crate) fn port(&self) -> u16 { - self.port.get() - } -} - -/// Errors that could be returned when allocating a port -#[derive(Debug, Clone, Copy, Error)] -pub enum LocalPortAllocationError { - #[error("Port {0} is already in use")] - AlreadyInUse(u16), - #[error("No free ports are available")] - NoAvailableFreePorts, -} diff --git a/litebox/src/net/mod.rs b/litebox/src/net/mod.rs index ff4aa3cb95..3610444865 100644 --- a/litebox/src/net/mod.rs +++ b/litebox/src/net/mod.rs @@ -3,36 +3,31 @@ //! Network-related functionality -use alloc::vec; -use alloc::vec::Vec; -use core::net::{Ipv4Addr, SocketAddr, SocketAddrV4}; -use core::sync::atomic::{AtomicBool, Ordering}; - -use crate::event::Events; -use crate::net::socket_channel::NetworkProxy; -use crate::platform::{Instant, TimeProvider}; -use crate::sync::RawSyncPrimitivesProvider; -use crate::{LiteBox, platform, sync}; +use alloc::{sync::Arc, vec, vec::Vec}; +use core::{ + net::SocketAddr, + sync::atomic::{AtomicBool, Ordering}, +}; use bitflags::bitflags; -use smoltcp::socket::{tcp, udp}; + +use crate::{ + LiteBox, + net::socket_channel::NetworkProxy, + platform::{self, TimeProvider}, + sync::{self, RawSyncPrimitivesProvider}, +}; mod broker_socket; pub mod errors; -pub mod local_ports; -mod phy; pub mod socket_channel; pub use broker_socket::{BrokerTcpSocket, BrokerUdpSocket}; -#[cfg(test)] -mod tests; - use errors::{ AcceptError, BindError, CloseError, ConnectError, ListenError, LocalAddrError, ReceiveError, RemoteAddrError, SendError, SocketError, }; -use local_ports::{LocalPort, LocalPortAllocator}; impl From for SocketError { fn from(error: crate::broker::error::BrokerObjectError) -> Self { @@ -49,16 +44,9 @@ impl From for SocketError { } } -/// IP address for LiteBox interface -// TODO: Make this configurable -const INTERFACE_IP_ADDR: Ipv4Addr = Ipv4Addr::new(10, 0, 0, 2); - -/// IP address for the gateway -// TODO: Make this configurable -const GATEWAY_IP_ADDR: Ipv4Addr = Ipv4Addr::new(10, 0, 0, 1); const ICMP_PROTOCOL_NUMBER: u8 = 1; -/// Maximum size of rx/tx buffers for sockets +/// Socket buffer size reported through socket option interfaces. pub const SOCKET_BUFFER_SIZE: usize = 65536 * 4; /// Maximum bytes one socket receive syscall stages before returning. pub const SOCKET_RECEIVE_OPERATION_SIZE: usize = 0x80_000; @@ -77,724 +65,68 @@ pub enum ShutdownDirection { Both, } -/// Limits maximum number of packets in a buffer -#[cfg(test)] -const MAX_PACKET_COUNT: usize = 32; - -/// TCP connection timeout. -const TCP_CONNECT_TIMEOUT: smoltcp::time::Duration = smoltcp::time::Duration::from_secs(75); - -/// The `Network` provides access to all networking related functionality provided by LiteBox. -/// -/// A LiteBox `Network` is parametric in the platform it runs on. +/// The `Network` provides access to broker-owned networking functionality. /// -/// An important decision that must be made by a user of a `Network` is decided by -/// [`set_platform_interaction`](Self::set_platform_interaction), whose docs explain this further. -/// -/// A user of `Network` who cares about [events](crate::event) should call -/// [`attach_socket_proxy`](Self::attach_socket_proxy) -/// to set up a proxy for each socket created, so that events can be notified properly. +/// [`attach_socket_proxy`](Self::attach_socket_proxy) provides a per-socket interface for +/// nonblocking I/O and readiness notification. pub struct Network where - Platform: - platform::IPInterfaceProvider + platform::TimeProvider + sync::RawSyncPrimitivesProvider, + Platform: platform::TimeProvider + sync::RawSyncPrimitivesProvider, { litebox: LiteBox, - /// The set of sockets - socket_set: smoltcp::iface::SocketSet<'static>, - /// The actual "physical" device, that connects to the platform - device: phy::Device, - /// The smoltcp network interface - interface: smoltcp::iface::Interface, - /// Initial instant of creation, used as an arbitrary stop point from when time begins - zero_time: Platform::Instant, - /// An allocator for local ports - // TODO: Maybe we should have separate allocators for TCP, UDP, ...? - local_port_allocator: LocalPortAllocator, - /// Whether outside interaction is automatic or manual - platform_interaction: PlatformInteraction, - /// FDs that are queued for eventual closure + /// FDs that are queued for eventual closure while another operation pins their entry. queued_for_closure: Vec>, - /// Sockets that are closing in the background - closing_in_background: Vec, } impl Network where - Platform: - platform::IPInterfaceProvider + platform::TimeProvider + sync::RawSyncPrimitivesProvider, + Platform: platform::TimeProvider + sync::RawSyncPrimitivesProvider, { - /// Construct a new `Network` instance - /// - /// This function is expected to only be invoked once per platform, as an initialization step, - /// and the created `Network` handle is expected to be shared across all usage over the - /// system. + /// Construct a new `Network` instance. pub fn new(litebox: &LiteBox) -> Self { - let mut device = phy::Device::new(litebox.x.platform); - let config = smoltcp::iface::Config::new(smoltcp::wire::HardwareAddress::Ip); - let mut interface = - smoltcp::iface::Interface::new(config, &mut device, smoltcp::time::Instant::ZERO); - interface.update_ip_addrs(|ip_addrs| { - match ip_addrs.push(smoltcp::wire::IpCidr::new( - smoltcp::wire::IpAddress::Ipv4(INTERFACE_IP_ADDR), - 24, - )) { - Ok(()) => {} - Err(_) => unreachable!(), - } - }); - match interface - .routes_mut() - .add_default_ipv4_route(GATEWAY_IP_ADDR) - { - Ok(None) => {} - _ => unreachable!(), - } Self { litebox: litebox.clone(), - socket_set: smoltcp::iface::SocketSet::new(vec![]), - device, - interface, - zero_time: litebox.x.platform.now(), - local_port_allocator: LocalPortAllocator::new(), - platform_interaction: PlatformInteraction::Automatic, queued_for_closure: vec![], - closing_in_background: vec![], } } } -/// [`SocketHandle`] stores all relevant information for a specific [`SocketFd`], for easy access -/// from [`SocketFd`], _except_ the `Socket` itself which is stored in the [`Network::socket_set`]. +/// Per-open-file-description state for a broker-owned socket. pub(crate) struct SocketHandle { - /// Whether this socket handle is going away soon (i.e., `close` has been invoked upon it but - /// it lingers for a bit to allow pending data to be sent). - consider_closed: bool, - /// The handle into the `socket_set`, absent for broker-owned sockets. - handle: Option, - /// Broker-owned socket state, absent for locally implemented sockets. - broker_socket: Option>, - // Protocol-specific data - specific: ProtocolSpecific, - /// The proxy associated with this socket to enable lock-free data transfer - /// and event notification - proxy: Option>>, + broker_socket: BrokerSocket, + /// Whether the final TCP close should be abortive. + immediate_close: AtomicBool, + /// The proxy associated with this socket for I/O and event notification. + proxy: Option>>, } enum BrokerSocket { - Tcp(alloc::sync::Arc>), - Udp(alloc::sync::Arc>), + Tcp(Arc>), + Udp(Arc>), } impl Clone for BrokerSocket { fn clone(&self) -> Self { match self { - Self::Tcp(socket) => Self::Tcp(alloc::sync::Arc::clone(socket)), - Self::Udp(socket) => Self::Udp(alloc::sync::Arc::clone(socket)), + Self::Tcp(socket) => Self::Tcp(Arc::clone(socket)), + Self::Udp(socket) => Self::Udp(Arc::clone(socket)), } } } impl SocketHandle { - fn smoltcp_handle(&self) -> smoltcp::iface::SocketHandle { - self.handle - .expect("broker-owned socket must not enter the smoltcp path") - } - - /// Convenience function to perform an operation depending on the socket type - fn with_smoltcp_socket( - &self, - socket_set: &smoltcp::iface::SocketSet<'static>, - tcp: TCP, - udp: UDP, - ) -> R - where - TCP: FnOnce(&tcp::Socket) -> R, - UDP: FnOnce(&udp::Socket) -> R, - { - match self.protocol() { - crate::net::Protocol::Tcp => { - let tcp_socket = socket_set.get::(self.smoltcp_handle()); - tcp(tcp_socket) - } - crate::net::Protocol::Udp => { - let udp_socket = socket_set.get::(self.smoltcp_handle()); - udp(udp_socket) - } - crate::net::Protocol::Icmp | crate::net::Protocol::Raw { protocol: _ } => { - unimplemented!() - } - } - } - - // Convenience function to perform a mutable operation depending on the socket type - fn with_smoltcp_socket_mut( - &mut self, - socket_set: &mut smoltcp::iface::SocketSet<'static>, - tcp: TCP, - udp: UDP, - ) -> R - where - TCP: FnOnce(&mut tcp::Socket) -> R, - UDP: FnOnce(&mut udp::Socket) -> R, - { - match self.protocol() { - crate::net::Protocol::Tcp => { - let tcp_socket = socket_set.get_mut::(self.smoltcp_handle()); - tcp(tcp_socket) - } - crate::net::Protocol::Udp => { - let udp_socket = socket_set.get_mut::(self.smoltcp_handle()); - udp(udp_socket) - } - crate::net::Protocol::Icmp | crate::net::Protocol::Raw { protocol: _ } => { - unimplemented!() - } - } - } -} - -impl core::ops::Deref - for SocketHandle -{ - type Target = ProtocolSpecific; - fn deref(&self) -> &Self::Target { - &self.specific - } -} -impl core::ops::DerefMut - for SocketHandle -{ - fn deref_mut(&mut self) -> &mut Self::Target { - &mut self.specific - } -} - -/// The [`ProtocolSpecific`] stores socket-type-specific data -#[expect( - dead_code, - reason = "these might eventually get used, they exist for completeness sake" -)] -pub(crate) enum ProtocolSpecific { - Tcp(TcpSpecific), - Udp(UdpSpecific), - Icmp(IcmpSpecific), - Raw(RawSpecific), -} - -/// Socket-specific data for TCP sockets -pub(crate) struct TcpSpecific { - /// A local port associated with this socket, if any - local_port: Option, - /// Server socket specific data - server_socket: Option, - /// Whether to immediately close the socket when closed (i.e., no graceful FIN handshake) - immediate_close: AtomicBool, - /// Timestamp when `connect` was initiated - connect_initiated_at_us: Option, -} - -/// Socket-specific data for TCP server sockets -struct TcpServerSpecific { - /// IP listening endpoint, if used as a server socket - ip_listen_endpoint: smoltcp::wire::IpListenEndpoint, - /// Specified backlog via `listen`, no packets can be `accept`ed unless this is `Some` - backlog: Option, - /// Handles into the top-level `socket_set` for when things are `accept`ed. - socket_set_handles: Vec, -} - -impl TcpServerSpecific { - fn refill_to_backlog(&mut self, socket_set: &mut smoltcp::iface::SocketSet) { - let backlog = self.backlog.unwrap(); - for _ in self.socket_set_handles.len()..backlog.into() { - let mut listening_socket = tcp::Socket::new( - smoltcp::storage::RingBuffer::new(vec![0u8; SOCKET_BUFFER_SIZE]), - smoltcp::storage::RingBuffer::new(vec![0u8; SOCKET_BUFFER_SIZE]), - ); - match listening_socket.listen(self.ip_listen_endpoint) { - Ok(()) => {} - Err(tcp::ListenError::InvalidState) => { - // Impossible, because we _just_ created a new tcp::Socket, which begins - // in CLOSED state. - unreachable!() - } - Err(tcp::ListenError::Unaddressable) => { - // Impossible, since listen endpoint port is non 0. - unreachable!() - } - } - self.socket_set_handles - .push(socket_set.add(listening_socket)); - } - } -} - -/// Socket-specific data for UDP sockets -pub(crate) struct UdpSpecific { - /// Remote endpoint - /// - /// If `connect`-ed, this is the remote endpoint to which packets are sent by default. - remote_endpoint: Option, -} - -/// Socket-specific data for ICMP sockets -pub(crate) struct IcmpSpecific {} - -/// Socket-specific data for RAW sockets -pub(crate) struct RawSpecific { - protocol: u8, -} - -#[expect( - dead_code, - reason = "the dead ones exist for completeness sake, might eventually get used" -)] -impl ProtocolSpecific { - /// Get the [`Protocol`] for this socket fn protocol(&self) -> Protocol { - match self { - ProtocolSpecific::Tcp(_) => Protocol::Tcp, - ProtocolSpecific::Udp(_) => Protocol::Udp, - ProtocolSpecific::Icmp(_) => Protocol::Icmp, - ProtocolSpecific::Raw(RawSpecific { protocol, .. }) => Protocol::Raw { - protocol: *protocol, - }, - } - } - - /// Obtain a reference to the tcp-socket-specific data. Panics if non-TCP. - fn tcp(&self) -> &TcpSpecific { - match self { - ProtocolSpecific::Tcp(specific) => specific, - _ => unreachable!(), - } - } - - /// Obtain a mutable reference to the tcp-socket-specific data. Panics if non-TCP. - fn tcp_mut(&mut self) -> &mut TcpSpecific { - match self { - ProtocolSpecific::Tcp(specific) => specific, - _ => unreachable!(), - } - } - - /// Obtain a reference to the udp-socket-specific data. Panics if non-UDP. - fn udp(&self) -> &UdpSpecific { - match self { - ProtocolSpecific::Udp(specific) => specific, - _ => unreachable!(), - } - } - - /// Obtain a mutable reference to the udp-socket-specific data. Panics if non-UDP. - fn udp_mut(&mut self) -> &mut UdpSpecific { - match self { - ProtocolSpecific::Udp(specific) => specific, - _ => unreachable!(), - } - } - - /// Obtain a reference to the icmp-socket-specific data. Panics if non-ICMP. - fn icmp(&self) -> &IcmpSpecific { - match self { - ProtocolSpecific::Icmp(specific) => specific, - _ => unreachable!(), - } - } - - /// Obtain a mutable reference to the icmp-socket-specific data. Panics if non-ICMP. - fn icmp_mut(&mut self) -> &mut IcmpSpecific { - match self { - ProtocolSpecific::Icmp(specific) => specific, - _ => unreachable!(), - } - } - - /// Obtain a reference to the raw-socket-specific data. Panics if non-RAW. - fn raw(&self) -> &RawSpecific { - match self { - ProtocolSpecific::Raw(specific) => specific, - _ => unreachable!(), - } - } - - /// Obtain a mutable reference to the raw-socket-specific data. Panics if non-RAW. - fn raw_mut(&mut self) -> &mut RawSpecific { - match self { - ProtocolSpecific::Raw(specific) => specific, - _ => unreachable!(), + match self.broker_socket { + BrokerSocket::Tcp(_) => Protocol::Tcp, + BrokerSocket::Udp(_) => Protocol::Udp, } } } -/// Whether [`Network::perform_platform_interaction`] needs to be manually invoked or not. -pub enum PlatformInteraction { - /// Automatically (internally) invoked whenever any calls like `send`/`recv`/... are made. - Automatic, - /// Requires manually (periodically) invoking [`Network::perform_platform_interaction`] - Manual, -} - -/// Direction of polling for platform interaction -#[derive(Clone, Copy)] -enum PollDirection { - /// Ingress (receiving) direction - Ingress, - /// Egress (sending) direction - Egress, - /// Both directions - Both, -} - -/// Advice on when to invoke [`Network::perform_platform_interaction`] again. -/// -/// It is perfectly ok to ignore this advice by calling things sooner (say, in a tight loop). -/// Specifically, it is harmless (but wastes energy) to call for interaction again sooner than -/// advised. In contrast, it _may_ be harmful (impacting quality of service) to call it later than -/// requested. -#[derive(Clone, Copy, Debug)] -pub enum PlatformInteractionReinvocationAdvice { - /// It is likely helpful to call again immediately, without any delay. The function has returned - /// control back to you to prevent unbounded length waits (crucial to prevent in - /// non-pre-emptible environments), but otherwise has more work it anticipates it can do. - CallAgainImmediately, - /// You don't need to call again until more packets arrive on the device's receive side (i.e., `timeout` is `None`), - /// or the given `timeout` expires. - WaitOnDeviceOrSocketInteraction { - timeout: Option, - }, -} -impl PlatformInteractionReinvocationAdvice { - /// Convenience function to match against [`Self::CallAgainImmediately`] - #[must_use] - pub fn call_again_immediately(self) -> bool { - matches!(self, Self::CallAgainImmediately) - } -} - impl Network where - Platform: - platform::IPInterfaceProvider + platform::TimeProvider + sync::RawSyncPrimitivesProvider, + Platform: platform::TimeProvider + sync::RawSyncPrimitivesProvider, { - /// Sets the interaction with the outside world to `platform_interaction`. - /// - /// If this is set to automatic, then a user of the network does not need to worry about - /// scheduling or calling [`perform_platform_interaction`](Self::perform_platform_interaction). - /// However, this may reduce predictability in terms of how quickly LiteBox responds to calls, - /// since any network calls may incur non-trivial performance penalty. - /// - /// On the other hand, more performance can be had in scenarios that can support (say) a - /// separate thread that repeatedly invokes - /// [`perform_platform_interaction`](Self::perform_platform_interaction), or in scenarios where - /// the user wants greater control over _when_ processing is performed, if done synchronously. - /// - /// By default, for convenience, the default setting (if this function is not invoked) is - /// [`PlatformInteraction::Automatic`]. - pub fn set_platform_interaction(&mut self, platform_interaction: PlatformInteraction) { - self.platform_interaction = platform_interaction; - } - - /// Performs queued interactions with the outside world. - /// - /// # Panics - /// - /// This function panics if run without first using [`Self::set_platform_interaction`] to set - /// interactions to manual. - pub fn perform_platform_interaction(&mut self) -> PlatformInteractionReinvocationAdvice { - assert!( - matches!(self.platform_interaction, PlatformInteraction::Manual), - "Requires manual-mode interactions" - ); - match self.internal_perform_platform_interaction() { - smoltcp::iface::PollResult::SocketStateChanged => { - PlatformInteractionReinvocationAdvice::CallAgainImmediately - } - smoltcp::iface::PollResult::None => { - let poll_at = self.poll_at(); - PlatformInteractionReinvocationAdvice::WaitOnDeviceOrSocketInteraction { - timeout: poll_at, - } - } - } - } - - /// Return a _soft timeout_ (duration to wait) before calling [`Self::perform_platform_interaction`] again. - /// - /// Returns `None` if there is no pending timeout (i.e., no scheduled work requiring network operations). - fn poll_at(&mut self) -> Option { - let timestamp = self.now(); - self.interface - .poll_at(timestamp, &self.socket_set) - .map(|instant| { - if timestamp < instant { - let diff = instant - timestamp; - diff.into() - } else { - core::time::Duration::ZERO - } - }) - } - - /// (Internal-only API) Actually perform the queued interactions with the outside world. - fn internal_perform_platform_interaction(&mut self) -> smoltcp::iface::PollResult { - self.attempt_to_close_queued(); - self.remove_dead_sockets(); - self.close_pending_sockets(); - - // Drain all socket channel buffers before polling to ensure data flows - self.drain_all_socket_channel_buffers(); - self.interface - .poll(self.now(), &mut self.device, &mut self.socket_set) - } - - /// (Internal-only API) Perform the queued interactions. - fn automated_platform_interaction(&mut self, _direction: PollDirection) { - match self.platform_interaction { - PlatformInteraction::Automatic => { - self.internal_perform_platform_interaction(); - } - PlatformInteraction::Manual => {} - } - } - - /// Remove dead sockets that were closing in the background - fn remove_dead_sockets(&mut self) { - self.closing_in_background.retain(|socket_handle| { - let handle = *socket_handle; - let tcp_socket = self.socket_set.get::(handle); - // a socket in the CLOSED state with the remote endpoint set means that an outgoing RST packet is pending - if !tcp_socket.is_open() && tcp_socket.remote_endpoint().is_none() { - self.socket_set.remove(handle); - false - } else { - true - } - }); - } - - /// Close all finished sockets that are marked as closed but waiting for pending data to be sent - fn close_pending_sockets(&mut self) { - let table = self.litebox.descriptor_table(); - for (_, mut handle) in table.iter_mut::>() { - let socket_handle = &mut handle.entry; - if socket_handle.consider_closed { - if socket_handle.broker_socket.is_some() { - socket_handle.consider_closed = false; - continue; - } - // check if there is pending data to be sent - if let Some(proxy) = &socket_handle.proxy - && proxy.has_pending_tx() - { - continue; - } - - let closed = socket_handle.with_smoltcp_socket_mut( - &mut self.socket_set, - |tcp_socket| { - let has_pending_data = tcp_socket.may_send() && tcp_socket.send_queue() > 0; - if !has_pending_data { - tcp_socket.close(); - } - !has_pending_data - }, - |udp_socket| { - let has_pending_data = udp_socket.is_open() && udp_socket.send_queue() > 0; - if !has_pending_data { - udp_socket.close(); - } - !has_pending_data - }, - ); - if closed { - socket_handle.consider_closed = false; - } - } - } - } - - /// Drain all socket channel buffers - fn drain_all_socket_channel_buffers(&mut self) { - let now = self.now(); - let table = self.litebox.descriptor_table(); - for (_, entry) in table.iter::>() { - Self::drain_socket_channel_buffers(&mut self.socket_set, &entry.entry, now); - } - } - - /// Drain data between socket channels and smoltcp sockets. - /// - /// This transfers data from the TX ring buffer (user writes) to the smoltcp socket, - /// and from the smoltcp socket to the RX ring buffer (user reads). - /// - /// Should be called periodically by the network worker to keep data flowing. - fn drain_socket_channel_buffers( - socket_set: &mut smoltcp::iface::SocketSet<'static>, - socket_handle: &SocketHandle, - now: smoltcp::time::Instant, - ) { - if socket_handle.broker_socket.is_some() { - return; - } - let proxy = match &socket_handle.proxy { - Some(proxy) => proxy.as_ref(), - None => return, - }; - match (socket_handle.protocol(), proxy) { - (Protocol::Tcp, NetworkProxy::Stream(proxy)) => { - let tcp_socket = socket_set.get_mut::(socket_handle.smoltcp_handle()); - - // Drain TX buffer: from ring buffer directly to smoltcp - while tcp_socket.can_send() { - let sent = proxy - .pop_tx_data_with(|data| tcp_socket.send_slice(data).unwrap_or_default()); - if sent == 0 { - break; - } - } - - // Drain RX buffer: from smoltcp directly to ring buffer - while tcp_socket.can_recv() { - let received = proxy - .push_rx_data_with(|buf| tcp_socket.recv_slice(buf).unwrap_or_default()); - if received == 0 { - break; - } - } - if let tcp::State::Established = tcp_socket.state() { - proxy.set_state(socket_channel::SocketState::Connected); - proxy.clear_async_error(); - } - let tcp_specific = socket_handle.specific.tcp(); - // Update socket state in the channel - // server socket that is listening also has closed state - if !tcp_socket.is_open() && tcp_specific.server_socket.is_none() { - // Determine error based on previous socket state - match proxy.state() { - socket_channel::SocketState::Connecting => { - // Socket closed while connecting. Distinguish RST from timeout. - let error = match tcp_specific.connect_initiated_at_us { - Some(initiated_at) if now - initiated_at >= TCP_CONNECT_TIMEOUT => { - errors::SocketAsyncError::TimedOut - } - _ => errors::SocketAsyncError::ConnectionRefused, - }; - proxy.set_async_error(error); - proxy.set_state(socket_channel::SocketState::Error); - } - socket_channel::SocketState::Connected => { - // Connection was reset by peer - proxy.set_async_error(errors::SocketAsyncError::ConnectionReset); - proxy.set_state(socket_channel::SocketState::Closed); - } - _ => { - proxy.set_state(socket_channel::SocketState::Closed); - } - } - } - if proxy.state() == socket_channel::SocketState::Connected && !tcp_socket.may_recv() - { - proxy.set_peer_eof(); - } - - if let Some(server_socket) = tcp_specific.server_socket.as_ref() - && !proxy.is_readable() - { - server_socket - .socket_set_handles - .iter() - .any(|&h| { - let socket: &tcp::Socket = socket_set.get(h); - socket.state() == tcp::State::Established - }) - .then(|| { - proxy.set_readable(true); - proxy.notify_io_event(Events::IN); - }); - } - } - (Protocol::Udp, NetworkProxy::Datagram(udp_proxy)) => { - let udp_socket = socket_set.get_mut::(socket_handle.smoltcp_handle()); - let remote_endpoint = socket_handle.udp().remote_endpoint; - - // Drain TX queue: try to send datagrams, consume only on success - while udp_socket.can_send() { - // Try to send - consumes datagram only if closure returns true - let result = udp_proxy.try_send_datagram_with(|data, addr| { - let destination = addr - .map(|s| match s { - SocketAddr::V4(addr) => smoltcp::wire::IpEndpoint::from(addr), - SocketAddr::V6(_) => unimplemented!(), - }) - .or(remote_endpoint); - if let Some(endpoint) = destination { - udp_socket.send_slice(data, endpoint).is_ok() - } else { - // No destination - discard - true - } - }); - if result != Some(true) { - // Either queue empty or send failed - break; - } - } - - // Drain RX: receive from smoltcp, push to channel - while udp_socket.can_recv() { - let received = udp_proxy.try_recv_datagram_with(|| { - let (data, meta) = udp_socket.recv().ok()?; - let source_addr = match meta.endpoint.addr { - smoltcp::wire::IpAddress::Ipv4(ipv4) => SocketAddr::V4( - core::net::SocketAddrV4::new(ipv4, meta.endpoint.port), - ), - }; - Some((data.into(), source_addr)) - }); - if received.is_none() { - break; - } - } - } - (Protocol::Icmp | Protocol::Raw { .. }, _) => { - unimplemented!() - } - ( - Protocol::Tcp | Protocol::Udp, - NetworkProxy::BrokerStream(_) | NetworkProxy::BrokerDatagram(_), - ) => { - unreachable!() - } - _ => panic!("Mismatched protocol and proxy type"), - } - } -} - -impl Network -where - Platform: - platform::IPInterfaceProvider + platform::TimeProvider + sync::RawSyncPrimitivesProvider, -{ - /// Explicitly private-only function that returns the current (smoltcp) Instant, relative to the - /// initialized arbitrary 0-point in time. - fn now(&self) -> smoltcp::time::Instant { - smoltcp::time::Instant::from_micros( - // This conversion from u128 to i64 should practically never fail, since 2^63 - // microseconds is roughly 250 years. If a system has been up for that long, then it - // deserves to panic. - i64::try_from( - self.device - .platform - .now() - .duration_since(&self.zero_time) - .as_micros(), - ) - .unwrap(), - ) - } - /// Creates a broker-owned TCP or UDP socket. /// /// By default, the created socket has no associated proxy; to set a proxy, use @@ -820,195 +152,87 @@ where } }; - Ok(self.new_socket_fd(protocol, None, Some(broker_socket))) - } - - #[cfg(test)] - fn local_socket(&mut self, protocol: Protocol) -> Result, SocketError> { - let handle = match protocol { - Protocol::Tcp => Some(self.socket_set.add(tcp::Socket::new( - smoltcp::storage::RingBuffer::new(vec![0u8; SOCKET_BUFFER_SIZE]), - smoltcp::storage::RingBuffer::new(vec![0u8; SOCKET_BUFFER_SIZE]), - ))), - Protocol::Udp => Some(self.socket_set.add(udp::Socket::new( - smoltcp::storage::PacketBuffer::new( - vec![smoltcp::storage::PacketMetadata::EMPTY; MAX_PACKET_COUNT], - vec![0u8; SOCKET_BUFFER_SIZE], - ), - smoltcp::storage::PacketBuffer::new( - vec![smoltcp::storage::PacketMetadata::EMPTY; MAX_PACKET_COUNT], - vec![0u8; SOCKET_BUFFER_SIZE], - ), - ))), - Protocol::Icmp => { - return Err(SocketError::UnsupportedProtocol(ICMP_PROTOCOL_NUMBER)); - } - Protocol::Raw { protocol } => { - return Err(SocketError::UnsupportedProtocol(protocol)); - } - }; - - Ok(self.new_socket_fd(protocol, handle, None)) - } - - fn new_socket_fd( - &mut self, - protocol: Protocol, - handle: Option, - broker_socket: Option>, - ) -> SocketFd { - self.new_socket_fd_for(SocketHandle { - consider_closed: false, - handle, + Ok(self.litebox.descriptor_table_mut().insert(SocketHandle { broker_socket, - specific: match protocol { - Protocol::Tcp => ProtocolSpecific::Tcp(TcpSpecific { - local_port: None, - server_socket: None, - immediate_close: AtomicBool::new(false), - connect_initiated_at_us: None, - }), - Protocol::Udp => ProtocolSpecific::Udp(UdpSpecific { - remote_endpoint: None, - }), - Protocol::Icmp => unimplemented!(), - Protocol::Raw { protocol: _ } => unimplemented!(), - }, + immediate_close: AtomicBool::new(false), proxy: None, - }) - } - - /// Creates a new [`SocketFd`] for a newly-created [`SocketHandle`]. - fn new_socket_fd_for(&mut self, socket_handle: SocketHandle) -> SocketFd { - self.litebox.descriptor_table_mut().insert(socket_handle) + })) } - /// Creates and attaches the userspace I/O proxy matching the backend of `fd`. - /// - /// Locally implemented sockets receive a channel that transfers data to and from smoltcp. - /// Broker-owned sockets instead receive a proxy for their existing broker socket. The - /// returned [`Arc`](alloc::sync::Arc) is the same proxy stored by the network subsystem for - /// event notification and backend interaction. + /// Creates and attaches the userspace I/O proxy matching `fd`. /// /// If a proxy is already attached, this returns another reference to that proxy. /// - /// Returns `None` if `fd` is invalid or its protocol does not support userspace I/O proxies. + /// Returns `None` if `fd` is invalid. #[must_use] pub fn attach_socket_proxy( &self, fd: &SocketFd, - ) -> Option>> { + ) -> Option>> { let descriptor_table = self.litebox.descriptor_table(); let mut entry = descriptor_table.get_entry_mut(fd)?; let socket_handle = &mut entry.entry; if let Some(proxy) = &socket_handle.proxy { - return Some(alloc::sync::Arc::clone(proxy)); + return Some(Arc::clone(proxy)); } - let proxy = if let Some(socket) = &socket_handle.broker_socket { - match socket { - BrokerSocket::Tcp(socket) => { - NetworkProxy::BrokerStream(alloc::sync::Arc::clone(socket)) - } - BrokerSocket::Udp(socket) => { - NetworkProxy::BrokerDatagram(alloc::sync::Arc::clone(socket)) - } - } - } else { - match socket_handle.protocol() { - Protocol::Tcp => NetworkProxy::Stream(socket_channel::StreamSocketChannel::new()), - Protocol::Udp => { - NetworkProxy::Datagram(socket_channel::DatagramSocketChannel::new()) - } - Protocol::Raw { .. } => NetworkProxy::Raw, - Protocol::Icmp => return None, - } + let proxy = match &socket_handle.broker_socket { + BrokerSocket::Tcp(socket) => NetworkProxy::BrokerStream(Arc::clone(socket)), + BrokerSocket::Udp(socket) => NetworkProxy::BrokerDatagram(Arc::clone(socket)), }; - let proxy = alloc::sync::Arc::new(proxy); - socket_handle.proxy = Some(alloc::sync::Arc::clone(&proxy)); + let proxy = Arc::new(proxy); + socket_handle.proxy = Some(Arc::clone(&proxy)); Some(proxy) } - /// Close the socket at `fd` + /// Close the socket at `fd`. pub fn close( &mut self, fd: &SocketFd, behavior: CloseBehavior, ) -> Result<(), CloseError> { - let mut dt = self.litebox.descriptor_table_mut(); - dt.with_entry_mut(fd, |entry| { - let socket_handle = &entry.entry; - if let crate::net::Protocol::Tcp = socket_handle.protocol() { - socket_handle.tcp().immediate_close.store( - matches!(behavior, CloseBehavior::Immediate), - Ordering::SeqCst, - ); - } - }) - .ok_or(CloseError::InvalidFd)?; - // We close immediately if we can - match dt - .close_and_duplicate_if_shared(fd, |entry| { - match behavior { - CloseBehavior::Immediate | CloseBehavior::Graceful => return true, - CloseBehavior::GracefulIfNoPendingData => {} - } - // check if there is pending data to be sent - let socket_handle = &entry.entry; - if let Some(proxy) = &socket_handle.proxy - && proxy.has_pending_tx() - { - return false; + let mut descriptor_table = self.litebox.descriptor_table_mut(); + descriptor_table + .with_entry_mut(fd, |entry| { + if matches!(entry.entry.protocol(), Protocol::Tcp) { + entry.entry.immediate_close.store( + matches!(behavior, CloseBehavior::Immediate), + Ordering::SeqCst, + ); } - if socket_handle.broker_socket.is_some() { - return true; - } - !socket_handle.with_smoltcp_socket( - &self.socket_set, - |tcp_socket| tcp_socket.may_send() && tcp_socket.send_queue() > 0, - |udp_socket| udp_socket.is_open() && udp_socket.send_queue() > 0, - ) }) + .ok_or(CloseError::InvalidFd)?; + + match descriptor_table + .close_and_duplicate_if_shared(fd, |_| true) .ok_or(CloseError::InvalidFd)? { super::fd::CloseResult::Closed(socket_handle) => { - // Can immediately close it out. - drop(dt); - self.close_handle(socket_handle.entry); + drop(descriptor_table); + Self::close_handle(socket_handle.entry); } - super::fd::CloseResult::Duplicated(dup_fd) => { - // It seems like there might be other duplicates around (e.g., due to `dup`), so we - // can't immediately close it out. - // We attempt to queue it for future closure and then just return. - self.queued_for_closure.push(dup_fd); - drop(dt); + super::fd::CloseResult::Duplicated(duplicate) => { + self.queued_for_closure.push(duplicate); + drop(descriptor_table); self.attempt_to_close_queued(); } - super::fd::CloseResult::Deferred => { - let Some(()) = dt.with_entry_mut(fd, |entry| entry.entry.consider_closed = true) - else { - unreachable!() - }; - return Err(CloseError::DataPending); - } + super::fd::CloseResult::Deferred => unreachable!(), } Ok(()) } - /// Attempt to close as many queued-to-close FDs as possible. Returns `true` iff any of them - /// were closed. + /// Attempt to close as many queued FDs as possible. fn attempt_to_close_queued(&mut self) -> bool { if self.queued_for_closure.is_empty() { - // fast path return false; } - let mut dt = self.litebox.descriptor_table_mut(); - let entries = dt.drain_entries_full_covered_by(&mut self.queued_for_closure); - drop(dt); + let mut descriptor_table = self.litebox.descriptor_table_mut(); + let entries = descriptor_table.drain_entries_full_covered_by(&mut self.queued_for_closure); + drop(descriptor_table); if entries.is_empty() { return false; } for entry in entries { - self.close_handle(entry.entry); + Self::close_handle(entry.entry); } true } @@ -1019,67 +243,19 @@ where !self.queued_for_closure.is_empty() } - /// Close the `socket_handle` - fn close_handle(&mut self, socket_handle: SocketHandle) { + fn close_handle(socket_handle: SocketHandle) { let SocketHandle { - consider_closed: _, - handle, broker_socket, - mut specific, - proxy, + immediate_close, + proxy: _, } = socket_handle; - if let Some(socket) = broker_socket { - match socket { - BrokerSocket::Tcp(socket) => { - let abortive = specific.tcp().immediate_close.load(Ordering::SeqCst); - socket.close(abortive); - } - BrokerSocket::Udp(socket) => socket.close(), - } - return; + match broker_socket { + BrokerSocket::Tcp(socket) => socket.close(immediate_close.load(Ordering::SeqCst)), + BrokerSocket::Udp(socket) => socket.close(), } - let handle = handle.expect("local socket must have a smoltcp handle"); - match specific.protocol() { - Protocol::Raw { .. } | Protocol::Icmp => { - // There is no close/abort for raw and icmp sockets - let _ = self.socket_set.remove(handle); - } - Protocol::Udp => { - let smoltcp::socket::Socket::Udp(mut socket) = self.socket_set.remove(handle) - else { - unreachable!() - }; - self.local_port_allocator - .deallocate_port(socket.endpoint().port); - socket.close(); - } - Protocol::Tcp => { - let tcp_specific = specific.tcp_mut(); - if let Some(server_socket) = tcp_specific.server_socket.take() { - // remove all listening sockets in the backlog - for handle in server_socket.socket_set_handles { - let _ = self.socket_set.remove(handle); - } - } - if let Some(local_port) = tcp_specific.local_port.take() { - self.local_port_allocator.deallocate(local_port); - } - let tcp_socket: &mut tcp::Socket = self.socket_set.get_mut(handle); - if tcp_specific.immediate_close.load(Ordering::Relaxed) { - tcp_socket.abort(); - } else { - tcp_socket.close(); - } - self.closing_in_background.push(handle); - } - } - if let Some(proxy) = proxy { - proxy.set_state(socket_channel::SocketState::Closed); - } - self.automated_platform_interaction(PollDirection::Both); } - /// Initiate a connection to an IP address + /// Initiate a connection to an IP address. /// /// When `check_progress` is false, this function attempts to initiate a connection. /// Otherwise, this function checks the progress of an ongoing connection. @@ -1094,182 +270,49 @@ where }; let descriptor_table = self.litebox.descriptor_table(); - let mut table_entry = descriptor_table - .get_entry_mut(fd) + let entry = descriptor_table + .get_entry(fd) .ok_or(ConnectError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; - let now = self.now(); - let ret = match socket_handle.protocol() { - Protocol::Tcp => { - if let Some(BrokerSocket::Tcp(socket)) = &socket_handle.broker_socket { - if check_progress { - socket.check_connect_progress() - } else { - socket.start_connect(*addr) - } + match &entry.entry.broker_socket { + BrokerSocket::Tcp(socket) => { + if check_progress { + socket.check_connect_progress() } else { - let check_state = |state: tcp::State| -> Result<(), ConnectError> { - match state { - tcp::State::Established => { - // already connected - Ok(()) - } - tcp::State::Closed | tcp::State::TimeWait => { - Err(ConnectError::InvalidState) - } - tcp::State::SynSent => Err(ConnectError::InProgress), - s => unimplemented!("state: {:?}", s), - } - }; - - let socket: &mut tcp::Socket = - self.socket_set.get_mut(socket_handle.smoltcp_handle()); - if check_progress { - check_state(socket.state()) - } else { - let local_port = self.local_port_allocator.ephemeral_port()?; - let local_endpoint: smoltcp::wire::IpListenEndpoint = - local_port.port().into(); - let addr: smoltcp::wire::IpEndpoint = (*addr).into(); - match socket.connect(self.interface.context(), addr, local_endpoint) { - Ok(()) => { - socket.set_timeout(Some(TCP_CONNECT_TIMEOUT)); - let tcp_specific = socket_handle.tcp_mut(); - tcp_specific.connect_initiated_at_us = Some(now); - let old_port = tcp_specific.local_port.replace(local_port); - if old_port.is_some() { - // Need to think about how to handle this situation - unimplemented!() - } - check_state(socket.state()) - } - Err(tcp::ConnectError::InvalidState) => unreachable!(), - Err(tcp::ConnectError::Unaddressable) => { - self.local_port_allocator.deallocate(local_port); - Err(ConnectError::Unaddressable) - } - } - } + socket.start_connect(*addr) } } - Protocol::Udp => { - if let Some(BrokerSocket::Udp(socket)) = &socket_handle.broker_socket { - if check_progress { - socket.check_connect_progress() - } else { - socket.start_connect(*addr) - } + BrokerSocket::Udp(socket) => { + if check_progress { + socket.check_connect_progress() } else { - if addr.port() == 0 { - return Err(ConnectError::Unaddressable); - } - let socket: &mut udp::Socket = - self.socket_set.get_mut(socket_handle.smoltcp_handle()); - if !socket.is_open() { - let local_port = self.local_port_allocator.ephemeral_port()?; - let local_endpoint: smoltcp::wire::IpListenEndpoint = - local_port.port().into(); - let Ok(()) = socket.bind(local_endpoint) else { - unreachable!("binding to a free port cannot fail") - }; - } - let addr: smoltcp::wire::IpEndpoint = (*addr).into(); - socket_handle.udp_mut().remote_endpoint = Some(addr); - Ok(()) + socket.start_connect(*addr) } } - Protocol::Icmp => unimplemented!(), - Protocol::Raw { protocol: _ } => unimplemented!(), - }; - - let mut result = ret; - if let Some(proxy) = &socket_handle.proxy { - match ret { - Ok(()) => proxy.set_state(socket_channel::SocketState::Connected), - Err(ConnectError::InProgress) => { - proxy.set_state(socket_channel::SocketState::Connecting); - } - Err(ConnectError::InvalidState) - if matches!(socket_handle.protocol(), Protocol::Tcp) => - { - // Distinguish timeout from RST using elapsed time - match socket_handle.tcp().connect_initiated_at_us { - Some(initiated_at) if now - initiated_at >= TCP_CONNECT_TIMEOUT => { - proxy.set_async_error(errors::SocketAsyncError::TimedOut); - result = Err(ConnectError::TimedOut); - } - _ => proxy.set_async_error(errors::SocketAsyncError::ConnectionRefused), - } - } - Err(ConnectError::Unaddressable | ConnectError::InvalidState) => { - proxy.set_async_error(errors::SocketAsyncError::ConnectionRefused); - } - Err(_) => {} - } } - drop(table_entry); - drop(descriptor_table); - - self.automated_platform_interaction(PollDirection::Both); - result } /// Get the local address and port a socket is bound to. pub fn get_local_addr(&self, fd: &SocketFd) -> Result { let descriptor_table = self.litebox.descriptor_table(); - let mut table_entry = descriptor_table - .get_entry_mut(fd) + let entry = descriptor_table + .get_entry(fd) .ok_or(LocalAddrError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; - - if let Some(socket) = &socket_handle.broker_socket { - return Ok(match socket { - BrokerSocket::Tcp(socket) => socket.local_addr(), - BrokerSocket::Udp(socket) => socket.local_addr(), - }); - } - - match socket_handle.protocol() { - Protocol::Tcp => { - let socket: &tcp::Socket = self.socket_set.get(socket_handle.smoltcp_handle()); - match socket.local_endpoint() { - Some(endpoint) => match endpoint.addr { - smoltcp::wire::IpAddress::Ipv4(ipv4) => { - Ok(SocketAddr::V4(SocketAddrV4::new(ipv4, endpoint.port))) - } - }, - None => Ok(SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0))), - } - } - Protocol::Udp => { - let socket: &udp::Socket = self.socket_set.get(socket_handle.smoltcp_handle()); - let local_endpoint = socket.endpoint(); - match local_endpoint.addr { - Some(smoltcp::wire::IpAddress::Ipv4(ipv4)) => { - Ok(SocketAddr::V4(SocketAddrV4::new(ipv4, local_endpoint.port))) - } - None => Ok(SocketAddr::V4(SocketAddrV4::new( - Ipv4Addr::UNSPECIFIED, - local_endpoint.port, - ))), - } - } - Protocol::Icmp => unimplemented!(), - Protocol::Raw { protocol: _ } => unimplemented!(), - } + Ok(match &entry.entry.broker_socket { + BrokerSocket::Tcp(socket) => socket.local_addr(), + BrokerSocket::Udp(socket) => socket.local_addr(), + }) } /// Get the remote address and port a socket is connected to, if any. pub fn get_remote_addr(&self, fd: &SocketFd) -> Result { let descriptor_table = self.litebox.descriptor_table(); - let mut table_entry = descriptor_table - .get_entry_mut(fd) + let entry = descriptor_table + .get_entry(fd) .ok_or(RemoteAddrError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; - self.get_remote_addr_for_handle(socket_handle) + Self::get_remote_addr_for_handle(&entry.entry) } - /// Shuts down one or both directions of a broker-owned socket. + /// Shuts down one or both directions of a socket. pub fn shutdown( &self, fd: &SocketFd, @@ -1279,17 +322,12 @@ where let entry = descriptor_table .get_entry(fd) .ok_or(errors::ShutdownError::InvalidFd)?; - let socket = entry - .entry - .broker_socket - .as_ref() - .ok_or(errors::ShutdownError::UnsupportedOperation)?; let mode = match direction { ShutdownDirection::Read => litebox_broker_protocol::socket::ShutdownMode::Read, ShutdownDirection::Write => litebox_broker_protocol::socket::ShutdownMode::Write, ShutdownDirection::Both => litebox_broker_protocol::socket::ShutdownMode::Both, }; - match socket { + match &entry.entry.broker_socket { BrokerSocket::Tcp(socket) => { if socket.is_listening() { return Err(errors::ShutdownError::Listening); @@ -1301,18 +339,13 @@ where .map_err(errors::ShutdownError::OperationFailed) } - /// Stops a broker-owned listener without closing its socket object. + /// Stops a listener without closing its socket object. pub fn stop_listening(&self, fd: &SocketFd) -> Result<(), errors::ShutdownError> { let descriptor_table = self.litebox.descriptor_table(); let entry = descriptor_table .get_entry(fd) .ok_or(errors::ShutdownError::InvalidFd)?; - let socket = entry - .entry - .broker_socket - .as_ref() - .ok_or(errors::ShutdownError::UnsupportedOperation)?; - let BrokerSocket::Tcp(socket) = socket else { + let BrokerSocket::Tcp(socket) = &entry.entry.broker_socket else { return Err(errors::ShutdownError::UnsupportedOperation); }; if !socket.is_listening() { @@ -1325,38 +358,16 @@ where .map_err(errors::ShutdownError::OperationFailed) } - /// Get the remote address and port a `SocketHandle` is connected to, if any. fn get_remote_addr_for_handle( - &self, socket_handle: &SocketHandle, ) -> Result { - if let Some(socket) = &socket_handle.broker_socket { - return match socket { - BrokerSocket::Tcp(socket) => socket.remote_addr(), - BrokerSocket::Udp(socket) => socket.remote_addr(), - }; - } - let endpoint = match socket_handle.protocol() { - Protocol::Tcp => self - .socket_set - .get::(socket_handle.smoltcp_handle()) - .remote_endpoint() - .ok_or(RemoteAddrError::NotConnected)?, - Protocol::Udp => socket_handle - .udp() - .remote_endpoint - .ok_or(RemoteAddrError::NotConnected)?, - Protocol::Icmp => unimplemented!(), - Protocol::Raw { protocol: _ } => unimplemented!(), - }; - match endpoint.addr { - smoltcp::wire::IpAddress::Ipv4(ipv4) => { - Ok(SocketAddr::V4(SocketAddrV4::new(ipv4, endpoint.port))) - } + match &socket_handle.broker_socket { + BrokerSocket::Tcp(socket) => socket.remote_addr(), + BrokerSocket::Udp(socket) => socket.remote_addr(), } } - /// Bind a socket to a specific address and port. If the port is 0, an ephemeral port is allocated. + /// Bind a socket to a specific address and port. pub fn bind( &mut self, fd: &SocketFd, @@ -1367,292 +378,82 @@ where }; let descriptor_table = self.litebox.descriptor_table(); - let mut table_entry = descriptor_table - .get_entry_mut(fd) + let entry_handle = descriptor_table + .entry_handle(fd) .ok_or(BindError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; - if let Some(socket) = socket_handle.broker_socket.clone() { - drop(table_entry); - drop(descriptor_table); - return match socket { - BrokerSocket::Tcp(socket) => socket.bind(*addr), - BrokerSocket::Udp(socket) => socket.bind(*addr), - }; - } - match socket_handle.protocol() { - Protocol::Tcp => { - if socket_handle.tcp().server_socket.is_some() { - return Err(BindError::AlreadyBound); - } - let lp = self - .local_port_allocator - .allocate_local_port(addr.port()) - .map_err(|_| BindError::PortAlreadyInUse(addr.port()))?; - let new_port = lp.port(); - let old_lp = socket_handle.tcp_mut().local_port.replace(lp); - if let Some(old) = old_lp { - self.local_port_allocator.deallocate(old); - // Currently unsure if the dealloc is sufficient and if we need to do - // anything else here (possibly return an error message due to trying to - // do things to a connected socket, not sure), so just marking as - // unimplemented for now to trigger a panic. - unimplemented!() - } - socket_handle.tcp_mut().server_socket = Some(TcpServerSpecific { - ip_listen_endpoint: smoltcp::wire::IpListenEndpoint { - addr: Some(smoltcp::wire::IpAddress::Ipv4(*addr.ip())), - port: new_port, - }, - backlog: None, - socket_set_handles: vec![], - }); - } - Protocol::Udp => { - let lp = self - .local_port_allocator - .allocate_local_port(addr.port()) - .map_err(|_| BindError::PortAlreadyInUse(addr.port()))?; - let local_endpoint = smoltcp::wire::IpListenEndpoint { - addr: Some(smoltcp::wire::IpAddress::Ipv4(*addr.ip())), - port: lp.port(), - }; - let socket: &mut udp::Socket = - self.socket_set.get_mut(socket_handle.smoltcp_handle()); - if let Err(e) = socket.bind(local_endpoint) { - self.local_port_allocator.deallocate(lp); - return Err(match e { - udp::BindError::InvalidState => BindError::AlreadyBound, - udp::BindError::Unaddressable => unreachable!(), - }); - } - } - Protocol::Icmp => unimplemented!(), - Protocol::Raw { protocol: _ } => unimplemented!(), - } - - drop(table_entry); + let entry = descriptor_table.get_entry(fd).ok_or(BindError::InvalidFd)?; + let socket = entry.entry.broker_socket.clone(); + drop(entry); drop(descriptor_table); - - self.automated_platform_interaction(PollDirection::Both); - Ok(()) + let result = match socket { + BrokerSocket::Tcp(socket) => socket.bind(*addr), + BrokerSocket::Udp(socket) => socket.bind(*addr), + }; + drop(entry_handle); + result } - /// Prepare a socket to accept incoming connections. Marks the socket as a passive socket, such - /// that it will be used to accept new connection requests via [`accept`](Self::accept). - /// - /// The `backlog` argument defines the maximum length to which the queue of pending connections - /// the `fd` may grow. This function is allowed to silently cap the value to a reasonable upper - /// bound. + /// Prepare a socket to accept incoming connections. pub fn listen(&mut self, fd: &SocketFd, backlog: u16) -> Result<(), ListenError> { let descriptor_table = self.litebox.descriptor_table(); - let mut table_entry = descriptor_table - .get_entry_mut(fd) + let entry_handle = descriptor_table + .entry_handle(fd) .ok_or(ListenError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; - if let Some(socket) = socket_handle.broker_socket.as_ref() { - let BrokerSocket::Tcp(socket) = socket else { - return Err(ListenError::UnsupportedOperation); - }; - let socket = alloc::sync::Arc::clone(socket); - drop(table_entry); - drop(descriptor_table); - return socket.listen( - u32::from(backlog).min(litebox_broker_protocol::socket::MAX_TCP_LISTEN_BACKLOG), - ); - } - if backlog == 0 { - // What should actually happen here? - unimplemented!() - } - - // This prevents users from overloading things too badly; 4096 is the upper limit with - // similar silent-cap behavior since Linux 5.4 (earlier versions capped even smaller, at - // 128, but we use the larger value to be more flexible). - // - // TODO: smoltcp performs a linear search through SocketSet when dispatching an incoming - // packet to the socket it belongs to, so having a large backlog can cause performance issues - // (see https://github.com/smoltcp-rs/smoltcp/issues/973). Restricting the backlog to a smaller - // value for now until we have a better solution. - let backlog = backlog.min(8); - - match &mut socket_handle.specific { - ProtocolSpecific::Tcp(handle) => { - if handle.server_socket.is_none() { - let local_port = - self.local_port_allocator - .ephemeral_port() - .map_err(|e| match e { - local_ports::LocalPortAllocationError::AlreadyInUse(_) => { - unreachable!() - } - local_ports::LocalPortAllocationError::NoAvailableFreePorts => { - ListenError::NoAvailableFreeEphemeralPorts - } - })?; - let port = local_port.port(); - let old_local_port = handle.local_port.replace(local_port); - if let Some(lp) = old_local_port { - self.local_port_allocator.deallocate(lp); - // Should anything else be done here? - unimplemented!() - } - handle.server_socket = Some(TcpServerSpecific { - ip_listen_endpoint: smoltcp::wire::IpListenEndpoint { - addr: Some(smoltcp::wire::IpAddress::v4(0, 0, 0, 0)), - port, - }, - backlog: None, - socket_set_handles: vec![], - }); - } - let Some(server_socket) = &mut handle.server_socket else { - unreachable!() - }; - if server_socket.ip_listen_endpoint.port == 0 { - return Err(ListenError::InvalidAddress); - } - if server_socket.backlog.is_some() || !server_socket.socket_set_handles.is_empty() { - // Need to change the amount of backlog; growing will just work, but truncating - // might need some effort to pick which ones to keep/drop - unimplemented!() - } else { - server_socket.backlog = Some(backlog); - server_socket.socket_set_handles = Vec::with_capacity(backlog.into()); - } - server_socket.refill_to_backlog(&mut self.socket_set); - } - ProtocolSpecific::Udp(_) => unimplemented!(), - ProtocolSpecific::Icmp(_) => unimplemented!(), - ProtocolSpecific::Raw(_) => unimplemented!(), - } - - if let Some(proxy) = &socket_handle.proxy { - proxy.set_state(socket_channel::SocketState::Listening); - } - - drop(table_entry); + let entry = descriptor_table + .get_entry(fd) + .ok_or(ListenError::InvalidFd)?; + let BrokerSocket::Tcp(socket) = &entry.entry.broker_socket else { + return Err(ListenError::UnsupportedOperation); + }; + let socket = Arc::clone(socket); + drop(entry); drop(descriptor_table); - - self.automated_platform_interaction(PollDirection::Ingress); - Ok(()) + let result = socket.listen( + u32::from(backlog).min(litebox_broker_protocol::socket::MAX_TCP_LISTEN_BACKLOG), + ); + drop(entry_handle); + result } /// Accept a new incoming connection on a listening socket. /// /// If `peer` is provided, it is filled with the remote address of the accepted connection. /// - /// Note that the returned new socket has no associated proxy; to set a proxy, use - /// [`attach_socket_proxy`](Self::attach_socket_proxy). + /// The returned socket has no associated proxy; use + /// [`attach_socket_proxy`](Self::attach_socket_proxy) to attach one. pub fn accept( &mut self, fd: &SocketFd, peer: Option<&mut SocketAddr>, ) -> Result, AcceptError> { - self.automated_platform_interaction(PollDirection::Both); let descriptor_table = self.litebox.descriptor_table(); - let mut table_entry = descriptor_table - .get_entry_mut(fd) + let entry_handle = descriptor_table + .entry_handle(fd) .ok_or(AcceptError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; - if let Some(listener) = socket_handle.broker_socket.as_ref() { - let BrokerSocket::Tcp(listener) = listener else { - return Err(AcceptError::UnsupportedOperation); - }; - let listener = alloc::sync::Arc::clone(listener); - drop(table_entry); - drop(descriptor_table); - let accepted = listener.accept()?; - if let Some(peer) = peer { - *peer = accepted.remote_addr().map_err(|_| { - AcceptError::OperationFailed(errors::SocketAsyncError::BackendFailure) - })?; - } - return Ok(self.new_socket_fd_for(SocketHandle { - consider_closed: false, - handle: None, - broker_socket: Some(BrokerSocket::Tcp(accepted)), - specific: ProtocolSpecific::Tcp(TcpSpecific { - local_port: None, - server_socket: None, - immediate_close: AtomicBool::new(false), - connect_initiated_at_us: None, - }), - proxy: None, - })); - } - match &mut socket_handle.specific { - ProtocolSpecific::Tcp(handle) => { - let Some(server_socket) = &mut handle.server_socket else { - return Err(AcceptError::NotListening); - }; - if server_socket.backlog.is_none() { - return Err(AcceptError::NotListening); - } - // (Purely an optimization) remove all handles that are closed, by only keeping ones - // that are not closed - server_socket.socket_set_handles.retain(|&h| { - let socket: &tcp::Socket = self.socket_set.get(h); - socket.is_open() - }); - // Find a socket that has progressed further in its TCP state machine, by finding a - // socket in an established state - let Some(position) = server_socket.socket_set_handles.iter().position(|&h| { - let socket: &tcp::Socket = self.socket_set.get(h); - socket.state() == tcp::State::Established - }) else { - if let Some(proxy) = &socket_handle.proxy { - // No connections are ready; make sure the readable flag is cleared - proxy.set_readable(false); - } - return Err(AcceptError::NoConnectionsReady); - }; - if let Some(proxy) = &socket_handle.proxy { - // reset the readable flag so that we send one [`Events::In`] event per accepted connection - proxy.set_readable(false); - } - // Pull that position out of the listening handles - let ready_handle = server_socket.socket_set_handles.swap_remove(position); - // Refill to the backlog, so that we can have more listening sockets again if needed - server_socket.refill_to_backlog(&mut self.socket_set); - // Grab the local port again, so we can put it into the new `TcpSpecific` - let local_port = handle - .local_port - .as_ref() - .map(|lp| self.local_port_allocator.allocate_same_local_port(lp)); - // Release the locks, needed to be able to use `self` below - drop(table_entry); - drop(descriptor_table); - // Create a new FD to hand it back out to the user - let handle = SocketHandle { - consider_closed: false, - handle: Some(ready_handle), - broker_socket: None, - specific: ProtocolSpecific::Tcp(TcpSpecific { - local_port, - server_socket: None, - immediate_close: AtomicBool::new(false), - connect_initiated_at_us: None, - }), - proxy: None, - }; - if let Some(peer) = peer { - let Ok(remote_addr) = self.get_remote_addr_for_handle(&handle) else { - unreachable!("a connected TCP socket must have a remote address") - }; - *peer = remote_addr; - } - Ok(self.new_socket_fd_for(handle)) - } - ProtocolSpecific::Udp(_) => unimplemented!(), - ProtocolSpecific::Icmp(_) => unimplemented!(), - ProtocolSpecific::Raw(_) => unimplemented!(), - } + let entry = descriptor_table + .get_entry(fd) + .ok_or(AcceptError::InvalidFd)?; + let BrokerSocket::Tcp(listener) = &entry.entry.broker_socket else { + return Err(AcceptError::UnsupportedOperation); + }; + let listener = Arc::clone(listener); + drop(entry); + drop(descriptor_table); + let accepted = listener.accept()?; + drop(entry_handle); + if let Some(peer) = peer { + *peer = accepted.remote_addr().map_err(|_| { + AcceptError::OperationFailed(errors::SocketAsyncError::BackendFailure) + })?; + } + Ok(self.litebox.descriptor_table_mut().insert(SocketHandle { + broker_socket: BrokerSocket::Tcp(accepted), + immediate_close: AtomicBool::new(false), + proxy: None, + })) } /// Send data over a socket, optionally specifying the destination address. - /// - /// If the socket is connection-mode and the destination address is provided, - /// `Err(SendError::UnnecessaryDestinationAddress)` is returned. pub fn send( &mut self, fd: &SocketFd, @@ -1660,100 +461,35 @@ where flags: SendFlags, destination: Option, ) -> Result { - let descriptor_table = self.litebox.descriptor_table(); - let mut table_entry = descriptor_table - .get_entry_mut(fd) - .ok_or(SendError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; if !flags.is_empty() { unimplemented!() } - if let Some(socket) = &socket_handle.broker_socket { - let outcome = match socket { - BrokerSocket::Tcp(socket) => { - if destination.is_some() { - return Err(SendError::UnnecessaryDestinationAddress); - } - socket.try_write(buf) - } - BrokerSocket::Udp(socket) => socket.try_write(buf, destination), - }; - return outcome.map_err(|error| match error { - socket_channel::ChannelWriteError::BufferFull => SendError::BufferFull, - socket_channel::ChannelWriteError::MessageTooLong => SendError::MessageTooLong, - socket_channel::ChannelWriteError::Unaddressable => SendError::Unaddressable, - socket_channel::ChannelWriteError::DestinationAddressRequired => { - SendError::DestinationAddressRequired - } - socket_channel::ChannelWriteError::WriteShutdown - | socket_channel::ChannelWriteError::NotConnected - | socket_channel::ChannelWriteError::ConnectionClosed - | socket_channel::ChannelWriteError::Socket(_) => SendError::SocketInInvalidState, - }); - } - - let ret = match socket_handle.protocol() { - Protocol::Tcp => { + let descriptor_table = self.litebox.descriptor_table(); + let entry = descriptor_table.get_entry(fd).ok_or(SendError::InvalidFd)?; + let outcome = match &entry.entry.broker_socket { + BrokerSocket::Tcp(socket) => { if destination.is_some() { - // TCP is connection-oriented, so no destination address should be provided return Err(SendError::UnnecessaryDestinationAddress); } - self.socket_set - .get_mut::(socket_handle.smoltcp_handle()) - .send_slice(buf) - .map_err(|tcp::SendError::InvalidState| SendError::SocketInInvalidState) - } - Protocol::Udp => { - let destination = destination - .map(|s| match s { - SocketAddr::V4(addr) => smoltcp::wire::IpEndpoint::from(addr), - SocketAddr::V6(_) => unimplemented!(), - }) - .or_else(|| socket_handle.udp().remote_endpoint); - let Some(remote_endpoint) = destination else { - return Err(SendError::DestinationAddressRequired); - }; - let udp_socket: &mut udp::Socket = - self.socket_set.get_mut(socket_handle.smoltcp_handle()); - if !udp_socket.is_open() { - let local_port = self - .local_port_allocator - .ephemeral_port() - .map_err(SendError::PortAllocationFailure)?; - let port = local_port.port(); - let Ok(()) = - udp_socket.bind(smoltcp::wire::IpListenEndpoint { addr: None, port }) - else { - self.local_port_allocator.deallocate(local_port); - unreachable!("binding to a free port cannot fail") - }; - } - udp_socket - .send_slice(buf, remote_endpoint) - .map(|()| buf.len()) - .map_err(|e| match e { - udp::SendError::BufferFull => SendError::BufferFull, - udp::SendError::Unaddressable => SendError::Unaddressable, - }) + socket.try_write(buf) } - Protocol::Icmp => unimplemented!(), - Protocol::Raw { protocol: _ } => unimplemented!(), + BrokerSocket::Udp(socket) => socket.try_write(buf, destination), }; - - drop(table_entry); - drop(descriptor_table); - - self.automated_platform_interaction(PollDirection::Egress); - ret + outcome.map_err(|error| match error { + socket_channel::ChannelWriteError::BufferFull => SendError::BufferFull, + socket_channel::ChannelWriteError::MessageTooLong => SendError::MessageTooLong, + socket_channel::ChannelWriteError::Unaddressable => SendError::Unaddressable, + socket_channel::ChannelWriteError::DestinationAddressRequired => { + SendError::DestinationAddressRequired + } + socket_channel::ChannelWriteError::WriteShutdown + | socket_channel::ChannelWriteError::NotConnected + | socket_channel::ChannelWriteError::ConnectionClosed + | socket_channel::ChannelWriteError::Socket(_) => SendError::SocketInInvalidState, + }) } - /// Receive data from a connected socket. - /// - /// If the `source_addr` is `Some` and the underlying protocol provides a source address, it will be updated. - /// e.g., UDP does provide the source address, while TCP does not (because it is connection-oriented, - /// once it is established, both ends should already know each other's addresses). - /// - /// On success, returns the number of bytes received. + /// Receive data from a socket. pub fn receive( &mut self, fd: &SocketFd, @@ -1761,122 +497,39 @@ where flags: ReceiveFlags, source_addr: Option<&mut Option>, ) -> Result { - // Note that we do an earlier-than-usual automated interaction to ingress packets since it - // doesn't hurt to do this too often (other than wasting energy), and this allows us to - // possibly get packets where we might otherwise return with size 0 on the `receive`. - self.automated_platform_interaction(PollDirection::Ingress); - let descriptor_table = self.litebox.descriptor_table(); - let mut table_entry = descriptor_table - .get_entry_mut(fd) - .ok_or(ReceiveError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; if flags.intersects( (ReceiveFlags::DONTWAIT | ReceiveFlags::TRUNC | ReceiveFlags::DISCARD).complement(), ) { unimplemented!("flags: {:?}", flags); } - if let Some(socket) = &socket_handle.broker_socket { - let outcome = match socket { - BrokerSocket::Tcp(socket) => socket.try_read(buf, flags, source_addr), - BrokerSocket::Udp(socket) => { - socket.try_read(buf, flags, source_addr).map(|received| { - if flags.intersects(ReceiveFlags::TRUNC | ReceiveFlags::DISCARD) { - received - } else { - received.min(buf.len()) - } - }) - } - }; - return outcome.or_else(|error| match error { - socket_channel::ChannelReadError::WouldBlock => Ok(0), - socket_channel::ChannelReadError::ReadShutdown - | socket_channel::ChannelReadError::ConnectionClosed => { - Err(ReceiveError::OperationFinished) - } - socket_channel::ChannelReadError::NotConnected - | socket_channel::ChannelReadError::Socket(_) => { - Err(ReceiveError::SocketInInvalidState) - } - }); - } - - let ret = match socket_handle.protocol() { - Protocol::Tcp => { - if let Some(source_addr) = source_addr { - // TCP is connection-oriented, so no need to provide a source address - *source_addr = None; - } - let tcp_socket = self - .socket_set - .get_mut::(socket_handle.smoltcp_handle()); - if flags.contains(ReceiveFlags::TRUNC) { - unimplemented!("TRUNC flag for tcp"); - } - if flags.contains(ReceiveFlags::DISCARD) { - let discard_slice = - |tcp_socket: &mut tcp::Socket<'_>| -> Result { - // See [`tcp::Socket::recv_slice`] and [`tcp::Socket::recv`] for why we do two `recv` calls. - // Basically, the socket buffer is implemented as a ring buffer, and if the data to be read - // wraps around, a single `recv` call will not be able to read all the data. - let size1 = tcp_socket.recv(|data| (data.len(), data.len()))?; - let size2 = tcp_socket.recv(|data| (data.len(), data.len()))?; - Ok(size1 + size2) - }; - discard_slice(tcp_socket) + let descriptor_table = self.litebox.descriptor_table(); + let entry = descriptor_table + .get_entry(fd) + .ok_or(ReceiveError::InvalidFd)?; + let outcome = match &entry.entry.broker_socket { + BrokerSocket::Tcp(socket) => socket.try_read(buf, flags, source_addr), + BrokerSocket::Udp(socket) => socket.try_read(buf, flags, source_addr).map(|received| { + if flags.intersects(ReceiveFlags::TRUNC | ReceiveFlags::DISCARD) { + received } else { - tcp_socket.recv_slice(buf) + received.min(buf.len()) } - .map_err(|e| match e { - tcp::RecvError::InvalidState => ReceiveError::SocketInInvalidState, - tcp::RecvError::Finished => ReceiveError::OperationFinished, - }) + }), + }; + outcome.or_else(|error| match error { + socket_channel::ChannelReadError::WouldBlock => Ok(0), + socket_channel::ChannelReadError::ReadShutdown + | socket_channel::ChannelReadError::ConnectionClosed => { + Err(ReceiveError::OperationFinished) } - Protocol::Udp => { - let udp_socket = self - .socket_set - .get_mut::(socket_handle.smoltcp_handle()); - match udp_socket.recv() { - Ok((data, meta)) => { - if let Some(source_addr) = source_addr { - let remote_addr = match meta.endpoint.addr { - smoltcp::wire::IpAddress::Ipv4(ipv4_addr) => { - SocketAddr::V4(SocketAddrV4::new(ipv4_addr, meta.endpoint.port)) - } - }; - *source_addr = Some(remote_addr); - } - let n = if flags.contains(ReceiveFlags::DISCARD) { - data.len() - } else { - let length = data.len().min(buf.len()); - buf[..length].copy_from_slice(&data[..length]); - if flags.contains(ReceiveFlags::TRUNC) { - // return the real size of the packet or datagram, - // even when it was longer than the passed buffer. - data.len() - } else { - length - } - }; - Ok(n) - } - Err(udp::RecvError::Exhausted) => Ok(0), - Err(udp::RecvError::Truncated) => unreachable!(), - } + socket_channel::ChannelReadError::NotConnected + | socket_channel::ChannelReadError::Socket(_) => { + Err(ReceiveError::SocketInInvalidState) } - Protocol::Icmp => unimplemented!(), - Protocol::Raw { protocol: _ } => unimplemented!(), - }; - - drop(table_entry); - drop(descriptor_table); - - self.automated_platform_interaction(PollDirection::Ingress); - ret + }) } - /// Set TCP options + /// Set TCP options. pub fn set_tcp_option( &mut self, fd: &SocketFd, @@ -1886,49 +539,21 @@ where let entry_handle = descriptor_table .entry_handle(fd) .ok_or(errors::SetTcpOptionError::InvalidFd)?; - let mut table_entry = descriptor_table - .get_entry_mut(fd) + let entry = descriptor_table + .get_entry(fd) .ok_or(errors::SetTcpOptionError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; - if let Some(BrokerSocket::Tcp(socket)) = &socket_handle.broker_socket { - let socket = alloc::sync::Arc::clone(socket); - drop(table_entry); - drop(descriptor_table); - return socket.set_tcp_option(data); - } + let BrokerSocket::Tcp(socket) = &entry.entry.broker_socket else { + return Err(errors::SetTcpOptionError::NotTcpSocket); + }; + let socket = Arc::clone(socket); + drop(entry); + drop(descriptor_table); + let result = socket.set_tcp_option(data); drop(entry_handle); - match socket_handle.protocol() { - Protocol::Tcp => { - let tcp_socket = self - .socket_set - .get_mut::(socket_handle.smoltcp_handle()); - match data { - TcpOptionData::NODELAY(nodelay) => { - tcp_socket.set_nagle_enabled(!nodelay); - } - TcpOptionData::KEEPALIVE(keepalive) => { - tcp_socket.set_keep_alive( - keepalive.then_some(smoltcp::time::Duration::from_secs(2 * 60 * 60)), - ); - } - TcpOptionData::KEEPINTVL(keepalive) => { - tcp_socket.set_keep_alive(keepalive.map(smoltcp::time::Duration::from)); - } - TcpOptionData::CONGESTION(congestion) => match congestion { - CongestionControl::None => { - tcp_socket.set_congestion_control(tcp::CongestionControl::None); - } - _ => unimplemented!(), - }, - } - Ok(()) - } - Protocol::Udp | Protocol::Icmp | Protocol::Raw { .. } => { - Err(errors::SetTcpOptionError::NotTcpSocket) - } - } + result } - /// Get TCP options + + /// Get TCP options. pub fn get_tcp_option( &self, fd: &SocketFd, @@ -1938,47 +563,22 @@ where let entry_handle = descriptor_table .entry_handle(fd) .ok_or(errors::GetTcpOptionError::InvalidFd)?; - let mut table_entry = descriptor_table - .get_entry_mut(fd) + let entry = descriptor_table + .get_entry(fd) .ok_or(errors::GetTcpOptionError::InvalidFd)?; - let socket_handle = &mut table_entry.entry; - if let Some(BrokerSocket::Tcp(socket)) = &socket_handle.broker_socket { - let socket = alloc::sync::Arc::clone(socket); - drop(table_entry); - drop(descriptor_table); - return socket.get_tcp_option(name); - } + let BrokerSocket::Tcp(socket) = &entry.entry.broker_socket else { + return Err(errors::GetTcpOptionError::NotTcpSocket); + }; + let socket = Arc::clone(socket); + drop(entry); + drop(descriptor_table); + let result = socket.get_tcp_option(name); drop(entry_handle); - match socket_handle.protocol() { - Protocol::Tcp => { - let tcp_socket = self - .socket_set - .get::(socket_handle.smoltcp_handle()); - match name { - TcpOptionName::NODELAY => { - Ok(TcpOptionData::NODELAY(!tcp_socket.nagle_enabled())) - } - TcpOptionName::KEEPALIVE => { - Ok(TcpOptionData::KEEPALIVE(tcp_socket.keep_alive().is_some())) - } - TcpOptionName::KEEPINTVL => Ok(TcpOptionData::KEEPINTVL( - tcp_socket.keep_alive().map(core::time::Duration::from), - )), - TcpOptionName::CONGESTION => Ok(TcpOptionData::CONGESTION( - match tcp_socket.congestion_control() { - tcp::CongestionControl::None => CongestionControl::None, - }, - )), - } - } - Protocol::Udp | Protocol::Icmp | Protocol::Raw { .. } => { - Err(errors::GetTcpOptionError::NotTcpSocket) - } - } + result } } -/// Protocols for sockets supported by LiteBox +/// Protocols for sockets supported by LiteBox. #[non_exhaustive] pub enum Protocol { Tcp, @@ -2031,25 +631,20 @@ bitflags! { } } -/// Socket options for TCP +/// Socket options for TCP. #[non_exhaustive] pub enum TcpOptionName { - /// If set, disable the Nagle algorithm. This means that - /// segments are always sent as soon as possible, even if there - /// is only a small amount of data. + /// If set, disable the Nagle algorithm. NODELAY, /// Enable sending of keep-alive messages. KEEPALIVE, /// Interval between keep-alive probes. KEEPINTVL, - /// TCP congestion control algorithm + /// TCP congestion control algorithm. CONGESTION, } -/// Data for TCP options -/// -/// Note it should be paired with the correct [`TcpOptionName`] variant. -/// For example, `TcpOptionName::NODELAY` should be paired with `TcpOptionData::NODELAY(true)`. +/// Data for TCP options. #[non_exhaustive] pub enum TcpOptionData { NODELAY(bool), @@ -2058,7 +653,7 @@ pub enum TcpOptionData { CONGESTION(CongestionControl), } -/// TCP Congestion Control Algorithms +/// TCP congestion control algorithms. #[non_exhaustive] pub enum CongestionControl { None, @@ -2070,15 +665,16 @@ pub enum CongestionControl { pub enum CloseBehavior { /// Close the socket immediately (i.e., abortive close). Immediate, - /// Close the socket in background and return immediately + /// Close the socket gracefully. Graceful, - /// Close the socket in background only if there is not unsent data remaining, - /// else return an error. + /// Close gracefully only if no data is pending transmission. + /// + /// Broker socket operations are synchronous, so this is equivalent to [`Self::Graceful`]. GracefulIfNoPendingData, } crate::fd::enable_fds_for_subsystem! { - @Platform: { platform::IPInterfaceProvider + platform::TimeProvider + sync::RawSyncPrimitivesProvider }; + @Platform: { platform::TimeProvider + sync::RawSyncPrimitivesProvider }; Network; @Platform: { platform::TimeProvider + sync::RawSyncPrimitivesProvider }; SocketHandle; diff --git a/litebox/src/net/phy.rs b/litebox/src/net/phy.rs deleted file mode 100644 index ff4d85c6e9..0000000000 --- a/litebox/src/net/phy.rs +++ /dev/null @@ -1,103 +0,0 @@ -// Copyright (c) Microsoft Corporation. -// Licensed under the MIT license. - -//! Connection to the physical (i.e., "lower") side for networking. - -// TODO(jayb): Do we need to wrap/unwrap the IPv4 header here, or is a better place within the -// implementer of the `platform::IPInterfaceProvider` trait? - -use crate::platform; - -/// The maximum transmission unit for a device -pub(crate) const DEVICE_MTU: usize = 1600; - -pub(crate) struct Device { - pub(crate) platform: &'static Platform, - receive_buffer: [u8; DEVICE_MTU], - send_buffer: [u8; DEVICE_MTU], -} - -impl Device { - pub(crate) fn new(platform: &'static Platform) -> Self { - Self { - platform, - receive_buffer: [0u8; DEVICE_MTU], - send_buffer: [0u8; DEVICE_MTU], - } - } -} - -impl smoltcp::phy::Device for Device { - type RxToken<'a> - = RxToken<'a> - where - Self: 'a; - type TxToken<'a> - = TxToken<'a, Platform> - where - Self: 'a; - - fn receive( - &mut self, - _timestamp: smoltcp::time::Instant, - ) -> Option<(Self::RxToken<'_>, Self::TxToken<'_>)> { - match self.platform.receive_ip_packet(&mut self.receive_buffer) { - Ok(size) => Some(( - RxToken { - buffer: &self.receive_buffer[..size], - }, - TxToken { - platform: self.platform, - buffer: &mut self.send_buffer, - }, - )), - Err(platform::ReceiveError::WouldBlock) => None, - } - } - - fn transmit(&mut self, _timestamp: smoltcp::time::Instant) -> Option> { - Some(TxToken { - platform: self.platform, - buffer: &mut self.send_buffer, - }) - } - - fn capabilities(&self) -> smoltcp::phy::DeviceCapabilities { - let mut caps = smoltcp::phy::DeviceCapabilities::default(); - caps.medium = smoltcp::phy::Medium::Ip; - caps.max_transmission_unit = DEVICE_MTU; - caps - } -} - -pub(crate) struct RxToken<'a> { - buffer: &'a [u8], -} - -impl smoltcp::phy::RxToken for RxToken<'_> { - fn consume(self, f: F) -> R - where - F: FnOnce(&[u8]) -> R, - { - f(self.buffer) - } -} - -pub(crate) struct TxToken<'a, Platform: platform::IPInterfaceProvider> { - platform: &'a Platform, - buffer: &'a mut [u8], -} - -impl smoltcp::phy::TxToken for TxToken<'_, Platform> { - fn consume(self, len: usize, f: F) -> R - where - F: FnOnce(&mut [u8]) -> R, - { - let packet = &mut self.buffer[..len]; - let res = f(packet); - self.platform - .send_ip_packet(packet) - .expect("Sending IP packet failed"); - res - } -} diff --git a/litebox/src/net/socket_channel.rs b/litebox/src/net/socket_channel.rs index 665adc2728..0dd535f049 100644 --- a/litebox/src/net/socket_channel.rs +++ b/litebox/src/net/socket_channel.rs @@ -1,131 +1,29 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT license. -//! Lock-free socket channels using ringbuf for decoupling user I/O from network processing. -//! -//! This module provides a channel-based design for socket data transfer that eliminates -//! lock contention between user threads (performing read/write) and the network worker -//! (processing packets via smoltcp). -//! -//! # Architecture -//! -//! ```text -//! SocketChannel -//! ┌──────────────────────────────────────────────────────────────┐ -//! │ │ -//! │ ┌────────────────────────────────────────────────────────┐ │ -//! │ │ RX Ring Buffer │ │ -//! │ │ (lock-free) │ │ -//! │ └────────────────────────────────────────────────────────┘ │ -//! │ ▲ │ │ -//! │ │ push │ pop │ -//! │ │ ▼ │ -//! │ Network Worker User Thread │ -//! │ (smoltcp) (read) │ -//! │ │ -//! │ ┌────────────────────────────────────────────────────────┐ │ -//! │ │ TX Ring Buffer │ │ -//! │ │ (lock-free) │ │ -//! │ └────────────────────────────────────────────────────────┘ │ -//! │ │ ▲ │ -//! │ │ pop │ push │ -//! │ ▼ │ │ -//! │ Network Worker User Thread │ -//! │ (smoltcp) (write) │ -//! │ │ -//! │ ┌────────────────────────────────────────────────────────┐ │ -//! │ │ State flags (atomic: ready, closed, error, etc.) │ │ -//! │ └────────────────────────────────────────────────────────┘ │ -//! └──────────────────────────────────────────────────────────────┘ -//! ``` -//! -//! # Benefits -//! -//! - **No lock contention**: User read/write operations and network processing can proceed -//! concurrently without blocking each other. +//! Socket I/O and readiness proxies. -use alloc::boxed::Box; -use core::{ - net::SocketAddr, - sync::atomic::{AtomicBool, AtomicU16, AtomicU32, AtomicUsize, Ordering}, -}; - -use ringbuf::{ - HeapCons, HeapProd, HeapRb, - traits::{Consumer as _, Observer as _, Producer as _, Split as _}, -}; +use core::net::SocketAddr; -use crate::sync::{Mutex, RawSyncPrimitivesProvider}; use crate::{ - event::{Events, IOPollable, observer::Observer, polling::Pollee}, - net::ReceiveFlags, + event::{Events, IOPollable, observer::Observer}, platform::TimeProvider, + sync::RawSyncPrimitivesProvider, }; -/// Generates common async socket error accessor methods for channel types -/// that contain an `inner` field with a `socket_error: SocketAsyncErrorState`. -macro_rules! impl_socket_async_error_accessors { - () => { - /// Set the async socket error. - pub(super) fn set_async_error(&self, error: super::errors::SocketAsyncError) { - self.inner.socket_error.set(error); - } - - /// Clear the async socket error. - #[allow(dead_code)] - pub(super) fn clear_async_error(&self) { - let _ = self.inner.socket_error.get(true); - } - - /// Read and optionally clear the async socket error. - fn get_async_error(&self, clear: bool) -> Option { - self.inner.socket_error.get(clear) - } - }; -} - -/// Atomic storage for [`SocketAsyncError`] -/// -/// [`SocketAsyncError`]: super::errors::SocketAsyncError -struct SocketAsyncErrorState { - /// Socket error stored as raw u32; 0 means no error. - value: AtomicU32, -} - -impl SocketAsyncErrorState { - fn new() -> Self { - Self { - value: AtomicU32::new(0), - } - } - - fn set(&self, error: super::errors::SocketAsyncError) { - self.value.store(error as u32, Ordering::Release); - } - - fn get(&self, clear: bool) -> Option { - let raw = if clear { - self.value.swap(0, Ordering::AcqRel) - } else { - self.value.load(Ordering::Acquire) - }; - super::errors::SocketAsyncError::from_u32(raw) - } -} - -/// Socket state flags stored atomically +/// Socket state used by the Linux shim when publishing state transitions. #[derive(Debug, Clone, Copy, PartialEq, Eq)] #[repr(u32)] pub enum SocketState { - /// Socket is in initial state. + /// Socket is in its initial state. Initial = 0, - /// Socket is connecting (TCP SYN sent) + /// Socket is connecting. Connecting = 1, - /// Socket is connected and ready for data transfer + /// Socket is connected and ready for data transfer. Connected = 2, - /// Socket is listening for incoming connections + /// Socket is listening for incoming connections. Listening = 3, - /// Socket encountered an error + /// Socket encountered an error. Error = 4, /// Socket is closed. Closed = 5, @@ -167,31 +65,8 @@ pub enum ChannelWriteError { Socket(super::errors::SocketAsyncError), } -impl From for SocketState { - fn from(v: u32) -> Self { - match v { - 0 => SocketState::Initial, - 1 => SocketState::Connecting, - 2 => SocketState::Connected, - 3 => SocketState::Listening, - 5 => SocketState::Closed, - _ => SocketState::Error, - } - } -} - -/// A proxy for network socket operations that decouples user I/O from network processing. -/// -/// This enum wraps different socket handle types (stream, datagram, raw) and provides -/// a unified interface for the network layer to interact with sockets. -/// The proxy enables lock-free communication between user threads and the network worker. +/// A proxy for broker-owned network socket operations and readiness. pub enum NetworkProxy { - /// Stream (TCP) socket proxy - Stream(StreamSocketChannel), - /// Datagram (UDP) socket proxy - Datagram(DatagramSocketChannel), - /// Raw socket proxy (not yet implemented) - Raw, /// Broker-owned TCP stream proxy. BrokerStream(alloc::sync::Arc>), /// Broker-owned UDP datagram proxy. @@ -200,48 +75,30 @@ pub enum NetworkProxy { impl NetworkProxy { /// Set the socket state. - /// - /// For stream sockets, this sets the connection state. - /// For datagram sockets, setting state to `Connected` marks the socket as connected. pub fn set_state(&self, state: SocketState) { match self { - NetworkProxy::Stream(channel) => channel.set_state(state), - NetworkProxy::Datagram(channel) => { - // For datagram sockets, track connected state - channel.set_connected(state == SocketState::Connected); - } - NetworkProxy::Raw => {} - NetworkProxy::BrokerStream(socket) => socket.set_state(state), - NetworkProxy::BrokerDatagram(socket) => socket.set_state(state), + Self::BrokerStream(socket) => socket.set_state(state), + Self::BrokerDatagram(socket) => socket.set_state(state), } } - /// Set the async socket error. + /// Set the asynchronous socket error. pub fn set_async_error(&self, error: super::errors::SocketAsyncError) { match self { - NetworkProxy::Stream(channel) => channel.set_async_error(error), - NetworkProxy::Datagram(channel) => channel.set_async_error(error), - NetworkProxy::Raw => {} - NetworkProxy::BrokerStream(socket) => socket.set_async_error(error), - NetworkProxy::BrokerDatagram(socket) => socket.set_async_error(error), + Self::BrokerStream(socket) => socket.set_async_error(error), + Self::BrokerDatagram(socket) => socket.set_async_error(error), } } - /// Read and optionally clear the async socket error. + /// Read and optionally clear the asynchronous socket error. pub fn get_async_error(&self, clear: bool) -> Option { match self { - NetworkProxy::Stream(channel) => channel.get_async_error(clear), - NetworkProxy::Datagram(channel) => channel.get_async_error(clear), - NetworkProxy::Raw => None, - NetworkProxy::BrokerStream(socket) => socket.get_async_error(clear), - NetworkProxy::BrokerDatagram(socket) => socket.get_async_error(clear), + Self::BrokerStream(socket) => socket.get_async_error(clear), + Self::BrokerDatagram(socket) => socket.get_async_error(clear), } } - /// Attempt to read data from the socket into the provided buffer. - /// - /// Returns the number of bytes read, or an error if the operation would block - /// or the socket is in an invalid state. + /// Attempt to read data from the socket. pub fn try_read( &self, buf: &mut [u8], @@ -249,18 +106,12 @@ impl NetworkProxy source_addr: Option<&mut Option>, ) -> Result { match self { - NetworkProxy::Stream(channel) => channel.try_read(buf, flags, source_addr), - NetworkProxy::Datagram(channel) => channel.try_read(buf, flags, source_addr), - NetworkProxy::Raw => unimplemented!(), - NetworkProxy::BrokerStream(socket) => socket.try_read(buf, flags, source_addr), - NetworkProxy::BrokerDatagram(socket) => socket.try_read(buf, flags, source_addr), + Self::BrokerStream(socket) => socket.try_read(buf, flags, source_addr), + Self::BrokerDatagram(socket) => socket.try_read(buf, flags, source_addr), } } - /// Attempt to write data to the socket from the provided buffer. - /// - /// Returns the number of bytes written, or an error if the buffer is full - /// or the socket is in an invalid state. + /// Attempt to write data to the socket. pub fn try_write( &self, buf: &[u8], @@ -276,1412 +127,24 @@ impl NetworkProxy return Err(ChannelWriteError::Unaddressable); } match self { - NetworkProxy::Stream(channel) => channel.try_write(buf), - NetworkProxy::Datagram(channel) => channel.send_to(buf, destination), - NetworkProxy::Raw => unimplemented!(), - NetworkProxy::BrokerStream(socket) => socket.try_write(buf), - NetworkProxy::BrokerDatagram(socket) => socket.try_write(buf, destination), + Self::BrokerStream(socket) => socket.try_write(buf), + Self::BrokerDatagram(socket) => socket.try_write(buf, destination), } } } + impl IOPollable for NetworkProxy { fn register_observer(&self, observer: alloc::sync::Weak>, mask: Events) { match self { - NetworkProxy::Stream(channel) => channel.register_observer(observer, mask), - NetworkProxy::Datagram(channel) => channel.register_observer(observer, mask), - NetworkProxy::Raw => {} - NetworkProxy::BrokerStream(socket) => socket.register_observer(observer, mask), - NetworkProxy::BrokerDatagram(socket) => socket.register_observer(observer, mask), + Self::BrokerStream(socket) => socket.register_observer(observer, mask), + Self::BrokerDatagram(socket) => socket.register_observer(observer, mask), } } fn check_io_events(&self) -> Events { match self { - NetworkProxy::Stream(channel) => channel.check_io_events(), - NetworkProxy::Datagram(channel) => channel.check_io_events(), - NetworkProxy::Raw => unimplemented!(), - NetworkProxy::BrokerStream(socket) => socket.check_io_events(), - NetworkProxy::BrokerDatagram(socket) => socket.check_io_events(), - } - } -} - -impl NetworkProxy { - /// Manually set the readable state. - /// - /// This is used for server sockets to indicate that a connection is ready to accept. - pub(super) fn set_readable(&self, readable: bool) { - match self { - NetworkProxy::Stream(channel) => channel.set_readable(readable), - NetworkProxy::Datagram(channel) => channel.set_readable(readable), - NetworkProxy::Raw | NetworkProxy::BrokerStream(_) | NetworkProxy::BrokerDatagram(_) => { - } - } - } - - /// Check if there is data pending in the TX buffer to be sent. - pub(super) fn has_pending_tx(&self) -> bool { - match self { - NetworkProxy::Stream(channel) => channel.has_pending_tx(), - NetworkProxy::Datagram(channel) => channel.has_pending_tx(), - NetworkProxy::Raw | NetworkProxy::BrokerStream(_) | NetworkProxy::BrokerDatagram(_) => { - false - } - } - } -} - -/// A channel for stream (TCP) socket communication. -/// -/// This channel provides lock-free data transfer between user threads and the network -/// worker. User threads write to the TX buffer and read from the RX buffer, while -/// the network worker drains TX to smoltcp and fills RX from smoltcp. -pub struct StreamSocketChannel { - inner: StreamChannelInner, -} - -/// Internal state for a stream socket channel. -struct StreamChannelInner { - /// RX producer (network worker writes here) - rx_prod: Mutex>, - /// RX consumer (user reads from here) - rx_cons: Mutex>, - /// TX producer (user writes here) - tx_prod: Mutex>, - /// TX consumer (network worker reads from here) - tx_cons: Mutex>, - - /// Current socket state - state: AtomicU32, - /// Whether the read side is shut down (SHUT_RD) - read_shutdown: AtomicBool, - /// Whether the write side is shut down (SHUT_WR) - write_shutdown: AtomicBool, - /// Bytes available in RX buffer (for quick poll checks) - rx_available: AtomicUsize, - /// The peer closed its write side. - peer_eof: AtomicBool, - /// Space available in TX buffer (for quick poll checks) - tx_available: AtomicUsize, - - /// Socket error. - socket_error: SocketAsyncErrorState, - - /// Event notification - pollee: Pollee, -} - -impl StreamChannelInner { - /// Create a new stream channel with the specified RX and TX buffer capacities. - fn new(rx_capacity: usize, tx_capacity: usize) -> Self { - let rx_rb: HeapRb = HeapRb::new(rx_capacity); - let (rx_prod, rx_cons) = rx_rb.split(); - - let tx_rb: HeapRb = HeapRb::new(tx_capacity); - let (tx_prod, tx_cons) = tx_rb.split(); - - Self { - rx_prod: Mutex::new(rx_prod), - rx_cons: Mutex::new(rx_cons), - tx_prod: Mutex::new(tx_prod), - tx_cons: Mutex::new(tx_cons), - - state: AtomicU32::new(SocketState::Initial as u32), - read_shutdown: AtomicBool::new(false), - write_shutdown: AtomicBool::new(false), - rx_available: AtomicUsize::new(0), - peer_eof: AtomicBool::new(false), - tx_available: AtomicUsize::new(tx_capacity), - - socket_error: SocketAsyncErrorState::new(), - - pollee: Pollee::new(), - } - } - - /// Get the current socket state. - fn state(&self) -> SocketState { - SocketState::from(self.state.load(Ordering::Acquire)) - } - - /// Set the socket state. - fn set_state(&self, state: SocketState) { - self.state.store(state as u32, Ordering::Release); - } -} - -impl Default for StreamSocketChannel { - fn default() -> Self { - Self::new() - } -} - -impl StreamSocketChannel { - /// Create a new stream socket channel with default buffer sizes. - /// - /// The receive side can stage one maximum-size receive operation, while - /// the transmit side uses [`super::SOCKET_BUFFER_SIZE`]. - pub fn new() -> Self { - Self::new_with_capacity( - super::SOCKET_RECEIVE_OPERATION_SIZE, - super::SOCKET_BUFFER_SIZE, - ) - } - - /// Create a new stream socket channel with specified buffer capacities. - /// - /// # Arguments - /// - /// * `rx_capacity` - Size of the receive buffer in bytes - /// * `tx_capacity` - Size of the transmit buffer in bytes - pub fn new_with_capacity(rx_capacity: usize, tx_capacity: usize) -> Self { - let inner = StreamChannelInner::new(rx_capacity, tx_capacity); - StreamSocketChannel { inner } - } - - /// Read data from the socket into the provided buffer. - /// - /// This reads from the RX ring buffer without blocking. - /// Returns the number of bytes read, or an error if the socket is closed - /// or not connected. - pub fn try_read( - &self, - buf: &mut [u8], - flags: super::ReceiveFlags, - source_addr: Option<&mut Option>, - ) -> Result { - if self.inner.read_shutdown.load(Ordering::Acquire) { - return Err(ChannelReadError::ReadShutdown); - } - if buf.is_empty() { - if let Some(source_addr) = source_addr { - *source_addr = None; - } - return Ok(0); - } - - let mut rx_cons = self.inner.rx_cons.lock(); - let (n, consumed) = if flags.contains(super::ReceiveFlags::PEEK) { - (rx_cons.peek_slice(buf), false) - } else if flags.contains(super::ReceiveFlags::DISCARD) { - (rx_cons.skip(buf.len()), true) - } else if flags.contains(super::ReceiveFlags::TRUNC) { - let n1 = rx_cons.pop_slice(buf); - let n2 = rx_cons.clear(); - (n1 + n2, true) - } else { - (rx_cons.pop_slice(buf), true) - }; - - if let Some(source_addr) = source_addr { - // TCP is connection-oriented, so no need to provide a source address - *source_addr = None; - } - - // Update available count - if consumed { - self.inner.rx_available.fetch_sub(n, Ordering::Release); - } - - if n > 0 { - return Ok(n); - } - match self.inner.state() { - SocketState::Connected if self.inner.peer_eof.load(Ordering::Acquire) => { - Err(ChannelReadError::ConnectionClosed) - } - SocketState::Connected => Err(ChannelReadError::WouldBlock), - SocketState::Closed | SocketState::Error => Err(ChannelReadError::ConnectionClosed), - _ => Err(ChannelReadError::NotConnected), - } - } - - /// Write data to the socket from the provided buffer. - /// - /// This writes to the TX ring buffer without blocking. The data will be - /// drained by the network worker and sent via smoltcp. - /// - /// Returns the number of bytes written, or an error if the socket is closed, - /// not connected, or the buffer is full. - pub fn try_write(&self, buf: &[u8]) -> Result { - if self.inner.write_shutdown.load(Ordering::Acquire) { - return Err(ChannelWriteError::WriteShutdown); - } - - match self.state() { - SocketState::Connected => {} - SocketState::Closed | SocketState::Error => { - return Err(ChannelWriteError::ConnectionClosed); - } - _ => return Err(ChannelWriteError::NotConnected), - } - - let mut tx_prod = self.inner.tx_prod.lock(); - let n = tx_prod.push_slice(buf); - - if n > 0 { - // Update available count - self.inner.tx_available.fetch_sub(n, Ordering::Release); - Ok(n) - } else { - Err(ChannelWriteError::BufferFull) - } - } - - /// Check if the socket is writable (has buffer space). - pub fn is_writable(&self) -> bool { - self.inner.tx_available.load(Ordering::Acquire) > 0 - } - - /// Shutdown the read side of the socket. - pub fn shutdown_read(&self) { - self.inner.read_shutdown.store(true, Ordering::Release); - } - - /// Shutdown the write side of the socket. - pub fn shutdown_write(&self) { - self.inner.write_shutdown.store(true, Ordering::Release); - } -} - -impl IOPollable - for StreamSocketChannel -{ - fn register_observer(&self, observer: alloc::sync::Weak>, mask: Events) { - self.inner.pollee.register_observer(observer, mask); - } - - fn check_io_events(&self) -> Events { - let mut events = Events::empty(); - - if self.is_readable() { - events |= Events::IN; - } - if self.inner.peer_eof.load(Ordering::Acquire) { - events |= Events::IN | Events::RDHUP; - } - - match self.inner.state() { - SocketState::Initial | SocketState::Closed => events |= Events::HUP | Events::OUT, - SocketState::Error => events |= Events::ERR | Events::OUT, - SocketState::Connected if self.is_writable() => events |= Events::OUT, - _ => {} - } - - events - } -} - -impl StreamSocketChannel { - /// Push received data from the network into the RX buffer using zero-copy access. - /// - /// The closure receives mutable slices directly into the ring buffer. - /// Returns the total number of bytes written. - /// - /// The closure should return how many bytes it wrote to each slice. - pub(super) fn push_rx_data_with(&self, mut f: F) -> usize - where - F: FnMut(&mut [u8]) -> usize, - { - let mut rx_prod = self.inner.rx_prod.lock(); - let (first, second) = (*rx_prod).vacant_slices_mut(); - - // SAFETY: We're treating maybe_uninit slices as &mut [u8]. - // This is safe because: - // 1. u8 has no drop implementation or invalid bit patterns - // 2. The closure will write to the slices before we advance the write index - // 3. We only advance by the number of bytes actually written - let first: &mut [u8] = unsafe { - core::slice::from_raw_parts_mut(first.as_mut_ptr().cast::(), first.len()) - }; - let second: &mut [u8] = unsafe { - core::slice::from_raw_parts_mut(second.as_mut_ptr().cast::(), second.len()) - }; - - let mut total = 0; - - // Fill first slice - if !first.is_empty() { - let written = f(first); - total += written; - } - - // Fill second slice if we have filled all of the first - if total == first.len() && !second.is_empty() { - let written = f(second); - total += written; - } - - if total > 0 { - unsafe { (*rx_prod).advance_write_index(total) }; - self.inner.rx_available.fetch_add(total, Ordering::Release); - self.inner.pollee.notify_observers(Events::IN); - } - - total - } - - /// Pop data from the TX buffer to send over the network. - /// - /// Called by the network worker when smoltcp is ready to send. - /// Returns the number of bytes popped into `buf`. - #[cfg(test)] - pub(super) fn pop_tx_data(&self, buf: &mut [u8]) -> usize { - let mut tx_cons = self.inner.tx_cons.lock(); - let n = tx_cons.pop_slice(buf); - - if n > 0 { - self.inner.tx_available.fetch_add(n, Ordering::Release); - self.inner.pollee.notify_observers(Events::OUT); - } - - n - } - - /// Pop data from the TX buffer using zero-copy access. - /// - /// The closure receives slices of data directly from the ring buffer. - /// Returns the total number of bytes consumed. - /// - /// The closure should return how many bytes it consumed from each slice. - /// This allows partial consumption (e.g., if smoltcp's send buffer is full). - pub(super) fn pop_tx_data_with(&self, mut f: F) -> usize - where - F: FnMut(&[u8]) -> usize, - { - let tx_cons = self.inner.tx_cons.lock(); - let (first, second) = tx_cons.as_slices(); - - let mut total = 0; - - // Process first slice - if !first.is_empty() { - let consumed = f(first); - total += consumed; - } - - // Process second slice if we have consumed all of the first - if total == first.len() && !second.is_empty() { - let consumed = f(second); - total += consumed; - } - - if total > 0 { - unsafe { tx_cons.advance_read_index(total) }; - self.inner.tx_available.fetch_add(total, Ordering::Release); - self.inner.pollee.notify_observers(Events::OUT); - } - - total - } - - /// Check if the socket has data available for reading. - pub(super) fn is_readable(&self) -> bool { - self.inner.rx_available.load(Ordering::Acquire) > 0 - } - - /// Manually set the readable state. - /// - /// This is used for server sockets to indicate that a connection is ready to accept. - pub(super) fn set_readable(&self, readable: bool) { - if readable { - self.inner.rx_available.store(1, Ordering::Release); - } else { - self.inner.rx_available.store(0, Ordering::Release); - } - } - - /// Check if there is data in the TX buffer waiting to be sent. - pub(super) fn has_pending_tx(&self) -> bool { - let tx_cons = self.inner.tx_cons.lock(); - !tx_cons.is_empty() - } - - /// Get the available space in the RX buffer. - /// - /// This indicates how many bytes can be pushed before the buffer is full. - #[cfg(test)] - pub(super) fn rx_space(&self) -> usize { - let rx_prod = self.inner.rx_prod.lock(); - rx_prod.vacant_len() - } - - /// Set the socket state and notify observers of state changes. - /// - /// State transitions trigger appropriate event notifications: - /// - `Connected` -> `Events::OUT` (socket is now writable) - /// - `Closed` -> `Events::HUP` (hang up) - /// - `Error` -> `Events::ERR` (error condition) - pub(super) fn set_state(&self, state: SocketState) { - let old_state = self.inner.state(); - if old_state == state { - return; - } - self.inner.set_state(state); - - // Notify user of state changes - match state { - SocketState::Connected => { - self.inner.pollee.notify_observers(Events::OUT); - } - SocketState::Closed => { - self.inner.pollee.notify_observers(Events::HUP); - } - SocketState::Error => { - self.inner.pollee.notify_observers(Events::ERR); - } - _ => {} - } - } - - /// Record that the peer closed its write side. - pub(super) fn set_peer_eof(&self) { - if !self.inner.peer_eof.swap(true, Ordering::AcqRel) { - self.inner - .pollee - .notify_observers(Events::IN | Events::RDHUP); - } - } - - /// Get the current socket state. - pub(super) fn state(&self) -> SocketState { - self.inner.state() - } - - /// Notify observers of an I/O event. - pub(super) fn notify_io_event(&self, events: Events) { - self.inner.pollee.notify_observers(events); - } - - impl_socket_async_error_accessors!(); -} - -/// A datagram message for UDP-like sockets. -/// -/// Each datagram carries its payload and an optional address: -/// - For received datagrams: the source address -/// - For sent datagrams: the destination address (or `None` if using a connected socket) -#[derive(Clone, Debug)] -pub struct DatagramMessage { - /// The data payload - pub data: Box<[u8]>, - /// Source address (for RX) or destination address (for TX) - pub addr: Option, -} - -/// A channel for datagram (UDP) socket communication. -/// -/// Unlike [`StreamSocketChannel`], this channel operates on discrete messages -/// rather than a byte stream. Each datagram is queued independently and includes -/// its associated address. -/// -/// # Capacity -/// -/// The channel has a fixed queue size for datagrams. -pub struct DatagramSocketChannel { - inner: DatagramChannelInner, -} - -/// Internal state for a datagram socket channel. -/// TODO: seperate `data` and `addr` into two ring buffers to avoid memory allocation? -struct DatagramChannelInner { - /// RX producer (network worker writes here) - rx_prod: Mutex>, - /// RX consumer (user reads from here) - rx_cons: Mutex>, - /// TX producer (user writes here) - tx_prod: Mutex>, - /// TX consumer (network worker reads from here) - tx_cons: Mutex>, - - /// Messages available in RX - rx_count: AtomicUsize, - /// Space available in TX - tx_space: AtomicUsize, - - /// Local port the socket is bound to (0 if unbound). - /// This is set atomically when auto-binding during sendto. - local_port: AtomicU16, - - /// Whether the socket is connected to a remote endpoint. - /// For UDP, this indicates that a default destination has been set via connect(). - is_connected: AtomicBool, - - /// Socket error. - socket_error: SocketAsyncErrorState, - - /// Event notification - pollee: Pollee, -} - -/// Maximum number of datagrams in queue -const DEFAULT_DATAGRAM_QUEUE_SIZE: usize = 64; - -impl DatagramChannelInner { - /// Create a new datagram channel with the specified queue size. - fn new(queue_size: usize) -> Self { - let rx_rb: HeapRb = HeapRb::new(queue_size); - let (rx_prod, rx_cons) = rx_rb.split(); - - let tx_rb: HeapRb = HeapRb::new(queue_size); - let (tx_prod, tx_cons) = tx_rb.split(); - - Self { - rx_prod: Mutex::new(rx_prod), - rx_cons: Mutex::new(rx_cons), - tx_prod: Mutex::new(tx_prod), - tx_cons: Mutex::new(tx_cons), - - rx_count: AtomicUsize::new(0), - tx_space: AtomicUsize::new(queue_size), - - local_port: AtomicU16::new(0), - is_connected: AtomicBool::new(false), - - socket_error: SocketAsyncErrorState::new(), - - pollee: Pollee::new(), - } - } -} - -impl Default - for DatagramSocketChannel -{ - fn default() -> Self { - Self::new() - } -} - -impl DatagramSocketChannel { - /// Create a new datagram socket channel with default queue size. - /// - /// The channel is created with default queue size (64 messages). - pub fn new() -> Self { - Self::new_with_capacity(DEFAULT_DATAGRAM_QUEUE_SIZE) - } - - /// Create a new datagram socket channel with specified queue size. - /// - /// # Arguments - /// - /// * `queue_size` - Maximum number of datagrams that can be queued - pub fn new_with_capacity(queue_size: usize) -> Self { - let inner = DatagramChannelInner::new(queue_size); - DatagramSocketChannel { inner } - } - - /// Receive a datagram from the socket. - /// - /// Copies the datagram payload into `buf` and optionally returns the source address. - /// If the datagram is larger than `buf`, behavior depends on `flags`. - /// Returns the original message size (which may exceed `buf.len()`). - pub fn try_read( - &self, - buf: &mut [u8], - flags: super::ReceiveFlags, - source_addr: Option<&mut Option>, - ) -> Result { - let mut rx_cons = self.inner.rx_cons.lock(); - - if flags.contains(ReceiveFlags::PEEK) { - let (first, second) = rx_cons.as_slices(); - let Some(msg) = first.first().or_else(|| second.first()) else { - return Err(ChannelReadError::WouldBlock); - }; - if let Some(source_addr) = source_addr { - *source_addr = msg.addr; - } - if !flags.contains(ReceiveFlags::DISCARD) { - let to_copy = core::cmp::min(buf.len(), msg.data.len()); - buf[..to_copy].copy_from_slice(&msg.data[..to_copy]); - } - return Ok(msg.data.len()); - } - - if let Some(msg) = rx_cons.try_pop() { - let DatagramMessage { data, addr } = msg; - if let Some(source_addr) = source_addr { - *source_addr = addr; - } - if !flags.contains(ReceiveFlags::DISCARD) { - let to_copy = core::cmp::min(buf.len(), data.len()); - buf[..to_copy].copy_from_slice(&data[..to_copy]); - } - self.inner.rx_count.fetch_sub(1, Ordering::Release); - Ok(data.len()) - } else { - Err(ChannelReadError::WouldBlock) - } - } - - /// Send a datagram to the specified address. - /// - /// The datagram is queued for transmission by the network worker. - /// Returns the number of bytes queued (always `data.len()` on success). - pub fn send_to( - &self, - data: &[u8], - addr: Option, - ) -> Result { - if addr.is_none() && !self.inner.is_connected.load(Ordering::Acquire) { - // No destination specified and socket is not connected - return Err(ChannelWriteError::DestinationAddressRequired); - } - - let size = data.len(); - let msg = DatagramMessage { - data: data.into(), - addr, - }; - let mut tx_prod = self.inner.tx_prod.lock(); - - match tx_prod.try_push(msg) { - Ok(()) => { - self.inner.tx_space.fetch_sub(1, Ordering::Release); - Ok(size) - } - Err(_) => Err(ChannelWriteError::BufferFull), - } - } - - /// Check if the socket is readable. - pub fn is_readable(&self) -> bool { - self.inner.rx_count.load(Ordering::Acquire) > 0 - } - - /// Check if the socket is writable. - pub fn is_writable(&self) -> bool { - self.inner.tx_space.load(Ordering::Acquire) > 0 - } - - /// Get the local port the socket is bound to. - /// - /// Returns 0 if the socket is not yet bound. - pub fn local_port(&self) -> u16 { - self.inner.local_port.load(Ordering::Acquire) - } - - /// Set the local port the socket is bound to. - /// - /// This should be called when the socket is bound (either explicitly or via auto-binding). - /// Uses compare-and-swap to ensure only one thread can set the port. - /// - /// Returns `Ok(())` if the port was set successfully, or `Err(current_port)` if - /// a port was already set. - pub fn set_local_port(&self, port: u16) -> Result<(), u16> { - debug_assert!(port != 0, "Port 0 is not a valid bound port"); - match self - .inner - .local_port - .compare_exchange(0, port, Ordering::AcqRel, Ordering::Acquire) - { - Ok(_) => Ok(()), - Err(current) => Err(current), + Self::BrokerStream(socket) => socket.check_io_events(), + Self::BrokerDatagram(socket) => socket.check_io_events(), } } } - -impl IOPollable - for DatagramSocketChannel -{ - fn register_observer(&self, observer: alloc::sync::Weak>, mask: Events) { - self.inner.pollee.register_observer(observer, mask); - } - - fn check_io_events(&self) -> Events { - let mut events = Events::empty(); - - if self.inner.rx_count.load(Ordering::Acquire) > 0 { - events |= Events::IN; - } - - if self.inner.tx_space.load(Ordering::Acquire) > 0 { - events |= Events::OUT; - } - - events - } -} - -impl DatagramSocketChannel { - /// Try to receive a datagram using a closure that provides the data. - /// - /// The closure should return `Some((data, source_addr))` if a datagram was received, - /// or `None` if no datagram is available. The datagram is pushed to the RX queue - /// only if the queue has space. - /// - /// Returns: - /// - `Some(len)` if a datagram was received and pushed (len is data length) - /// - `None` if the closure returned `None` or the queue is full - pub(super) fn try_recv_datagram_with(&self, f: F) -> Option - where - F: FnOnce() -> Option<(Box<[u8]>, SocketAddr)>, - { - let mut rx_prod = self.inner.rx_prod.lock(); - - if rx_prod.is_full() { - return None; - } - - let (data, source_addr) = f()?; - let len = data.len(); - - let msg = DatagramMessage { - data, - addr: Some(source_addr), - }; - - match rx_prod.try_push(msg) { - Ok(()) => { - self.inner.rx_count.fetch_add(1, Ordering::Release); - self.inner.pollee.notify_observers(Events::IN); - Some(len) - } - Err(_) => None, - } - } - - /// Try to send the next datagram using a closure, consuming it only on success. - /// - /// The closure receives the data slice and optional destination address, and should - /// return `true` if the send succeeded (datagram will be consumed) or `false` if - /// the send failed (datagram remains in queue for retry). - /// - /// Returns: - /// - `Some(true)` if a datagram was sent and consumed - /// - `Some(false)` if a datagram was peeked but send failed (still in queue) - /// - `None` if the queue is empty - pub(super) fn try_send_datagram_with(&self, f: F) -> Option - where - F: FnOnce(&[u8], Option) -> bool, - { - let mut tx_cons = self.inner.tx_cons.lock(); - - let msg = tx_cons.iter().next()?; - let success = f(&msg.data, msg.addr); - - if success { - // Send succeeded, consume the datagram - let consumed = tx_cons.try_pop(); - assert!(consumed.is_some()); - self.inner.tx_space.fetch_add(1, Ordering::Release); - self.inner.pollee.notify_observers(Events::OUT); - } - - Some(success) - } - - /// Check if the RX queue is full (cannot accept more datagrams). - #[cfg(test)] - pub(super) fn is_rx_full(&self) -> bool { - let rx_prod = self.inner.rx_prod.lock(); - rx_prod.is_full() - } - - /// Check if there are datagrams waiting to be sent. - pub(super) fn has_pending_tx(&self) -> bool { - let tx_cons = self.inner.tx_cons.lock(); - !tx_cons.is_empty() - } - - /// Manually set the readable state. - pub(super) fn set_readable(&self, readable: bool) { - if readable { - self.inner.rx_count.store(1, Ordering::Release); - } else { - self.inner.rx_count.store(0, Ordering::Release); - } - } - - /// Set the connected state of the datagram socket. - fn set_connected(&self, connected: bool) { - self.inner.is_connected.store(connected, Ordering::Release); - } - - impl_socket_async_error_accessors!(); -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::platform::mock::MockPlatform; - - type TestPlatform = MockPlatform; - - // ==================== StreamSocketChannel Tests ==================== - - #[test] - fn stream_channel_initial_state() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - - assert_eq!(channel.state(), SocketState::Initial); - - // Should not be readable initially - assert!(!channel.is_readable()); - - // Should be writable (buffer is empty) - assert!(channel.is_writable()); - - // No pending TX data - assert!(!channel.has_pending_tx()); - } - - #[test] - fn stream_channel_push_rx_and_read() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - channel.set_state(SocketState::Connected); - - // Push data from "network" side - let data = b"Hello, World!"; - let pushed = channel.push_rx_data_with(|buf: &mut [u8]| { - let to_copy = core::cmp::min(buf.len(), data.len()); - buf[..to_copy].copy_from_slice(&data[..to_copy]); - to_copy - }); - assert_eq!(pushed, data.len()); - - // Should now be readable - assert!(channel.is_readable()); - - // Read from "user" side - let mut buf = [0u8; 32]; - let read = channel - .try_read(&mut buf, super::super::ReceiveFlags::empty(), None) - .unwrap(); - assert_eq!(read, data.len()); - assert_eq!(&buf[..read], data); - - // Should no longer be readable - assert!(!channel.is_readable()); - } - - #[test] - fn stream_channel_empty_read_returns_immediately() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - channel.set_state(SocketState::Connected); - - assert_eq!( - channel - .try_read(&mut [], super::super::ReceiveFlags::empty(), None) - .unwrap(), - 0 - ); - } - - #[test] - fn stream_channel_write_and_pop_tx() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - channel.set_state(SocketState::Connected); - - // Write from "user" side - let data = b"Hello, Network!"; - let written = channel.try_write(data).unwrap(); - assert_eq!(written, data.len()); - - // Should have pending TX - assert!(channel.has_pending_tx()); - - // Pop from "network" side - let mut buf = [0u8; 32]; - let popped = channel.pop_tx_data(&mut buf); - assert_eq!(popped, data.len()); - assert_eq!(&buf[..popped], data); - - // No more pending TX - assert!(!channel.has_pending_tx()); - } - - #[test] - fn stream_channel_not_connected() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - - // Try to read while not connected - let mut buf = [0u8; 32]; - let result = channel.try_read(&mut buf, super::super::ReceiveFlags::empty(), None); - assert!(matches!(result, Err(ChannelReadError::NotConnected))); - - // Try to write while not connected - let data = b"test"; - let result = channel.try_write(data); - assert!(matches!(result, Err(ChannelWriteError::NotConnected))); - } - - #[test] - fn stream_channel_shutdown_read() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - channel.set_state(SocketState::Connected); - - // Push some data - let data = b"data"; - channel.push_rx_data_with(|buf: &mut [u8]| { - let to_copy = core::cmp::min(buf.len(), data.len()); - buf[..to_copy].copy_from_slice(&data[..to_copy]); - to_copy - }); - - // Shutdown read side - channel.shutdown_read(); - - // Should fail to read - let mut buf = [0u8; 32]; - let result = channel.try_read(&mut buf, super::super::ReceiveFlags::empty(), None); - assert!(matches!(result, Err(ChannelReadError::ReadShutdown))); - } - - #[test] - fn stream_channel_closed_after_connected_drains_rx_before_eof() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - channel.set_state(SocketState::Connected); - - let data = b"data"; - channel.push_rx_data_with(|buf: &mut [u8]| { - let to_copy = core::cmp::min(buf.len(), data.len()); - buf[..to_copy].copy_from_slice(&data[..to_copy]); - to_copy - }); - channel.set_state(SocketState::Closed); - - let mut buf = [0u8; 32]; - let read = channel - .try_read(&mut buf, super::super::ReceiveFlags::empty(), None) - .unwrap(); - assert_eq!(read, data.len()); - assert_eq!(&buf[..read], data); - - let result = channel.try_read(&mut buf, super::super::ReceiveFlags::empty(), None); - assert!(matches!(result, Err(ChannelReadError::ConnectionClosed))); - } - - #[test] - fn stream_channel_peer_eof_drains_rx_and_reports_rdhup() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - channel.set_state(SocketState::Connected); - channel.push_rx_data_with(|buf| { - buf[..4].copy_from_slice(b"data"); - 4 - }); - channel.set_peer_eof(); - - let events = channel.check_io_events(); - assert!(events.contains(Events::IN)); - assert!(events.contains(Events::RDHUP)); - - let mut buf = [0u8; 4]; - assert_eq!( - channel - .try_read(&mut buf, super::super::ReceiveFlags::empty(), None) - .unwrap(), - 4 - ); - assert_eq!(&buf, b"data"); - assert!(matches!( - channel.try_read(&mut buf, super::super::ReceiveFlags::empty(), None), - Err(ChannelReadError::ConnectionClosed) - )); - } - - #[test] - fn stream_channel_shutdown_write() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - channel.set_state(SocketState::Connected); - - // Shutdown write side - channel.shutdown_write(); - - // Should fail to write - let result = channel.try_write(b"data"); - assert!(matches!(result, Err(ChannelWriteError::WriteShutdown))); - } - - #[test] - fn stream_channel_rx_space() { - let capacity = 1024; - let channel: StreamSocketChannel = - StreamSocketChannel::new_with_capacity(capacity, capacity); - channel.set_state(SocketState::Connected); - - // Initially all space is available - assert_eq!(channel.rx_space(), capacity); - - // Push some data - let pushed = channel.push_rx_data_with(|buf: &mut [u8]| { - let to_write = core::cmp::min(buf.len(), 100); - buf[..to_write].fill(0); - to_write - }); - assert_eq!(pushed, 100); - - // Space should decrease - assert_eq!(channel.rx_space(), capacity - 100); - } - - #[test] - fn stream_channel_partial_read() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - channel.set_state(SocketState::Connected); - - // Push 100 bytes - let pushed = channel.push_rx_data_with(|buf: &mut [u8]| { - let to_write = core::cmp::min(buf.len(), 100); - buf[..to_write].fill(42); - to_write - }); - assert_eq!(pushed, 100); - - // Read only 50 bytes - let mut buf = [0u8; 50]; - let read = channel - .try_read(&mut buf, super::super::ReceiveFlags::empty(), None) - .unwrap(); - assert_eq!(read, 50); - assert!(buf.iter().all(|&b| b == 42)); - - // Should still be readable (50 bytes remaining) - assert!(channel.is_readable()); - - // Read remaining - let read = channel - .try_read(&mut buf, super::super::ReceiveFlags::empty(), None) - .unwrap(); - assert_eq!(read, 50); - } - - #[test] - fn stream_channel_io_events() { - let channel: StreamSocketChannel = StreamSocketChannel::new(); - - let events = channel.check_io_events(); - assert!(events.contains(Events::HUP)); - - // Connected with empty RX and available TX - channel.set_state(SocketState::Connected); - let events = channel.check_io_events(); - assert!(!events.contains(Events::IN)); // No data to read - assert!(events.contains(Events::OUT)); // Can write - - // Push data to RX - let data = b"data"; - channel.push_rx_data_with(|buf: &mut [u8]| { - let to_copy = core::cmp::min(buf.len(), data.len()); - buf[..to_copy].copy_from_slice(&data[..to_copy]); - to_copy - }); - let events = channel.check_io_events(); - assert!(events.contains(Events::IN)); // Data available - assert!(events.contains(Events::OUT)); // Can still write - } - - // ==================== DatagramSocketChannel Tests ====================================== - - const DUMMY_ADDR: core::net::SocketAddr = core::net::SocketAddr::V4( - core::net::SocketAddrV4::new(core::net::Ipv4Addr::LOCALHOST, 1234), - ); - - #[test] - fn datagram_channel_initial_state() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - // Should not be readable initially - assert!(!channel.is_readable()); - - // Should be writable (queue is empty) - assert!(channel.is_writable()); - - // No pending TX - assert!(!channel.has_pending_tx()); - - // RX not full - assert!(!channel.is_rx_full()); - } - - #[test] - fn datagram_channel_send_and_receive() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - // Send a datagram (user side) - let addr = Some(core::net::SocketAddr::V4(core::net::SocketAddrV4::new( - core::net::Ipv4Addr::new(10, 0, 0, 1), - 8080, - ))); - let result = channel.send_to(b"Hello, UDP!", addr); - assert!(result.is_ok()); - assert_eq!(result.unwrap(), 11); - - // Should have pending TX - assert!(channel.has_pending_tx()); - - // Pop from network side using try_send_datagram_with - let mut received_data = None; - let mut received_addr = None; - let result = channel.try_send_datagram_with(|data, dest| { - received_data = Some(data.to_vec()); - received_addr = Some(dest); - true // consume the datagram - }); - assert_eq!(result, Some(true)); - assert_eq!(received_data.unwrap(), b"Hello, UDP!"); - assert_eq!(received_addr.unwrap(), addr); - - // No more pending TX - assert!(!channel.has_pending_tx()); - } - - #[test] - fn datagram_channel_push_and_read() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - // Push a datagram (network side) using try_recv_datagram_with - let addr = core::net::SocketAddr::V4(core::net::SocketAddrV4::new( - core::net::Ipv4Addr::new(192, 168, 1, 1), - 1234, - )); - let result = - channel.try_recv_datagram_with(|| Some((Box::from(*b"Incoming packet"), addr))); - assert_eq!(result, Some(15)); - - // Should be readable - assert!(channel.is_readable()); - - // Read from user side - let mut buf = [0u8; 64]; - let mut source = None; - let read = channel - .try_read( - &mut buf, - super::super::ReceiveFlags::empty(), - Some(&mut source), - ) - .unwrap(); - assert_eq!(read, 15); - assert_eq!(&buf[..read], b"Incoming packet"); - assert_eq!(source, Some(addr)); - } - - #[test] - fn datagram_channel_peek_preserves_the_message() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - channel - .try_recv_datagram_with(|| Some((Box::from(*b"peek"), DUMMY_ADDR))) - .unwrap(); - - let mut buf = [0u8; 4]; - let mut source = None; - assert_eq!( - channel - .try_read( - &mut buf, - super::super::ReceiveFlags::PEEK, - Some(&mut source), - ) - .unwrap(), - 4 - ); - assert_eq!(&buf, b"peek"); - assert_eq!(source, Some(DUMMY_ADDR)); - assert!(channel.is_readable()); - assert_eq!( - channel - .try_read(&mut buf, super::super::ReceiveFlags::empty(), None) - .unwrap(), - 4 - ); - assert!(!channel.is_readable()); - } - - #[test] - fn datagram_channel_peek_preserves_an_empty_message() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - channel - .try_recv_datagram_with(|| { - Some((alloc::vec::Vec::new().into_boxed_slice(), DUMMY_ADDR)) - }) - .unwrap(); - - let mut buf = []; - assert_eq!( - channel - .try_read(&mut buf, super::super::ReceiveFlags::PEEK, None,) - .unwrap(), - 0 - ); - assert!(channel.is_readable()); - assert_eq!( - channel - .try_read(&mut buf, super::super::ReceiveFlags::empty(), None) - .unwrap(), - 0 - ); - assert!(!channel.is_readable()); - } - - #[test] - fn datagram_channel_read_empty() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - // Try to read when empty - let mut buf = [0u8; 64]; - let result = channel.try_read(&mut buf, super::super::ReceiveFlags::empty(), None); - assert!(matches!(result, Err(ChannelReadError::WouldBlock))); - } - - #[test] - fn datagram_channel_queue_full() { - let queue_size = 4; - let channel: DatagramSocketChannel = - DatagramSocketChannel::new_with_capacity(queue_size); - - // Fill the TX queue - for i in 0..queue_size { - let result = channel.send_to(&alloc::vec![0; i], Some(DUMMY_ADDR)); - assert!(result.is_ok()); - } - - // Next send should fail - let result = channel.send_to(&[99], Some(DUMMY_ADDR)); - assert!(matches!(result, Err(ChannelWriteError::BufferFull))); - } - - #[test] - fn datagram_channel_unconnected_send_without_address() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - // Sending without an address on an unconnected socket should fail - let result = channel.send_to(&[1, 2, 3], None); - assert!(matches!( - result, - Err(ChannelWriteError::DestinationAddressRequired) - )); - } - - #[test] - fn datagram_channel_rx_full() { - let queue_size = 4; - let channel: DatagramSocketChannel = - DatagramSocketChannel::new_with_capacity(queue_size); - - // Fill the RX queue - for i in 0..queue_size { - let data: Box<[u8]> = alloc::vec![0; i].into_boxed_slice(); - let result = channel.try_recv_datagram_with(|| Some((data, DUMMY_ADDR))); - assert!(result.is_some()); - } - - // Queue should be full - assert!(channel.is_rx_full()); - - // Next push should fail (returns None when full) - let result = channel.try_recv_datagram_with(|| Some((Box::from([99u8]), DUMMY_ADDR))); - assert!(result.is_none()); - } - - #[test] - fn datagram_channel_truncation() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - // Push a large datagram - let data: Box<[u8]> = alloc::vec![42u8; 100].into_boxed_slice(); - channel - .try_recv_datagram_with(|| Some((data.clone(), DUMMY_ADDR))) - .unwrap(); - - // Read with a small buffer (no TRUNC flag) - let mut buf = [0u8; 10]; - let read = channel - .try_read(&mut buf, super::super::ReceiveFlags::empty(), None) - .unwrap(); - assert_eq!(read, data.len()); - } - - #[test] - fn datagram_channel_trunc_flag() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - let dummy_addr = core::net::SocketAddr::V4(core::net::SocketAddrV4::new( - core::net::Ipv4Addr::LOCALHOST, - 1234, - )); - - // Push a large datagram - let data: Box<[u8]> = alloc::vec![42u8; 100].into_boxed_slice(); - channel - .try_recv_datagram_with(|| Some((data.clone(), dummy_addr))) - .unwrap(); - - // Read with TRUNC flag - should return actual packet size - let mut buf = [0u8; 10]; - let read = channel - .try_read(&mut buf, super::super::ReceiveFlags::TRUNC, None) - .unwrap(); - assert_eq!(read, 100); // Returns actual datagram size - } - - #[test] - fn datagram_channel_io_events() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - let dummy_addr = core::net::SocketAddr::V4(core::net::SocketAddrV4::new( - core::net::Ipv4Addr::LOCALHOST, - 1234, - )); - - // Initially: no IN, has OUT - let events = channel.check_io_events(); - assert!(!events.contains(Events::IN)); - assert!(events.contains(Events::OUT)); - - // Push a datagram using try_recv_datagram_with - channel - .try_recv_datagram_with(|| Some((Box::from(*b"test"), dummy_addr))) - .unwrap(); - - // Now has IN - let events = channel.check_io_events(); - assert!(events.contains(Events::IN)); - assert!(events.contains(Events::OUT)); - } - - #[test] - fn datagram_channel_try_send_failure() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - // Send a datagram (user side) - let addr = Some(core::net::SocketAddr::V4(core::net::SocketAddrV4::new( - core::net::Ipv4Addr::new(10, 0, 0, 1), - 8080, - ))); - channel.send_to(b"Hello!", addr).unwrap(); - - // Try to send but fail (return false) - let result = channel.try_send_datagram_with(|_data, _dest| { - false // simulate send failure - }); - assert_eq!(result, Some(false)); - - // Datagram should still be in queue - assert!(channel.has_pending_tx()); - - // Try again and succeed - let result = channel.try_send_datagram_with(|data, dest| { - assert_eq!(data, b"Hello!"); - assert_eq!(dest, addr); - true - }); - assert_eq!(result, Some(true)); - - // No more pending TX - assert!(!channel.has_pending_tx()); - } - - #[test] - fn datagram_channel_try_recv_none() { - let channel: DatagramSocketChannel = DatagramSocketChannel::new(); - - // Closure returns None - nothing should be pushed - let result = channel.try_recv_datagram_with(|| None); - assert!(result.is_none()); - - // Should still not be readable - assert!(!channel.is_readable()); - } -} diff --git a/litebox/src/net/tests.rs b/litebox/src/net/tests.rs deleted file mode 100644 index e82dd127aa..0000000000 --- a/litebox/src/net/tests.rs +++ /dev/null @@ -1,230 +0,0 @@ -// Copyright (c) Microsoft Corporation. -// Licensed under the MIT license. - -use platform::mock::MockPlatform; - -use super::*; - -use core::net::SocketAddrV4; -use core::str::FromStr; - -extern crate std; - -fn bidi_tcp_comms(mut network: Network, comms: fn(&mut Network)) { - // Create a listening socket - let listener_fd = network - .local_socket(Protocol::Tcp) - .expect("Failed to create TCP socket"); - let listen_addr = SocketAddr::V4(SocketAddrV4::from_str("10.0.0.2:8080").unwrap()); - - network - .bind(&listener_fd, &listen_addr) - .expect("Failed to bind TCP socket"); - network - .listen(&listener_fd, 1) - .expect("Failed to listen on TCP socket"); - - // Create a connecting socket - let client_fd = network - .local_socket(Protocol::Tcp) - .expect("Failed to create TCP socket"); - let err = network - .connect(&client_fd, &listen_addr, false) - .unwrap_err(); - assert!( - matches!(err, ConnectError::InProgress), - "Expected InProgress error, got {err:?}", - ); - - comms(&mut network); - - // Accept the connection on the listening socket - let server_fd = loop { - match network.accept(&listener_fd, None) { - Ok(fd) => break fd, - Err(AcceptError::NoConnectionsReady) => {} - Err(other) => panic!("Unexpected accept error: {other:?}"), - } - }; - - // Send data from client to server - let client_to_server_data = b"Hello from client!"; - let bytes_sent = network - .send(&client_fd, client_to_server_data, SendFlags::empty(), None) - .expect("Failed to send data"); - assert_eq!(bytes_sent, client_to_server_data.len()); - - comms(&mut network); - - // Receive data on the server - let mut server_buffer = [0u8; 1024]; - let bytes_received = network - .receive(&server_fd, &mut server_buffer, ReceiveFlags::empty(), None) - .expect("Failed to receive data"); - assert_eq!(&server_buffer[..bytes_received], client_to_server_data); - - // Send data from server to client - let server_to_client_data = b"Hello from server!"; - let bytes_sent = network - .send(&server_fd, server_to_client_data, SendFlags::empty(), None) - .expect("Failed to send data"); - assert_eq!(bytes_sent, server_to_client_data.len()); - - comms(&mut network); - - // Receive data on the client - let mut client_buffer = [0u8; 1024]; - let bytes_received = network - .receive(&client_fd, &mut client_buffer, ReceiveFlags::empty(), None) - .expect("Failed to receive data"); - assert_eq!(&client_buffer[..bytes_received], server_to_client_data); - - network.close(&client_fd, CloseBehavior::Immediate).unwrap(); - network.close(&server_fd, CloseBehavior::Immediate).unwrap(); - network - .close(&listener_fd, CloseBehavior::Immediate) - .unwrap(); -} - -#[test] -fn test_bidirectional_tcp_communication_default() { - let litebox = LiteBox::new(MockPlatform::new()); - let network = Network::new(&litebox); - bidi_tcp_comms(network, |_| {}); -} - -#[test] -fn test_bidirectional_tcp_communication_manual() { - let litebox = LiteBox::new(MockPlatform::new()); - let mut network = Network::new(&litebox); - network.set_platform_interaction(PlatformInteraction::Manual); - bidi_tcp_comms(network, |nw| { - while nw.perform_platform_interaction().call_again_immediately() {} - }); -} - -#[test] -fn test_bidirectional_tcp_communication_automatic() { - let litebox = LiteBox::new(MockPlatform::new()); - let mut network = Network::new(&litebox); - network.set_platform_interaction(PlatformInteraction::Automatic); - bidi_tcp_comms(network, |_| {}); -} - -#[test] -fn attach_socket_proxy_is_idempotent() { - let litebox = LiteBox::new(MockPlatform::new()); - let mut network = Network::new(&litebox); - let fd = network - .local_socket(Protocol::Tcp) - .expect("failed to create TCP socket"); - - let first = network - .attach_socket_proxy(&fd) - .expect("TCP socket must support a proxy"); - let second = network - .attach_socket_proxy(&fd) - .expect("TCP socket must retain its proxy"); - - assert!(alloc::sync::Arc::ptr_eq(&first, &second)); -} - -#[test] -fn pinned_immediate_close_remains_abortive() { - let litebox = LiteBox::new(MockPlatform::new()); - let mut network = Network::new(&litebox); - let fd = network - .local_socket(Protocol::Tcp) - .expect("failed to create TCP socket"); - let entry = litebox - .descriptor_table() - .entry_handle(&fd) - .expect("socket entry must exist"); - - network - .close(&fd, CloseBehavior::Immediate) - .expect("close must be queued while the entry is pinned"); - assert!( - entry.with_entry(|socket| socket.entry.tcp().immediate_close.load(Ordering::SeqCst)), - "queued close must retain abortive semantics" - ); - - drop(entry); - assert!(!network.finish_deferred_closes()); - assert!(network.queued_for_closure.is_empty()); -} - -#[test] -fn final_duplicate_close_updates_abortive_behavior() { - let litebox = LiteBox::new(MockPlatform::new()); - let mut network = Network::new(&litebox); - let fd = network - .local_socket(Protocol::Tcp) - .expect("failed to create TCP socket"); - let duplicate = litebox - .descriptor_table_mut() - .duplicate(&fd) - .expect("socket duplication must succeed"); - let entry = litebox - .descriptor_table() - .entry_handle(&duplicate) - .expect("socket entry must exist"); - - network - .close(&fd, CloseBehavior::Immediate) - .expect("first duplicate close must be queued"); - assert!(entry.with_entry(|socket| socket.entry.tcp().immediate_close.load(Ordering::SeqCst))); - - network - .close(&duplicate, CloseBehavior::Graceful) - .expect("final duplicate close must be queued while the entry is pinned"); - assert!( - !entry.with_entry(|socket| socket.entry.tcp().immediate_close.load(Ordering::SeqCst)), - "the final descriptor close must determine abortive behavior" - ); - - drop(entry); - assert!(!network.finish_deferred_closes()); - assert!(network.queued_for_closure.is_empty()); -} - -#[test] -fn final_duplicate_close_is_reaped_immediately() { - let litebox = LiteBox::new(MockPlatform::new()); - let mut network = Network::new(&litebox); - let fd = network - .local_socket(Protocol::Tcp) - .expect("failed to create TCP socket"); - let duplicate = litebox - .descriptor_table_mut() - .duplicate(&fd) - .expect("socket duplication must succeed"); - - network - .close(&fd, CloseBehavior::Graceful) - .expect("first duplicate close must be queued"); - assert_eq!(network.queued_for_closure.len(), 1); - network - .close(&duplicate, CloseBehavior::Graceful) - .expect("final duplicate close must succeed"); - assert!(network.queued_for_closure.is_empty()); -} - -#[test] -fn unsupported_sockets_fail_cleanly_without_broker() { - let litebox = LiteBox::new(MockPlatform::new()); - let mut network = Network::new(&litebox); - - assert!(matches!( - network.socket(Protocol::Tcp), - Err(SocketError::BrokerUnavailable) - )); - assert!(matches!( - network.socket(Protocol::Udp), - Err(SocketError::BrokerUnavailable) - )); - assert!(matches!( - network.socket(Protocol::Icmp), - Err(SocketError::UnsupportedProtocol(1)) - )); -} diff --git a/litebox/src/platform/mock.rs b/litebox/src/platform/mock.rs index 5eed9f861e..733b4ce805 100644 --- a/litebox/src/platform/mock.rs +++ b/litebox/src/platform/mock.rs @@ -26,14 +26,12 @@ use super::*; /// /// - Full determinism /// + time moves at one millisecond per "now" call -/// + IP packets are placed into a deterministic ring buffer and spin back around /// - Debuging output goes to stderr /// - Can pre-fill stdin and check stdout easily between invocations (see [`Self::stdin_queue`], /// [`Self::stdout_queue`], and [`Self::stderr_queue`]) /// - It will not mock you for using it during tests pub(crate) struct MockPlatform { current_time: AtomicU64, - ip_packets: RwLock>>, random: Mutex, pub(crate) stdin_queue: RwLock>>, pub(crate) stdout_queue: RwLock>>, @@ -46,7 +44,6 @@ impl MockPlatform { // order to give ourselves a statically lived platform easily. alloc::boxed::Box::leak(alloc::boxed::Box::new(MockPlatform { current_time: AtomicU64::new(0), - ip_packets: RwLock::new(VecDeque::new()), random: Mutex::new(crate::utils::rng::FastRng::new_from_seed( core::num::NonZeroU64::new(0x4d595df4d0f33173).unwrap(), )), @@ -189,25 +186,6 @@ impl RawMutexProvider for MockPlatform { type RawMutex = MockRawMutex; } -impl IPInterfaceProvider for MockPlatform { - fn send_ip_packet(&self, packet: &[u8]) -> Result<(), SendError> { - self.ip_packets.write().unwrap().push_back(packet.into()); - Ok(()) - } - - fn receive_ip_packet(&self, packet: &mut [u8]) -> Result { - if self.ip_packets.read().unwrap().is_empty() { - Err(ReceiveError::WouldBlock) - } else { - let mut ipp = self.ip_packets.write().unwrap(); - let v = ipp.pop_front().unwrap(); - assert!(v.len() <= packet.len()); - packet[..v.len()].copy_from_slice(&v); - Ok(v.len()) - } - } -} - #[derive(Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] pub(crate) struct MockInstant { time: u64, diff --git a/litebox/src/platform/mod.rs b/litebox/src/platform/mod.rs index 7ba721f752..db54d28d7e 100644 --- a/litebox/src/platform/mod.rs +++ b/litebox/src/platform/mod.rs @@ -27,7 +27,7 @@ pub use page_mgmt::PageManagementProvider; /// provided by it. _However_, most of the provided APIs within the provider act upon an `&self` to /// allow storage of any useful "globals" within it necessary. pub trait Provider: - RawMutexProvider + IPInterfaceProvider + TimeProvider + ArchSpecificProvider + RawPointerProvider + RawMutexProvider + TimeProvider + ArchSpecificProvider + RawPointerProvider { } @@ -232,36 +232,6 @@ pub enum UnblockedOrTimedOut { TimedOut, } -/// An IP packet interface to the outside world. -/// -/// This could be implemented via a `read`/`write` to a TUN device. -pub trait IPInterfaceProvider { - /// Send the IP packet. - /// - /// Returns `Ok(())` when entire packet is sent, or a [`SendError`] if it is unable to send the - /// entire packet. - fn send_ip_packet(&self, packet: &[u8]) -> Result<(), SendError>; - - /// Receive an IP packet into `packet`. - /// - /// Returns size of packet received, or a [`ReceiveError`] if unable to receive an entire - /// packet. - fn receive_ip_packet(&self, packet: &mut [u8]) -> Result; -} - -/// A non-exhaustive list of errors that can be thrown by [`IPInterfaceProvider::send_ip_packet`]. -#[derive(Error, Debug)] -#[non_exhaustive] -pub enum SendError {} - -/// A non-exhaustive list of errors that can be thrown by [`IPInterfaceProvider::receive_ip_packet`]. -#[derive(Error, Debug)] -#[non_exhaustive] -pub enum ReceiveError { - #[error("Receive operation would block")] - WouldBlock, -} - /// An interface to understanding time. pub trait TimeProvider { type Instant: Instant; diff --git a/litebox_broker_core/src/policy.rs b/litebox_broker_core/src/policy.rs index f12a9b2b10..7cc5e6ad74 100644 --- a/litebox_broker_core/src/policy.rs +++ b/litebox_broker_core/src/policy.rs @@ -3,6 +3,7 @@ use core::net::SocketAddrV4; +use crate::socket::PlatformSocketScope; use crate::{BrokerError, CallerCredential, ObjectRights}; use litebox_broker_protocol::socket::{ AddressFamily, CreateSocketRequest, IpProtocol, Ipv4Address, Port, SocketType, @@ -324,6 +325,36 @@ impl SocketPolicy { .any(|rule| rule.permits(caller_credential, address)), } } + + /// Derives the host-network binding scope required by this socket's authorized destinations. + fn socket_scope( + self, + caller_credential: CallerCredential, + request: CreateSocketRequest, + ) -> PlatformSocketScope { + let (( + SocketType::Stream, + IpProtocol::Tcp, + Self::TcpDestinationRules(policy) | Self::TcpUdpDestinationRules { tcp: policy, .. }, + ) + | ( + SocketType::Datagram, + IpProtocol::Udp, + Self::UdpDestinationRules(policy) | Self::TcpUdpDestinationRules { udp: policy, .. }, + )) = (request.socket_type, request.protocol, self) + else { + return PlatformSocketScope::LoopbackOnly; + }; + if policy.rules().iter().any(|rule| { + rule.caller_credential == caller_credential + && !(rule.destination.prefix_length() >= 8 + && rule.destination.network().0[0] == 127) + }) { + PlatformSocketScope::GeneralIpv4 + } else { + PlatformSocketScope::LoopbackOnly + } + } } fn copy_destination_rules( @@ -481,6 +512,14 @@ impl PolicyEngine { Err(BrokerError::PolicyDenied) } } + + pub(crate) fn socket_scope( + &self, + caller_credential: CallerCredential, + request: CreateSocketRequest, + ) -> PlatformSocketScope { + self.socket_policy.socket_scope(caller_credential, request) + } } impl Default for PolicyEngine { diff --git a/litebox_broker_core/src/session.rs b/litebox_broker_core/src/session.rs index aa4a65f031..fcc31a9ae8 100644 --- a/litebox_broker_core/src/session.rs +++ b/litebox_broker_core/src/session.rs @@ -6,7 +6,7 @@ use core::sync::atomic::{AtomicUsize, Ordering}; use crate::event::EventObject; use crate::pipe::PipeObject; -use crate::socket::SocketObject; +use crate::socket::{SessionSocketPorts, SocketObject}; use crate::{BrokerCore, BrokerError, Result}; use hashbrown::HashMap; use litebox_broker_protocol::ObjectHandle; @@ -77,6 +77,8 @@ pub struct BrokerSession { references: Mutex, /// Socket quota held by pending, live, and closing in-flight resources. pub(crate) reserved_sockets: Arc, + /// Guest-visible TCP and UDP port namespaces owned by this session. + pub(crate) socket_ports: SessionSocketPorts, } impl BrokerSession { @@ -95,6 +97,7 @@ impl BrokerSession { pending_handles: 0, }), reserved_sockets: Arc::new(AtomicUsize::new(0)), + socket_ports: SessionSocketPorts::default(), } } diff --git a/litebox_broker_core/src/socket.rs b/litebox_broker_core/src/socket.rs index f823f0508e..b0050eea6c 100644 --- a/litebox_broker_core/src/socket.rs +++ b/litebox_broker_core/src/socket.rs @@ -3,10 +3,11 @@ //! Broker-owned platform socket authority. -use alloc::sync::Arc; +use alloc::{sync::Arc, vec::Vec}; use core::net::{Ipv4Addr, SocketAddrV4}; use core::sync::atomic::{AtomicUsize, Ordering}; +use hashbrown::HashMap; use litebox_broker_protocol::ObjectHandle; use litebox_broker_protocol::readiness::ReadinessFlags; use litebox_broker_protocol::socket::{ @@ -15,6 +16,7 @@ use litebox_broker_protocol::socket::{ ShutdownMode, SocketConnectionStatus, SocketError, SocketOutcome, SocketStatusResponse, SocketType, TcpOptionName, TcpOptionValue, }; +use spin::Mutex; use spin::Once; use crate::readiness::{ReadinessRegistration, ReadinessSink}; @@ -22,13 +24,145 @@ use crate::session::{ObjectEntry, ObjectRights}; use crate::{BrokerError, BrokerSession, Result, SessionId}; const DEFAULT_TCP_LISTEN_ADDRESS: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0); +const DEFAULT_TCP_LOCAL_ADDRESS: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0); +const DEFAULT_UDP_LOCAL_ADDRESS: SocketAddrV4 = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 0); +const FIRST_EPHEMERAL_PORT: u16 = 49152; + +#[derive(Clone, Copy)] +enum GuestPortProtocol { + Tcp, + Udp, +} + +#[derive(Default)] +struct SessionSocketPortState { + tcp: HashMap, + udp: HashMap, + next_tcp_ephemeral: Option, + next_udp_ephemeral: Option, +} + +/// Per-session authority for guest-visible TCP and UDP port namespaces. +#[derive(Clone, Default)] +pub(crate) struct SessionSocketPorts { + state: Arc>, +} + +impl SessionSocketPorts { + fn reserve( + &self, + request: CreateSocketRequest, + requested_address: SocketAddrV4, + mut port_is_reserved: impl FnMut(u16) -> bool, + ) -> Result> { + let protocol = guest_port_protocol(request).ok_or(BrokerError::Internal)?; + let mut state = self.state.lock(); + let port = if requested_address.port() == 0 { + state.allocate_ephemeral(protocol, &mut port_is_reserved)? + } else if state + .bindings(protocol) + .contains_key(&requested_address.port()) + { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } else { + requested_address.port() + }; + let local_address = SocketAddrV4::new(*requested_address.ip(), port); + state + .bindings_mut(protocol) + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory)?; + if state + .bindings_mut(protocol) + .insert(port, local_address) + .is_some() + { + return Err(BrokerError::Internal); + } + drop(state); + Ok(SocketOutcome::Completed(( + local_address, + GuestPortReservation { + ports: self.clone(), + protocol, + port, + }, + ))) + } +} + +impl SessionSocketPortState { + fn bindings(&self, protocol: GuestPortProtocol) -> &HashMap { + match protocol { + GuestPortProtocol::Tcp => &self.tcp, + GuestPortProtocol::Udp => &self.udp, + } + } + + fn bindings_mut(&mut self, protocol: GuestPortProtocol) -> &mut HashMap { + match protocol { + GuestPortProtocol::Tcp => &mut self.tcp, + GuestPortProtocol::Udp => &mut self.udp, + } + } + + fn next_ephemeral_mut(&mut self, protocol: GuestPortProtocol) -> &mut Option { + match protocol { + GuestPortProtocol::Tcp => &mut self.next_tcp_ephemeral, + GuestPortProtocol::Udp => &mut self.next_udp_ephemeral, + } + } + + fn allocate_ephemeral( + &mut self, + protocol: GuestPortProtocol, + port_is_reserved: &mut impl FnMut(u16) -> bool, + ) -> Result { + let start = self + .next_ephemeral_mut(protocol) + .unwrap_or(FIRST_EPHEMERAL_PORT); + let mut port = start; + loop { + if !self.bindings(protocol).contains_key(&port) && !port_is_reserved(port) { + *self.next_ephemeral_mut(protocol) = Some(if port == u16::MAX { + FIRST_EPHEMERAL_PORT + } else { + port + 1 + }); + return Ok(port); + } + port = if port == u16::MAX { + FIRST_EPHEMERAL_PORT + } else { + port + 1 + }; + if port == start { + return Err(BrokerError::ResourceExhausted); + } + } + } +} + +struct GuestPortReservation { + ports: SessionSocketPorts, + protocol: GuestPortProtocol, + port: u16, +} + +impl Drop for GuestPortReservation { + fn drop(&mut self) { + self.ports + .state + .lock() + .bindings_mut(self.protocol) + .remove(&self.port); + } +} /// Platform socket and endpoint metadata returned by an accept operation. pub struct AcceptedPlatformSocket { /// Accepted nonblocking platform socket. pub socket: Arc, - /// Local endpoint of the accepted connection. - pub local_address: SocketAddrV4, /// Remote endpoint of the accepted connection. pub remote_address: SocketAddrV4, } @@ -42,6 +176,24 @@ pub enum PlatformConnectError { PeerIndeterminate(BrokerError), } +/// Whether a guest-local binding was requested explicitly or allocated implicitly. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum PlatformBindKind { + /// The guest explicitly requested this nonzero local port through `bind`. + Explicit, + /// The broker allocated the local port for another socket operation. + Implicit, +} + +/// Host-network scope required by an authorized broker socket. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum PlatformSocketScope { + /// The socket can communicate only within IPv4 loopback. + LoopbackOnly, + /// The socket can communicate with authorized non-loopback IPv4 destinations. + GeneralIpv4, +} + /// Broker socket and endpoint metadata returned by an accept operation. pub struct AcceptedBrokerSocket { /// Broker handle naming the accepted socket. @@ -63,12 +215,41 @@ pub struct ReceivedPlatformDatagram { pub source_address: SocketAddrV4, } +/// Owned platform result for one stream receive. +#[derive(Debug, PartialEq, Eq)] +pub enum PlatformStreamReceive { + /// Bytes received from the stream. + Received(Vec), + /// The stream's receive direction reached end of stream. + EndOfStream, +} + +/// Owned platform result for one datagram receive. +#[derive(Debug, PartialEq, Eq)] +pub struct PlatformDatagramReceive { + /// Received datagram prefix, truncated to the requested capacity. + pub data: Vec, + /// Original datagram length before truncation. + pub datagram_length: usize, + /// Source address of the datagram. + pub source_address: SocketAddrV4, +} + /// Broker-wide socket provider supplied by the host platform. /// /// The provider creates per-socket [`PlatformSocket`] resources and owns any /// bookkeeping shared across the sockets of a broker session. Operations on an /// individual socket belong to [`PlatformSocket`], not this shared provider. pub trait SocketProvider: Send + Sync { + /// Returns whether a guest port is reserved for an explicit host port mapping. + /// + /// Reserved ports are skipped by implicit guest-port allocation but remain + /// available to an explicit nonzero `bind` request. + fn reserves_guest_port(&self, request: CreateSocketRequest, port: u16) -> bool { + let _ = (request, port); + false + } + /// Creates one nonblocking socket resource for an authenticated session. /// /// The returned socket must not retain authority beyond its `Arc` lifetime. @@ -76,6 +257,7 @@ pub trait SocketProvider: Send + Sync { &self, session_id: SessionId, request: CreateSocketRequest, + scope: PlatformSocketScope, readiness: ReadinessRegistration, ) -> Result>; @@ -89,11 +271,13 @@ pub trait SocketProvider: Send + Sync { /// in flight to finish after its object handle closes. Dropping the final `Arc` /// releases the platform socket. pub trait PlatformSocket: Send + Sync { - /// Binds this socket to a local address. - fn bind(&self, address: SocketAddrV4) -> Result>; + /// Associates this socket with a broker-reserved guest-local address. + /// + /// The address has a nonzero port and is not a host socket endpoint. + fn bind(&self, address: SocketAddrV4, kind: PlatformBindKind) -> Result>; /// Makes this socket listen for incoming connections. - fn listen(&self, backlog: u32) -> Result>; + fn listen(&self, backlog: u32) -> Result>; /// Accepts one pending connection without waiting. fn accept( @@ -118,7 +302,7 @@ pub trait PlatformSocket: Send + Sync { /// /// A temporarily full socket returns [`BrokerError::WouldBlock`]. Ordinary /// network failures return [`SocketOutcome::Failed`]. - fn send(&self, data: &[u8], flags: SendFlags) -> Result>; + fn send(&self, data: Vec, flags: SendFlags) -> Result>; /// Sends one complete datagram without waiting for platform readiness. /// @@ -126,7 +310,7 @@ pub trait PlatformSocket: Send + Sync { /// the socket's connected peer. Partial successful sends are invalid. fn send_to( &self, - data: &[u8], + data: Vec, flags: SendFlags, destination: Option, ) -> Result>; @@ -134,15 +318,15 @@ pub trait PlatformSocket: Send + Sync { /// Receives bytes without waiting for platform readiness. /// /// A temporarily empty socket returns [`BrokerError::WouldBlock`]. End of - /// stream returns [`ReceiveSocketResponse::EndOfStream`]; `Received(0)` is - /// reserved for a zero-length input buffer handled by the core. + /// stream returns [`PlatformStreamReceive::EndOfStream`]. The core handles + /// zero-length receives without invoking the platform. fn receive( &self, - data: &mut [u8], + length: usize, flags: ReceiveFlags, peek_offset: u32, peek_length: u32, - ) -> Result>; + ) -> Result>; /// Receives one datagram without waiting for platform readiness. /// @@ -150,9 +334,9 @@ pub trait PlatformSocket: Send + Sync { /// is smaller, and zero-length datagrams are successful receives. fn receive_from( &self, - data: &mut [u8], + length: usize, flags: ReceiveFromFlags, - ) -> Result>; + ) -> Result>; /// Shuts down one or both socket directions. fn shutdown(&self, mode: ShutdownMode) -> Result>; @@ -182,6 +366,7 @@ impl SocketProvider for UnsupportedSocketProvider { &self, _session_id: SessionId, _request: CreateSocketRequest, + _scope: PlatformSocketScope, _readiness: ReadinessRegistration, ) -> Result> { Err(BrokerError::UnsupportedOperation) @@ -200,6 +385,10 @@ pub fn create( .core .policy .authorize_socket_create(session.caller_credential, request)?; + let scope = session + .core + .policy + .socket_scope(session.caller_credential, request); let quota = Arc::new(SocketQuotaReservation::new(session)?); let reference = session.reserve_object_reference(rights)?; let readiness = ReadinessRegistration::new_with_retirement_guard( @@ -211,10 +400,12 @@ pub fn create( platform_socket: Once::new(), readiness, _quota: quota, + port_reservation: Mutex::new(None), }); let platform_socket = match session.core.socket_provider.create( session.session_id, request, + scope, resource.readiness.clone(), ) { Ok(socket) => socket, @@ -265,10 +456,10 @@ pub fn connect( } if is_udp(create_request) { - return connect_datagram(&object, address); + return connect_datagram(session, &object, address); } - let resource = { + let (resource, needs_bind) = { let mut object = object.write(); let ObjectEntry::Socket(socket) = &mut *object else { return Err(BrokerError::InvalidRights); @@ -283,8 +474,47 @@ pub fn connect( return Ok(SocketOutcome::Completed(socket.connection_status)); } socket.connect_in_flight = true; - Arc::clone(&socket.resource) + (Arc::clone(&socket.resource), socket.local_address.is_none()) }; + if needs_bind { + match session.core.policy.authorize_socket_bind( + session.caller_credential, + create_request, + DEFAULT_TCP_LOCAL_ADDRESS, + ) { + Ok(()) => {} + Err(BrokerError::PolicyDenied) => { + finish_connect(&object, SocketConnectionStatus::Unconnected); + return Ok(SocketOutcome::Failed(SocketError::PolicyDenied)); + } + Err(error) => { + finish_connect(&object, SocketConnectionStatus::Unconnected); + return Err(error); + } + } + let binding = match reserve_and_bind( + session, + create_request, + &resource, + DEFAULT_TCP_LOCAL_ADDRESS, + PlatformBindKind::Implicit, + ) { + Ok(binding) => binding, + Err(error) => { + finish_connect(&object, SocketConnectionStatus::Unconnected); + return Err(error); + } + }; + match binding { + SocketOutcome::Completed((local_address, reservation)) => { + attach_binding(&object, local_address, reservation); + } + SocketOutcome::Failed(error) => { + finish_connect(&object, SocketConnectionStatus::Unconnected); + return Ok(SocketOutcome::Failed(error)); + } + } + } let status = match resource.connect(address) { Ok(SocketConnectionStatus::Unconnected) => { finish_connect(&object, SocketConnectionStatus::Failed(SocketError::Other)); @@ -316,8 +546,10 @@ pub fn bind( let ObjectEntry::Socket(socket) = &mut *object else { return Err(BrokerError::InvalidRights); }; - if socket.configuration_in_flight - || socket.connect_in_flight + if socket.configuration_in_flight { + return Err(BrokerError::WouldBlock); + } + if socket.connect_in_flight || socket.listening || socket.local_address.is_some() || socket.connection_status != SocketConnectionStatus::Unconnected @@ -334,29 +566,35 @@ pub fn bind( ) { Ok(()) => {} Err(BrokerError::PolicyDenied) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Ok(SocketOutcome::Failed(SocketError::PolicyDenied)); } Err(error) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Err(error); } } - let outcome = resource.bind(address); - match outcome { - Ok(SocketOutcome::Completed(local_address)) => { - finish_configuration(&object, Some(local_address), false); - Ok(SocketOutcome::Completed(local_address)) - } - Ok(SocketOutcome::Failed(error)) => { - finish_configuration(&object, None, false); - Ok(SocketOutcome::Failed(error)) - } + let bind_kind = if address.port() == 0 { + PlatformBindKind::Implicit + } else { + PlatformBindKind::Explicit + }; + let binding = match reserve_and_bind(session, create_request, &resource, address, bind_kind) { + Ok(binding) => binding, Err(error) => { - finish_configuration(&object, None, false); - Err(error) + finish_configuration(&object, None, None, false); + return Err(error); } - } + }; + let (local_address, reservation) = match binding { + SocketOutcome::Completed(binding) => binding, + SocketOutcome::Failed(error) => { + finish_configuration(&object, None, None, false); + return Ok(SocketOutcome::Failed(error)); + } + }; + finish_configuration(&object, Some(local_address), Some(reservation), false); + Ok(SocketOutcome::Completed(local_address)) } /// Makes a socket listen for incoming loopback connections. @@ -369,7 +607,7 @@ pub fn listen( return Err(BrokerError::UnsupportedOperation); } let object = session.authorized_object(handle, ObjectRights::WRITE)?; - let (resource, create_request, needs_bind) = { + let (resource, create_request, existing_local_address) = { let mut object = object.write(); let ObjectEntry::Socket(socket) = &mut *object else { return Err(BrokerError::InvalidRights); @@ -387,12 +625,13 @@ pub fn listen( ( Arc::clone(&socket.resource), socket.create_request, - socket.local_address.is_none(), + socket.local_address, ) }; - let mut local_address = None; - if needs_bind { + let mut local_address = existing_local_address; + let mut port_reservation = None; + if local_address.is_none() { match session.core.policy.authorize_socket_bind( session.caller_credential, create_request, @@ -400,39 +639,54 @@ pub fn listen( ) { Ok(()) => {} Err(BrokerError::PolicyDenied) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Ok(SocketOutcome::Failed(SocketError::PolicyDenied)); } Err(error) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Err(error); } } - match resource.bind(DEFAULT_TCP_LISTEN_ADDRESS) { - Ok(SocketOutcome::Completed(address)) => local_address = Some(address), - Ok(SocketOutcome::Failed(error)) => { - finish_configuration(&object, None, false); - return Ok(SocketOutcome::Failed(error)); - } + let binding = match reserve_and_bind( + session, + create_request, + &resource, + DEFAULT_TCP_LISTEN_ADDRESS, + PlatformBindKind::Implicit, + ) { + Ok(binding) => binding, Err(error) => { - finish_configuration(&object, None, false); + finish_configuration(&object, None, None, false); return Err(error); } + }; + match binding { + SocketOutcome::Completed((address, reservation)) => { + local_address = Some(address); + port_reservation = Some(reservation); + } + SocketOutcome::Failed(error) => { + finish_configuration(&object, None, None, false); + return Ok(SocketOutcome::Failed(error)); + } } } match resource.listen(backlog) { - Ok(SocketOutcome::Completed(address)) => { - local_address = Some(address); - finish_configuration(&object, local_address, true); + Ok(SocketOutcome::Completed(())) => { + let Some(address) = local_address else { + finish_configuration(&object, None, port_reservation, false); + return Err(BrokerError::Internal); + }; + finish_configuration(&object, Some(address), port_reservation, true); Ok(SocketOutcome::Completed(address)) } Ok(SocketOutcome::Failed(error)) => { - finish_configuration(&object, local_address, false); + finish_configuration(&object, local_address, port_reservation, false); Ok(SocketOutcome::Failed(error)) } Err(error) => { - finish_configuration(&object, local_address, false); + finish_configuration(&object, local_address, port_reservation, false); Err(error) } } @@ -445,7 +699,7 @@ pub fn accept( readiness_sink: Arc, ) -> Result> { let listener = session.authorized_object(handle, ObjectRights::WAIT)?; - let (listener_resource, create_request) = { + let (listener_resource, create_request, local_address) = { let listener = listener.read(); let ObjectEntry::Socket(socket) = &*listener else { return Err(BrokerError::InvalidRights); @@ -456,7 +710,11 @@ pub fn accept( if !socket.listening { return Ok(SocketOutcome::Failed(SocketError::NotConnected)); } - (Arc::clone(&socket.resource), socket.create_request) + ( + Arc::clone(&socket.resource), + socket.create_request, + socket.local_address.ok_or(BrokerError::Internal)?, + ) }; let rights = session .core @@ -473,6 +731,7 @@ pub fn accept( platform_socket: Once::new(), readiness, _quota: quota, + port_reservation: Mutex::new(None), }); let accepted = match listener_resource.accept(resource.readiness.clone()) { Ok(SocketOutcome::Completed(accepted)) => accepted, @@ -486,12 +745,11 @@ pub fn accept( } }; resource.platform_socket.call_once(|| accepted.socket); - let accepted_socket = - SocketObject::new_connected(resource, create_request, accepted.local_address); + let accepted_socket = SocketObject::new_connected(resource, create_request, local_address); let handle = reference.commit(ObjectEntry::Socket(accepted_socket))?; Ok(SocketOutcome::Completed(AcceptedBrokerSocket { handle, - local_address: accepted.local_address, + local_address, remote_address: accepted.remote_address, })) } @@ -506,13 +764,34 @@ pub fn send( if flags.has_unsupported_bits() { return Err(BrokerError::UnsupportedOperation); } - let (resource, create_request, _) = socket_state(session, handle, ObjectRights::WRITE)?; + send_owned(session, handle, copy_to_vec(data)?, flags) +} + +/// Sends owned bytes without waiting for readiness. +/// +/// Ownership is passed through to the platform so threaded implementations can +/// avoid copying the payload into a separate command buffer. +pub fn send_owned( + session: &BrokerSession, + handle: ObjectHandle, + data: Vec, + flags: SendFlags, +) -> Result> { + if flags.has_unsupported_bits() { + return Err(BrokerError::UnsupportedOperation); + } + let length = data.len(); + let AuthorizedSocketState { + resource, + create_request, + .. + } = authorized_socket_state(session, handle, ObjectRights::WRITE)?; if !is_tcp(create_request) { return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); } let outcome = resource.send(data, flags)?; if let SocketOutcome::Completed(sent) = outcome - && sent > data.len() + && sent > length { return Err(BrokerError::Internal); } @@ -530,8 +809,30 @@ pub fn send_to( if flags.has_unsupported_bits() || data.len() > MAX_UDP_DATAGRAM_SIZE as usize { return Err(BrokerError::UnsupportedOperation); } - let (resource, create_request, connection_status) = - socket_state(session, handle, ObjectRights::WRITE)?; + send_to_owned(session, handle, copy_to_vec(data)?, flags, destination) +} + +/// Sends one owned datagram without waiting for readiness. +/// +/// Ownership is passed through to the platform so threaded implementations can +/// avoid copying the payload into a separate command buffer. +pub fn send_to_owned( + session: &BrokerSession, + handle: ObjectHandle, + data: Vec, + flags: SendFlags, + destination: Option, +) -> Result> { + if flags.has_unsupported_bits() || data.len() > MAX_UDP_DATAGRAM_SIZE as usize { + return Err(BrokerError::UnsupportedOperation); + } + let length = data.len(); + let AuthorizedSocketState { + object, + resource, + create_request, + connection_status, + } = authorized_socket_state(session, handle, ObjectRights::WRITE)?; if !is_udp(create_request) { return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); } @@ -550,9 +851,12 @@ pub fn send_to( } else if connection_status != SocketConnectionStatus::Connected { return Ok(SocketOutcome::Failed(SocketError::NotConnected)); } + if let SocketOutcome::Failed(error) = ensure_datagram_bound(session, &object)? { + return Ok(SocketOutcome::Failed(error)); + } let outcome = resource.send_to(data, flags, destination)?; if let SocketOutcome::Completed(sent) = outcome - && sent != data.len() + && sent != length { return Err(BrokerError::Internal); } @@ -568,12 +872,38 @@ pub fn receive( peek_offset: u32, peek_length: u32, ) -> Result> { - if flags.has_unsupported_bits() { + match receive_owned(session, handle, data.len(), flags, peek_offset, peek_length)? { + SocketOutcome::Completed(PlatformStreamReceive::Received(received)) => { + data[..received.len()].copy_from_slice(&received); + Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received( + received + .len() + .try_into() + .map_err(|_| BrokerError::Internal)?, + ))) + } + SocketOutcome::Completed(PlatformStreamReceive::EndOfStream) => { + Ok(SocketOutcome::Completed(ReceiveSocketResponse::EndOfStream)) + } + SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), + } +} + +/// Receives stream bytes into an owned platform buffer without waiting. +pub fn receive_owned( + session: &BrokerSession, + handle: ObjectHandle, + length: usize, + flags: ReceiveFlags, + peek_offset: u32, + peek_length: u32, +) -> Result> { + if flags.has_unsupported_bits() || length > MAX_SOCKET_TRANSFER_SIZE as usize { return Err(BrokerError::UnsupportedOperation); } let peek = flags.contains(ReceiveFlags::PEEK); let end = peek_offset - .checked_add(data.len().try_into().map_err(|_| BrokerError::Internal)?) + .checked_add(length.try_into().map_err(|_| BrokerError::Internal)?) .ok_or(BrokerError::UnsupportedOperation)?; let canonical_peek_length = peek_length .checked_sub(peek_offset) @@ -581,22 +911,28 @@ pub fn receive( if (!peek && (peek_offset != 0 || peek_length != 0)) || (peek && (!peek_offset.is_multiple_of(MAX_SOCKET_TRANSFER_SIZE) - || canonical_peek_length != data.len().try_into().ok() + || canonical_peek_length != length.try_into().ok() || peek_length < end || peek_length > litebox_broker_protocol::socket::MAX_SOCKET_PEEK_SIZE)) { return Err(BrokerError::UnsupportedOperation); } - let (resource, create_request, _) = socket_state(session, handle, ObjectRights::WAIT)?; + let AuthorizedSocketState { + resource, + create_request, + .. + } = authorized_socket_state(session, handle, ObjectRights::WAIT)?; if !is_tcp(create_request) { return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); } - if data.is_empty() { - return Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received(0))); + if length == 0 { + return Ok(SocketOutcome::Completed(PlatformStreamReceive::Received( + Vec::new(), + ))); } - let outcome = resource.receive(data, flags, peek_offset, peek_length)?; - if let SocketOutcome::Completed(ReceiveSocketResponse::Received(received)) = outcome - && (received as usize > data.len() || received == 0) + let outcome = resource.receive(length, flags, peek_offset, peek_length)?; + if let SocketOutcome::Completed(PlatformStreamReceive::Received(received)) = &outcome + && (received.len() > length || received.is_empty()) { return Err(BrokerError::Internal); } @@ -610,17 +946,45 @@ pub fn receive_from( data: &mut [u8], flags: ReceiveFromFlags, ) -> Result> { - if flags.has_unsupported_bits() || data.len() > MAX_UDP_DATAGRAM_SIZE as usize { + match receive_from_owned(session, handle, data.len(), flags)? { + SocketOutcome::Completed(received) => { + data[..received.data.len()].copy_from_slice(&received.data); + Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { + received: received.data.len(), + datagram_length: received.datagram_length, + source_address: received.source_address, + })) + } + SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), + } +} + +/// Receives one datagram into an owned platform buffer without waiting. +pub fn receive_from_owned( + session: &BrokerSession, + handle: ObjectHandle, + length: usize, + flags: ReceiveFromFlags, +) -> Result> { + if flags.has_unsupported_bits() || length > MAX_UDP_DATAGRAM_SIZE as usize { return Err(BrokerError::UnsupportedOperation); } - let (resource, create_request, _) = socket_state(session, handle, ObjectRights::WAIT)?; + let AuthorizedSocketState { + object, + resource, + create_request, + connection_status: _, + } = authorized_socket_state(session, handle, ObjectRights::WAIT)?; if !is_udp(create_request) { return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); } - let outcome = resource.receive_from(data, flags)?; - if let SocketOutcome::Completed(received) = outcome - && (received.received > data.len() - || received.datagram_length < received.received + if let SocketOutcome::Failed(error) = ensure_datagram_bound(session, &object)? { + return Ok(SocketOutcome::Failed(error)); + } + let outcome = resource.receive_from(length, flags)?; + if let SocketOutcome::Completed(received) = &outcome + && (received.data.len() > length + || received.datagram_length < received.data.len() || received.datagram_length > MAX_UDP_DATAGRAM_SIZE as usize) { return Err(BrokerError::Internal); @@ -628,13 +992,26 @@ pub fn receive_from( Ok(outcome) } +fn copy_to_vec(data: &[u8]) -> Result> { + let mut owned = Vec::new(); + owned + .try_reserve_exact(data.len()) + .map_err(|_| BrokerError::OutOfMemory)?; + owned.extend_from_slice(data); + Ok(owned) +} + /// Sets a typed option on a broker-owned TCP socket. pub fn set_tcp_option( session: &BrokerSession, handle: ObjectHandle, value: TcpOptionValue, ) -> Result<()> { - let (resource, create_request, _) = socket_state(session, handle, ObjectRights::WRITE)?; + let AuthorizedSocketState { + resource, + create_request, + .. + } = authorized_socket_state(session, handle, ObjectRights::WRITE)?; if !is_tcp(create_request) { return Err(BrokerError::UnsupportedOperation); } @@ -647,7 +1024,11 @@ pub fn get_tcp_option( handle: ObjectHandle, name: TcpOptionName, ) -> Result { - let (resource, create_request, _) = socket_state(session, handle, ObjectRights::WAIT)?; + let AuthorizedSocketState { + resource, + create_request, + .. + } = authorized_socket_state(session, handle, ObjectRights::WAIT)?; if !is_tcp(create_request) { return Err(BrokerError::UnsupportedOperation); } @@ -743,7 +1124,7 @@ pub fn status(session: &BrokerSession, handle: ObjectHandle) -> Result Result>, + resource: Arc, + create_request: CreateSocketRequest, + connection_status: SocketConnectionStatus, +} + +fn authorized_socket_state( session: &BrokerSession, handle: ObjectHandle, required_rights: ObjectRights, -) -> Result<( - Arc, - CreateSocketRequest, - SocketConnectionStatus, -)> { +) -> Result { let object = session.authorized_object(handle, required_rights)?; - let object = object.read(); - let ObjectEntry::Socket(socket) = &*object else { - return Err(BrokerError::InvalidRights); + let (resource, create_request, connection_status) = { + let entry = object.read(); + let ObjectEntry::Socket(socket) = &*entry else { + return Err(BrokerError::InvalidRights); + }; + ( + Arc::clone(&socket.resource), + socket.create_request, + socket.connection_status, + ) }; - Ok(( - Arc::clone(&socket.resource), - socket.create_request, - socket.connection_status, - )) + Ok(AuthorizedSocketState { + object, + resource, + create_request, + connection_status, + }) +} + +/// Ensures an authorized UDP socket has a guest-local binding before datagram I/O. +/// +/// The caller supplies the object retained by its initial authorization so an +/// in-flight operation can complete even if the handle is concurrently closed. +/// An unbound socket is implicitly bound to a guest wildcard address while +/// excluding concurrent configuration, connect, and listen transitions. +fn ensure_datagram_bound( + session: &BrokerSession, + object: &spin::RwLock, +) -> Result> { + let (resource, create_request) = { + let mut object = object.write(); + let ObjectEntry::Socket(socket) = &mut *object else { + return Err(BrokerError::InvalidRights); + }; + if !is_udp(socket.create_request) { + return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); + } + if socket.local_address.is_some() { + return Ok(SocketOutcome::Completed(())); + } + if socket.configuration_in_flight || socket.connect_in_flight || socket.listening { + return Err(BrokerError::WouldBlock); + } + socket.configuration_in_flight = true; + (Arc::clone(&socket.resource), socket.create_request) + }; + // Linux-compatible implicit UDP binding uses a guest wildcard address. + // The platform maps it to a private backend endpoint rather than binding a + // wildcard host endpoint. + let binding = match reserve_and_bind( + session, + create_request, + &resource, + DEFAULT_UDP_LOCAL_ADDRESS, + PlatformBindKind::Implicit, + ) { + Ok(binding) => binding, + Err(error) => { + finish_configuration(object, None, None, false); + return Err(error); + } + }; + match binding { + SocketOutcome::Completed((local_address, reservation)) => { + finish_configuration(object, Some(local_address), Some(reservation), false); + Ok(SocketOutcome::Completed(())) + } + SocketOutcome::Failed(error) => { + finish_configuration(object, None, None, false); + Ok(SocketOutcome::Failed(error)) + } + } +} + +fn reserve_and_bind( + session: &BrokerSession, + create_request: CreateSocketRequest, + resource: &SocketResource, + requested_address: SocketAddrV4, + kind: PlatformBindKind, +) -> Result> { + let (local_address, reservation) = + match session + .socket_ports + .reserve(create_request, requested_address, |port| { + session + .core + .socket_provider + .reserves_guest_port(create_request, port) + })? { + SocketOutcome::Completed(binding) => binding, + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + }; + match resource.bind(local_address, kind)? { + SocketOutcome::Completed(()) => Ok(SocketOutcome::Completed((local_address, reservation))), + SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), + } } fn connect_datagram( + session: &BrokerSession, object: &spin::RwLock, address: SocketAddrV4, ) -> Result> { - let (resource, previous_status) = { + let (resource, create_request, previous_status, needs_bind) = { let mut object = object.write(); let ObjectEntry::Socket(socket) = &mut *object else { return Err(BrokerError::InvalidRights); }; - if socket.configuration_in_flight || socket.connect_in_flight || socket.listening { + if socket.configuration_in_flight { + return Err(BrokerError::WouldBlock); + } + if socket.connect_in_flight || socket.listening { return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); } socket.configuration_in_flight = true; - (Arc::clone(&socket.resource), socket.connection_status) + ( + Arc::clone(&socket.resource), + socket.create_request, + socket.connection_status, + socket.local_address.is_none(), + ) }; + if needs_bind { + // This wildcard address belongs only to the guest namespace; the + // platform chooses a private backend endpoint. + let binding = match reserve_and_bind( + session, + create_request, + &resource, + DEFAULT_UDP_LOCAL_ADDRESS, + PlatformBindKind::Implicit, + ) { + Ok(binding) => binding, + Err(error) => { + finish_datagram_connect(object, previous_status, false, None); + return Err(error); + } + }; + match binding { + SocketOutcome::Completed((local_address, reservation)) => { + attach_binding(object, local_address, reservation); + } + SocketOutcome::Failed(error) => { + finish_datagram_connect(object, previous_status, false, None); + return Ok(SocketOutcome::Failed(error)); + } + } + } match resource.connect(address) { Ok(SocketConnectionStatus::Connected) => { - finish_datagram_connect(object, SocketConnectionStatus::Connected, true); + finish_datagram_connect( + object, + SocketConnectionStatus::Connected, + true, + Some(address), + ); Ok(SocketOutcome::Completed(SocketConnectionStatus::Connected)) } Ok(SocketConnectionStatus::Failed(error)) => { - finish_datagram_connect(object, previous_status, false); + finish_datagram_connect(object, previous_status, false, None); Ok(SocketOutcome::Failed(error)) } Ok(_) => { @@ -882,11 +1390,12 @@ fn connect_datagram( object, SocketConnectionStatus::Failed(SocketError::Other), true, + None, ); Err(BrokerError::Internal) } Err(PlatformConnectError::PeerUnchanged(error)) => { - finish_datagram_connect(object, previous_status, false); + finish_datagram_connect(object, previous_status, false, None); Err(error) } Err(PlatformConnectError::PeerIndeterminate(error)) => { @@ -894,6 +1403,7 @@ fn connect_datagram( object, SocketConnectionStatus::Failed(SocketError::Other), true, + None, ); Err(error) } @@ -904,11 +1414,18 @@ fn finish_datagram_connect( object: &spin::RwLock, status: SocketConnectionStatus, peer_state_changed: bool, + connected_address: Option, ) { let mut object = object.write(); if let ObjectEntry::Socket(socket) = &mut *object { socket.configuration_in_flight = false; socket.connection_status = status; + if let (Some(local_address), Some(_)) = (socket.local_address, connected_address) + && local_address.ip().is_unspecified() + { + socket.local_address = + Some(SocketAddrV4::new(Ipv4Addr::LOCALHOST, local_address.port())); + } if peer_state_changed { // This distinguishes status snapshots from before a successful // peer replacement or an indeterminate platform failure. A status @@ -927,15 +1444,31 @@ fn finish_connect(object: &spin::RwLock, status: SocketConnectionSt } } +fn attach_binding( + object: &spin::RwLock, + local_address: SocketAddrV4, + port_reservation: GuestPortReservation, +) { + let mut object = object.write(); + if let ObjectEntry::Socket(socket) = &mut *object { + socket.local_address = Some(local_address); + socket.resource.set_port_reservation(port_reservation); + } +} + fn finish_configuration( object: &spin::RwLock, local_address: Option, + port_reservation: Option, listening: bool, ) { let mut object = object.write(); if let ObjectEntry::Socket(socket) = &mut *object { socket.configuration_in_flight = false; socket.local_address = socket.local_address.or(local_address); + if let Some(port_reservation) = port_reservation { + socket.resource.set_port_reservation(port_reservation); + } socket.listening |= listening; } } @@ -991,9 +1524,18 @@ pub(crate) struct SocketResource { platform_socket: Once>, readiness: ReadinessRegistration, _quota: Arc, + port_reservation: Mutex>, } impl SocketResource { + fn set_port_reservation(&self, reservation: GuestPortReservation) { + let mut slot = self.port_reservation.lock(); + debug_assert!(slot.is_none()); + if slot.is_none() { + *slot = Some(reservation); + } + } + fn platform_socket(&self) -> &dyn PlatformSocket { self.platform_socket .get() @@ -1008,11 +1550,11 @@ impl SocketResource { self.platform_socket().connect(address) } - fn bind(&self, address: SocketAddrV4) -> Result> { - self.platform_socket().bind(address) + fn bind(&self, address: SocketAddrV4, kind: PlatformBindKind) -> Result> { + self.platform_socket().bind(address, kind) } - fn listen(&self, backlog: u32) -> Result> { + fn listen(&self, backlog: u32) -> Result> { self.platform_socket().listen(backlog) } @@ -1023,13 +1565,13 @@ impl SocketResource { self.platform_socket().accept(readiness) } - fn send(&self, data: &[u8], flags: SendFlags) -> Result> { + fn send(&self, data: Vec, flags: SendFlags) -> Result> { self.platform_socket().send(data, flags) } fn send_to( &self, - data: &[u8], + data: Vec, flags: SendFlags, destination: Option, ) -> Result> { @@ -1038,21 +1580,21 @@ impl SocketResource { fn receive( &self, - data: &mut [u8], + length: usize, flags: ReceiveFlags, peek_offset: u32, peek_length: u32, - ) -> Result> { + ) -> Result> { self.platform_socket() - .receive(data, flags, peek_offset, peek_length) + .receive(length, flags, peek_offset, peek_length) } fn receive_from( &self, - data: &mut [u8], + length: usize, flags: ReceiveFromFlags, - ) -> Result> { - self.platform_socket().receive_from(data, flags) + ) -> Result> { + self.platform_socket().receive_from(length, flags) } fn shutdown(&self, mode: ShutdownMode) -> Result> { @@ -1090,6 +1632,16 @@ const fn is_udp(request: CreateSocketRequest) -> bool { ) } +const fn guest_port_protocol(request: CreateSocketRequest) -> Option { + if is_tcp(request) { + Some(GuestPortProtocol::Tcp) + } else if is_udp(request) { + Some(GuestPortProtocol::Udp) + } else { + None + } +} + impl Drop for SocketResource { fn drop(&mut self) { self.readiness.retire(); @@ -1156,9 +1708,132 @@ pub(crate) mod tests { state: Arc, } + #[test] + fn guest_port_namespaces_are_per_session_and_protocol() { + let first_session = SessionSocketPorts::default(); + let second_session = SessionSocketPorts::default(); + let address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + + let SocketOutcome::Completed((_, first_tcp)) = first_session + .reserve(create_request(), address, |_| false) + .unwrap() + else { + panic!("first TCP reservation failed"); + }; + assert!(matches!( + first_session.reserve(create_request(), address, |_| false), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + )); + assert!(matches!( + first_session.reserve(create_udp_request(), address, |_| false), + Ok(SocketOutcome::Completed(_)) + )); + assert!(matches!( + second_session.reserve(create_request(), address, |_| false), + Ok(SocketOutcome::Completed(_)) + )); + + drop(first_tcp); + assert!(matches!( + first_session.reserve(create_request(), address, |_| false), + Ok(SocketOutcome::Completed(_)) + )); + } + + #[test] + fn implicit_guest_port_allocation_skips_provider_reservations() { + let ports = SessionSocketPorts::default(); + let requested_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0); + let SocketOutcome::Completed((address, _reservation)) = ports + .reserve(create_request(), requested_address, |port| { + port == FIRST_EPHEMERAL_PORT + }) + .unwrap() + else { + panic!("ephemeral TCP reservation failed"); + }; + + assert_eq!(address.port(), FIRST_EPHEMERAL_PORT + 1); + } + + fn check_connected_loopback_udp_uses_the_canonical_guest_local_address(broker: &BrokerCore) { + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let handle = create( + &session, + create_udp_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + let destination = SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 2), 53); + + assert_eq!( + connect(&session, handle, destination), + Ok(SocketOutcome::Completed(SocketConnectionStatus::Connected)) + ); + assert_eq!( + status(&session, handle).unwrap().local_address, + Some(SocketAddrV4::new(Ipv4Addr::LOCALHOST, FIRST_EPHEMERAL_PORT,)) + ); + } + + fn check_concurrent_udp_implicit_bind_is_retryable_for_all_operations( + broker: &BrokerCore, + provider: &TestSocketProvider, + ) { + let session = Arc::new( + broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(), + ); + let handle = create( + &session, + create_udp_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + let destination = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53); + let (bind_started, release_bind) = provider.block_next_bind(); + let first_session = Arc::clone(&session); + let first = std::thread::spawn(move || { + send_to( + &first_session, + handle, + b"a", + SendFlags::NONE, + Some(destination), + ) + }); + bind_started.recv_timeout(Duration::from_secs(5)).unwrap(); + + assert_eq!( + send_to(&session, handle, b"b", SendFlags::NONE, Some(destination)), + Err(BrokerError::WouldBlock) + ); + assert_eq!( + bind( + &session, + handle, + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 80), + ), + Err(BrokerError::WouldBlock) + ); + assert_eq!( + connect(&session, handle, destination), + Err(BrokerError::WouldBlock) + ); + release_bind.send(()).unwrap(); + assert_eq!(first.join().unwrap(), Ok(SocketOutcome::Completed(1))); + assert_eq!( + send_to(&session, handle, b"c", SendFlags::NONE, Some(destination)), + Ok(SocketOutcome::Completed(1)) + ); + } + #[derive(Default)] struct TestSocketState { - creates: StdMutex>, + creates: StdMutex>, closed_sessions: StdMutex>, sent: StdMutex>, connect_calls: AtomicUsize, @@ -1166,6 +1841,7 @@ pub(crate) mod tests { status_responses: StdMutex>, status_block: StdMutex, mpsc::Receiver<()>)>>, binds: StdMutex>, + bind_block: StdMutex, mpsc::Receiver<()>)>>, listens: StdMutex>, listen_block: StdMutex, mpsc::Receiver<()>)>>, shutdown_calls: AtomicUsize, @@ -1197,6 +1873,13 @@ pub(crate) mod tests { fn fail_next_shutdown(&self) { self.state.fail_shutdown.store(true, Ordering::Relaxed); } + + fn block_next_bind(&self) -> (mpsc::Receiver<()>, mpsc::Sender<()>) { + let (started_tx, started_rx) = mpsc::channel(); + let (release_tx, release_rx) = mpsc::channel(); + *self.state.bind_block.lock().unwrap() = Some((started_tx, release_rx)); + (started_rx, release_tx) + } } impl SocketProvider for TestSocketProvider { @@ -1204,13 +1887,14 @@ pub(crate) mod tests { &self, session_id: SessionId, request: CreateSocketRequest, + scope: PlatformSocketScope, readiness: ReadinessRegistration, ) -> Result> { self.state .creates .lock() .unwrap() - .push((session_id, request)); + .push((session_id, request, scope)); if self.state.fail_create.swap(false, Ordering::Relaxed) { *self.state.failed_readiness.lock().unwrap() = Some(readiness); return Err(BrokerError::OutOfMemory); @@ -1243,27 +1927,27 @@ pub(crate) mod tests { } impl PlatformSocket for TestPlatformSocket { - fn bind(&self, address: SocketAddrV4) -> Result> { + fn bind( + &self, + address: SocketAddrV4, + _kind: PlatformBindKind, + ) -> Result> { self.state.binds.lock().unwrap().push(address); - let address = if address.port() == 0 { - SocketAddrV4::new(*address.ip(), 49152) - } else { - address - }; - Ok(SocketOutcome::Completed(address)) + if let Some((started, release)) = self.state.bind_block.lock().unwrap().take() { + started.send(()).unwrap(); + release.recv_timeout(Duration::from_secs(5)).unwrap(); + } + Ok(SocketOutcome::Completed(())) } - fn listen(&self, backlog: u32) -> Result> { + fn listen(&self, backlog: u32) -> Result> { self.state.listens.lock().unwrap().push(backlog); let listen_block = self.state.listen_block.lock().unwrap().take(); if let Some((started, release)) = listen_block { started.send(()).unwrap(); release.recv_timeout(Duration::from_secs(5)).unwrap(); } - Ok(SocketOutcome::Completed(SocketAddrV4::new( - Ipv4Addr::LOCALHOST, - 49152, - ))) + Ok(SocketOutcome::Completed(())) } fn accept( @@ -1300,44 +1984,42 @@ pub(crate) mod tests { } } - fn send(&self, data: &[u8], _flags: SendFlags) -> Result> { - self.state.sent.lock().unwrap().extend_from_slice(data); + fn send(&self, data: Vec, _flags: SendFlags) -> Result> { + self.state.sent.lock().unwrap().extend_from_slice(&data); Ok(SocketOutcome::Completed(data.len())) } fn send_to( &self, - data: &[u8], + data: Vec, _flags: SendFlags, _destination: Option, ) -> Result> { - self.state.sent.lock().unwrap().extend_from_slice(data); + self.state.sent.lock().unwrap().extend_from_slice(&data); Ok(SocketOutcome::Completed(data.len())) } fn receive( &self, - data: &mut [u8], + length: usize, _flags: ReceiveFlags, _peek_offset: u32, _peek_length: u32, - ) -> Result> { - let received = data.len().min(2); - data[..received].copy_from_slice(&[7, 9][..received]); - Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received( - u32::try_from(received).unwrap(), + ) -> Result> { + let received = length.min(2); + Ok(SocketOutcome::Completed(PlatformStreamReceive::Received( + [7, 9][..received].to_vec(), ))) } fn receive_from( &self, - data: &mut [u8], + length: usize, _flags: ReceiveFromFlags, - ) -> Result> { - let received = data.len().min(2); - data[..received].copy_from_slice(&[7, 9][..received]); - Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { - received, + ) -> Result> { + let received = length.min(2); + Ok(SocketOutcome::Completed(PlatformDatagramReceive { + data: [7, 9][..received].to_vec(), datagram_length: 4, source_address: SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49153), })) @@ -1408,6 +2090,9 @@ pub(crate) mod tests { check_socket_operations_and_policy(broker, provider); check_tcp_option_state_is_per_socket(broker); check_udp_socket_operations(broker, provider); + check_connected_loopback_udp_uses_the_canonical_guest_local_address(broker); + check_concurrent_udp_implicit_bind_is_retryable_for_all_operations(broker, provider); + check_udp_implicit_bind_survives_handle_close(broker); check_concurrent_udp_status_does_not_regress_connection(broker, provider); check_server_socket_operations(broker, provider); check_failed_listener_shutdown_preserves_state(broker, provider); @@ -1419,6 +2104,43 @@ pub(crate) mod tests { check_socket_quotas(broker); } + fn check_udp_implicit_bind_survives_handle_close(broker: &BrokerCore) { + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let handle = create( + &session, + create_udp_request(), + Arc::new(TestReadinessSink::default()), + ) + .unwrap(); + let AuthorizedSocketState { + object, + resource, + create_request, + connection_status: _, + } = authorized_socket_state(&session, handle, ObjectRights::WRITE).unwrap(); + assert!(is_udp(create_request)); + + session.close_object_reference(handle).unwrap(); + assert_eq!( + ensure_datagram_bound(&session, &object), + Ok(SocketOutcome::Completed(())) + ); + assert_eq!( + resource.send_to( + b"x".to_vec(), + SendFlags::NONE, + Some(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53)), + ), + Ok(SocketOutcome::Completed(1)) + ); + drop(object); + drop(resource); + assert_eq!(broker.reserved_sockets.load(Ordering::Relaxed), 0); + assert_eq!(session.reserved_sockets.load(Ordering::Relaxed), 0); + } + fn check_failed_create_rolls_back(broker: &BrokerCore, provider: &TestSocketProvider) { let session = broker .create_session(CallerCredential::Unauthenticated) @@ -1456,7 +2178,11 @@ pub(crate) mod tests { let handle = create(&session, create_request(), readiness.clone()).unwrap(); assert_eq!( provider.state.creates.lock().unwrap().last(), - Some(&(session.session_id, create_request())) + Some(&( + session.session_id, + create_request(), + PlatformSocketScope::LoopbackOnly, + )) ); assert_eq!(broker.reserved_sockets.load(Ordering::Relaxed), 1); assert_eq!(session.reserved_sockets.load(Ordering::Relaxed), 1); @@ -1489,7 +2215,7 @@ pub(crate) mod tests { status(&session, handle), Ok(SocketStatusResponse { status: SocketConnectionStatus::Connected, - local_address: None, + local_address: Some(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152)), pending_error: None, }) ); @@ -1511,6 +2237,17 @@ pub(crate) mod tests { receive(&session, handle, &mut data[..1], ReceiveFlags::PEEK, 1, 2), Err(BrokerError::UnsupportedOperation) ); + assert_eq!( + receive_owned( + &session, + handle, + MAX_SOCKET_TRANSFER_SIZE as usize + 1, + ReceiveFlags::NONE, + 0, + 0, + ), + Err(BrokerError::UnsupportedOperation) + ); assert_eq!( set_tcp_option(&session, handle, TcpOptionValue::NoDelay(true)), Ok(()) @@ -1698,7 +2435,7 @@ pub(crate) mod tests { status(&session, handle), Ok(SocketStatusResponse { status: SocketConnectionStatus::Failed(SocketError::Other), - local_address: None, + local_address: Some(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152)), pending_error: Some(SocketError::ConnectionRefused), }) ); @@ -1814,13 +2551,14 @@ pub(crate) mod tests { assert_eq!(broker.reserved_sockets.load(Ordering::Relaxed), 0); let auto_bound = create(&session, create_request(), readiness).unwrap(); + let auto_bound_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49153); assert_eq!( listen(&session, auto_bound, 0), - Ok(SocketOutcome::Completed(local_address)) + Ok(SocketOutcome::Completed(auto_bound_address)) ); assert_eq!( provider.state.binds.lock().unwrap().last(), - Some(&DEFAULT_TCP_LISTEN_ADDRESS) + Some(&auto_bound_address) ); session.close_object_reference(auto_bound).unwrap(); } @@ -1840,7 +2578,7 @@ pub(crate) mod tests { Arc::new(TestReadinessSink::default()), ) .unwrap(); - let local_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 49152); + let local_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152); *provider.state.status_responses.lock().unwrap() = std::collections::VecDeque::from([ SocketStatusResponse { status: SocketConnectionStatus::Unconnected, @@ -1911,7 +2649,7 @@ pub(crate) mod tests { next_status.pending_error, Some(SocketError::NetworkUnreachable) ); - assert_eq!(next_status.local_address, Some(next_local_address)); + assert_eq!(next_status.local_address, Some(local_address)); assert_eq!( status(&session, handle).unwrap().status, SocketConnectionStatus::Connected @@ -2154,7 +2892,7 @@ pub(crate) mod tests { status(&session, poisoned), Ok(SocketStatusResponse { status: SocketConnectionStatus::Failed(SocketError::Other), - local_address: None, + local_address: Some(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49153)), pending_error: None, }) ); @@ -2242,7 +2980,7 @@ pub(crate) mod tests { Ok(SocketOutcome::Completed(SocketConnectionStatus::Connecting)) ); - let local_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49153); + let platform_local_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49153); provider .state .status_responses @@ -2250,13 +2988,14 @@ pub(crate) mod tests { .unwrap() .push_back(SocketStatusResponse { status: SocketConnectionStatus::Failed(SocketError::TimedOut), - local_address: Some(local_address), + local_address: Some(platform_local_address), pending_error: None, }); + let guest_local_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152); let expected = SocketStatusResponse { status: SocketConnectionStatus::Failed(SocketError::TimedOut), - local_address: Some(local_address), + local_address: Some(guest_local_address), pending_error: None, }; assert_eq!(status(&session, handle), Ok(expected)); diff --git a/litebox_broker_host/src/lib.rs b/litebox_broker_host/src/lib.rs index ad4e325c59..13f2a16c14 100644 --- a/litebox_broker_host/src/lib.rs +++ b/litebox_broker_host/src/lib.rs @@ -438,8 +438,13 @@ fn handle_socket_request( shared_buffers .read(request.buffer.slot_index, &mut data) .map_err(|_| RequestFailure::Abort(ErrorCode::Internal))?; - match litebox_broker_core::socket::send(session, request.handle, &data, request.flags) - .map_err(RequestFailure::from)? + match litebox_broker_core::socket::send_owned( + session, + request.handle, + data, + request.flags, + ) + .map_err(RequestFailure::from)? { SocketOutcome::Completed(sent) => { let sent = sent @@ -464,10 +469,10 @@ fn handle_socket_request( shared_buffers .read(request.buffer.slot_index, &mut data) .map_err(|_| RequestFailure::Abort(ErrorCode::Internal))?; - match litebox_broker_core::socket::send_to( + match litebox_broker_core::socket::send_to_owned( session, request.handle, - &data, + data, request.flags, request.destination, ) @@ -502,27 +507,33 @@ fn handle_socket_request( return Err(RequestFailure::Abort(ErrorCode::MalformedRequest)); } let length = request.buffer.length as usize; - let mut data = Vec::new(); - if data.try_reserve_exact(length).is_err() { - return Err(RequestFailure::Respond(ErrorCode::OutOfMemory)); - } - data.resize(length, 0); - match litebox_broker_core::socket::receive( + match litebox_broker_core::socket::receive_owned( session, request.handle, - &mut data, + length, request.flags, request.peek_offset, request.peek_length, ) .map_err(RequestFailure::from)? { - SocketOutcome::Completed(response) => { - if let ReceiveSocketResponse::Received(received) = response { - shared_buffers - .write(request.buffer.slot_index, &data[..received as usize]) - .map_err(|_| RequestFailure::Abort(ErrorCode::Internal))?; - } + SocketOutcome::Completed(received) => { + let response = match received { + litebox_broker_core::socket::PlatformStreamReceive::Received(data) => { + let received = data.len(); + shared_buffers + .write(request.buffer.slot_index, &data) + .map_err(|_| RequestFailure::Abort(ErrorCode::Internal))?; + ReceiveSocketResponse::Received( + received + .try_into() + .map_err(|_| RequestFailure::Abort(ErrorCode::Internal))?, + ) + } + litebox_broker_core::socket::PlatformStreamReceive::EndOfStream => { + ReceiveSocketResponse::EndOfStream + } + }; Ok(SocketResponse::Receive(response)) } SocketOutcome::Failed(error) => Ok(SocketResponse::Failed(error)), @@ -534,26 +545,22 @@ fn handle_socket_request( return Err(RequestFailure::Abort(ErrorCode::MalformedRequest)); } let length = request.buffer.length as usize; - let mut data = Vec::new(); - if data.try_reserve_exact(length).is_err() { - return Err(RequestFailure::Respond(ErrorCode::OutOfMemory)); - } - data.resize(length, 0); - match litebox_broker_core::socket::receive_from( + match litebox_broker_core::socket::receive_from_owned( session, request.handle, - &mut data, + length, request.flags, ) .map_err(RequestFailure::from)? { SocketOutcome::Completed(received) => { shared_buffers - .write(request.buffer.slot_index, &data[..received.received]) + .write(request.buffer.slot_index, &received.data) .map_err(|_| RequestFailure::Abort(ErrorCode::Internal))?; Ok(SocketResponse::ReceiveFrom(ReceiveFromSocketResponse { received: received - .received + .data + .len() .try_into() .map_err(|_| RequestFailure::Abort(ErrorCode::Internal))?, datagram_length: received @@ -695,8 +702,8 @@ mod tests { use core::net::{Ipv4Addr, SocketAddrV4}; use litebox_broker_core::readiness::ReadinessRegistration; use litebox_broker_core::socket::{ - AcceptedPlatformSocket, PlatformConnectError, PlatformSocket, ReceivedPlatformDatagram, - SocketProvider, + AcceptedPlatformSocket, PlatformConnectError, PlatformDatagramReceive, PlatformSocket, + PlatformStreamReceive, SocketProvider, }; use litebox_broker_core::{ObjectRights, PolicyEngine, SessionId, SocketPolicy}; use litebox_broker_protocol::event::{ @@ -759,6 +766,7 @@ mod tests { &self, _session_id: SessionId, request: CreateSocketRequest, + _scope: litebox_broker_core::socket::PlatformSocketScope, readiness: ReadinessRegistration, ) -> litebox_broker_core::Result> { Ok(Arc::new(TestPlatformSocket { @@ -778,24 +786,14 @@ mod tests { impl PlatformSocket for TestPlatformSocket { fn bind( &self, - address: SocketAddrV4, - ) -> litebox_broker_core::Result> { - let address = if address.port() == 0 { - SocketAddrV4::new(*address.ip(), 49152) - } else { - address - }; - Ok(SocketOutcome::Completed(address)) + _address: SocketAddrV4, + _kind: litebox_broker_core::socket::PlatformBindKind, + ) -> litebox_broker_core::Result> { + Ok(SocketOutcome::Completed(())) } - fn listen( - &self, - _backlog: u32, - ) -> litebox_broker_core::Result> { - Ok(SocketOutcome::Completed(SocketAddrV4::new( - Ipv4Addr::LOCALHOST, - 49152, - ))) + fn listen(&self, _backlog: u32) -> litebox_broker_core::Result> { + Ok(SocketOutcome::Completed(())) } fn accept( @@ -821,7 +819,7 @@ mod tests { fn send( &self, - data: &[u8], + data: Vec, _flags: SendFlags, ) -> litebox_broker_core::Result> { Ok(SocketOutcome::Completed(data.len())) @@ -829,7 +827,7 @@ mod tests { fn send_to( &self, - data: &[u8], + data: Vec, _flags: SendFlags, _destination: Option, ) -> litebox_broker_core::Result> { @@ -838,27 +836,25 @@ mod tests { fn receive( &self, - data: &mut [u8], + length: usize, _flags: ReceiveFlags, _peek_offset: u32, _peek_length: u32, - ) -> litebox_broker_core::Result> { - let received = data.len().min(3); - data[..received].copy_from_slice(&[4, 5, 6][..received]); - Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received( - u32::try_from(received).unwrap(), + ) -> litebox_broker_core::Result> { + let received = length.min(3); + Ok(SocketOutcome::Completed(PlatformStreamReceive::Received( + [4, 5, 6][..received].to_vec(), ))) } fn receive_from( &self, - data: &mut [u8], + length: usize, _flags: ReceiveFromFlags, - ) -> litebox_broker_core::Result> { - let received = data.len().min(3); - data[..received].copy_from_slice(&[4, 5, 6][..received]); - Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { - received, + ) -> litebox_broker_core::Result> { + let received = length.min(3); + Ok(SocketOutcome::Completed(PlatformDatagramReceive { + data: [4, 5, 6][..received].to_vec(), datagram_length: 4, source_address: SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49153), })) @@ -1363,7 +1359,7 @@ mod tests { ), BrokerResult::Socket(SocketResponse::Status(SocketStatusResponse { status: SocketConnectionStatus::Connected, - local_address: None, + local_address: Some(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 49152)), pending_error: None, })) ); diff --git a/litebox_broker_platform_linux_userland/src/lib.rs b/litebox_broker_platform_linux_userland/src/lib.rs index 7be7622a8c..027be6e1af 100644 --- a/litebox_broker_platform_linux_userland/src/lib.rs +++ b/litebox_broker_platform_linux_userland/src/lib.rs @@ -11,4 +11,4 @@ mod socket; -pub use socket::LinuxSocketProvider; +pub use socket::{LinuxSocketProvider, SocketPortMapping, SocketPortMappingProtocol}; diff --git a/litebox_broker_platform_linux_userland/src/socket.rs b/litebox_broker_platform_linux_userland/src/socket.rs index a9262c1309..ac761eb6bf 100644 --- a/litebox_broker_platform_linux_userland/src/socket.rs +++ b/litebox_broker_platform_linux_userland/src/socket.rs @@ -3,11 +3,11 @@ //! Broker-owned Linux sockets driven by one epoll reactor. -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::fmt; use std::io::{Error, ErrorKind, Result as IoResult}; use std::mem::size_of; -use std::net::SocketAddrV4; +use std::net::{Ipv4Addr, SocketAddrV4}; use std::os::fd::OwnedFd; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::mpsc::{Receiver, SyncSender, TryRecvError, TrySendError, sync_channel}; @@ -16,16 +16,16 @@ use std::thread::{self, JoinHandle}; use std::time::Duration; use litebox_broker_core::socket::{ - AcceptedPlatformSocket, PlatformConnectError, PlatformSocket, ReceivedPlatformDatagram, - SocketProvider, + AcceptedPlatformSocket, PlatformBindKind, PlatformConnectError, PlatformDatagramReceive, + PlatformSocket, PlatformSocketScope, PlatformStreamReceive, SocketProvider, }; use litebox_broker_core::{BrokerError, Result as BrokerResult, SessionId}; use litebox_broker_protocol::readiness::ReadinessFlags; use litebox_broker_protocol::socket::{ AddressFamily, CreateSocketRequest, IpProtocol, MAX_SOCKET_PEEK_SIZE, MAX_SOCKET_TRANSFER_SIZE, - MAX_UDP_DATAGRAM_SIZE, ReceiveFlags, ReceiveFromFlags, ReceiveSocketResponse, SendFlags, - ShutdownMode, SocketConnectionStatus, SocketError, SocketOutcome, SocketStatusResponse, - SocketType, TcpOptionName, TcpOptionValue, + MAX_UDP_DATAGRAM_SIZE, ReceiveFlags, ReceiveFromFlags, SendFlags, ShutdownMode, + SocketConnectionStatus, SocketError, SocketOutcome, SocketStatusResponse, SocketType, + TcpOptionName, TcpOptionValue, }; use rustix::buffer::spare_capacity; use rustix::event::{EventfdFlags, PollFd, PollFlags, Timespec, epoll, eventfd, poll}; @@ -43,6 +43,15 @@ use litebox_broker_core::readiness::ReadinessRegistration; const WAKE_TOKEN: u64 = 0; const MAX_QUEUED_SOCKET_COMMANDS: usize = 64; const MAX_EPOLL_EVENTS: usize = 64; +const MAX_STALE_PUBLICATION_ITEMS: usize = 64; +const MAX_FILTERED_UDP_DATAGRAMS: usize = 64; +const MAX_UDP_PEERS_PER_SOCKET: usize = 256; +const FIRST_PRIVATE_UDP_PORT: u16 = 16384; +const LAST_PRIVATE_UDP_PORT: u16 = 32767; +const MAX_RETAINED_TRANSLATIONS: usize = 1 << 14; +// Retained source ports prevent stale datagrams from being attributed to later +// sockets, so this lifetime quota is released only when the session closes. +const MAX_PRIVATE_UDP_PORTS_PER_SESSION: usize = 1 << 10; /// Linux-userland socket provider. /// @@ -51,22 +60,114 @@ const MAX_EPOLL_EVENTS: usize = 64; /// immediate nonblocking operation, never for network readiness. pub struct LinuxSocketProvider { reactor: Arc, + port_mappings: Vec, +} + +/// Transport protocol for a host-to-guest socket port mapping. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum SocketPortMappingProtocol { + /// Map a host TCP endpoint to a guest TCP listener. + Tcp, + /// Map a host UDP endpoint to a guest UDP binding. + Udp, +} + +/// Explicit mapping from one host endpoint to a guest-local port. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct SocketPortMapping { + /// Transport protocol used by the mapping. + pub protocol: SocketPortMappingProtocol, + /// Guest-local port visible inside one broker session. + pub guest_port: u16, + /// Host endpoint mapped to the guest port. + pub host_address: SocketAddrV4, +} + +struct PortMappingState { + mapping: SocketPortMapping, + reservation: Option, + claimed_by: Option, } impl LinuxSocketProvider { /// Starts a provider whose reactor tracks at most `max_sockets` resources. pub fn new(max_sockets: usize) -> IoResult { + Self::new_with_port_mappings(max_sockets, &[]) + } + + /// Starts a provider and reserves every mapped host endpoint immediately. + pub fn new_with_port_mappings( + max_sockets: usize, + port_mappings: &[SocketPortMapping], + ) -> IoResult { + for (index, mapping) in port_mappings.iter().enumerate() { + if mapping.guest_port == 0 || mapping.host_address.port() == 0 { + return Err(Error::new( + ErrorKind::InvalidInput, + "mapped guest and host ports must be nonzero", + )); + } + if port_mappings[..index].iter().any(|existing| { + (existing.protocol, existing.guest_port) == (mapping.protocol, mapping.guest_port) + || (existing.protocol, existing.host_address) + == (mapping.protocol, mapping.host_address) + }) { + return Err(Error::new( + ErrorKind::InvalidInput, + "mapped guest ports and host endpoints must be unique per protocol", + )); + } + } + let port_mappings = port_mappings.to_vec(); Ok(Self { - reactor: Arc::new(ReactorClient::start(max_sockets)?), + reactor: Arc::new(ReactorClient::start(max_sockets, port_mappings.clone())?), + port_mappings, }) } } +fn create_port_mapping_reservation( + mapping: SocketPortMapping, + reuse_address: bool, + reuse_port: bool, +) -> core::result::Result { + let (socket_type, protocol) = match mapping.protocol { + SocketPortMappingProtocol::Tcp => (LinuxSocketType::STREAM, ipproto::TCP), + SocketPortMappingProtocol::Udp => (LinuxSocketType::DGRAM, ipproto::UDP), + }; + let socket = socket_with( + LinuxAddressFamily::INET, + socket_type, + LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, + Some(protocol), + )?; + if reuse_address { + sockopt::set_socket_reuseaddr(&socket, true)?; + } + if reuse_port { + sockopt::set_socket_reuseport(&socket, true)?; + } + bind(&socket, &mapping.host_address)?; + Ok(socket) +} + impl SocketProvider for LinuxSocketProvider { + fn reserves_guest_port(&self, request: CreateSocketRequest, port: u16) -> bool { + let protocol = match (request.socket_type, request.protocol) { + (SocketType::Stream, IpProtocol::Tcp) => SocketPortMappingProtocol::Tcp, + (SocketType::Datagram, IpProtocol::Udp) => SocketPortMappingProtocol::Udp, + _ => return false, + }; + self.port_mappings + .iter() + .any(|mapping| mapping.protocol == protocol && mapping.guest_port == port) + } + fn create( &self, - _session_id: SessionId, + session_id: SessionId, request: CreateSocketRequest, + scope: PlatformSocketScope, readiness: ReadinessRegistration, ) -> BrokerResult> { if socket_kind(request).is_none() { @@ -86,7 +187,9 @@ impl SocketProvider for LinuxSocketProvider { }); self.reactor.request(|response| ReactorCommand::Create { id, + session_id, request, + scope, readiness, snapshot, active, @@ -95,7 +198,9 @@ impl SocketProvider for LinuxSocketProvider { Ok(socket) } - fn close_session(&self, _session_id: SessionId) {} + fn close_session(&self, session_id: SessionId) { + self.reactor.close_session(session_id); + } } /// Broker-core-facing handle for a reactor-owned socket. @@ -110,15 +215,20 @@ struct LinuxSocket { } impl PlatformSocket for LinuxSocket { - fn bind(&self, address: SocketAddrV4) -> BrokerResult> { + fn bind( + &self, + address: SocketAddrV4, + kind: PlatformBindKind, + ) -> BrokerResult> { self.reactor.request(|response| ReactorCommand::Bind { id: self.id, address, + kind, response, }) } - fn listen(&self, backlog: u32) -> BrokerResult> { + fn listen(&self, backlog: u32) -> BrokerResult> { self.reactor.request(|response| ReactorCommand::Listen { id: self.id, backlog, @@ -150,7 +260,6 @@ impl PlatformSocket for LinuxSocket { SocketOutcome::Completed(accepted) => { Ok(SocketOutcome::Completed(AcceptedPlatformSocket { socket, - local_address: accepted.local_address, remote_address: accepted.remote_address, })) } @@ -165,33 +274,23 @@ impl PlatformSocket for LinuxSocket { self.reactor.connect(self.id, address) } - fn send(&self, data: &[u8], _flags: SendFlags) -> BrokerResult> { - let mut owned = Vec::new(); - owned - .try_reserve_exact(data.len()) - .map_err(|_| BrokerError::OutOfMemory)?; - owned.extend_from_slice(data); + fn send(&self, data: Vec, _flags: SendFlags) -> BrokerResult> { self.reactor.request(|response| ReactorCommand::Send { id: self.id, - data: owned, + data, response, }) } fn send_to( &self, - data: &[u8], + data: Vec, _flags: SendFlags, destination: Option, ) -> BrokerResult> { - let mut owned = Vec::new(); - owned - .try_reserve_exact(data.len()) - .map_err(|_| BrokerError::OutOfMemory)?; - owned.extend_from_slice(data); self.reactor.request(|response| ReactorCommand::SendTo { id: self.id, - data: owned, + data, destination, response, }) @@ -199,34 +298,28 @@ impl PlatformSocket for LinuxSocket { fn receive( &self, - data: &mut [u8], + length: usize, flags: ReceiveFlags, peek_offset: u32, peek_length: u32, - ) -> BrokerResult> { + ) -> BrokerResult> { let peek_offset = usize::try_from(peek_offset).map_err(|_| BrokerError::UnsupportedOperation)?; let peek_length = usize::try_from(peek_length).map_err(|_| BrokerError::UnsupportedOperation)?; match self.reactor.request(|response| ReactorCommand::Receive { id: self.id, - length: data.len(), + length, flags, peek_offset, peek_length, response, })? { - ReactorReceiveOutcome::Received(received) => { - data[..received.len()].copy_from_slice(&received); - Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received( - received - .len() - .try_into() - .map_err(|_| BrokerError::Internal)?, - ))) - } + ReactorReceiveOutcome::Received(received) => Ok(SocketOutcome::Completed( + PlatformStreamReceive::Received(received), + )), ReactorReceiveOutcome::EndOfStream => { - Ok(SocketOutcome::Completed(ReceiveSocketResponse::EndOfStream)) + Ok(SocketOutcome::Completed(PlatformStreamReceive::EndOfStream)) } ReactorReceiveOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), } @@ -234,14 +327,14 @@ impl PlatformSocket for LinuxSocket { fn receive_from( &self, - data: &mut [u8], + length: usize, flags: ReceiveFromFlags, - ) -> BrokerResult> { + ) -> BrokerResult> { match self .reactor .request(|response| ReactorCommand::ReceiveFrom { id: self.id, - length: data.len(), + length, flags, response, })? { @@ -249,14 +342,11 @@ impl PlatformSocket for LinuxSocket { data: received, datagram_length, source_address, - } => { - data[..received.len()].copy_from_slice(&received); - Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { - received: received.len(), - datagram_length, - source_address, - })) - } + } => Ok(SocketOutcome::Completed(PlatformDatagramReceive { + data: received, + datagram_length, + source_address, + })), ReactorReceiveFromOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), } } @@ -319,9 +409,19 @@ struct ReactorClient { } impl ReactorClient { - fn start(max_sockets: usize) -> IoResult { + fn start(max_sockets: usize, port_mappings: Vec) -> IoResult { let epoll_fd = epoll::create(epoll::CreateFlags::CLOEXEC)?; let wake = Arc::new(eventfd(0, EventfdFlags::CLOEXEC | EventfdFlags::NONBLOCK)?); + let port_mappings = port_mappings + .into_iter() + .map(|mapping| { + Ok(PortMappingState { + mapping, + reservation: Some(create_port_mapping_reservation(mapping, false, false)?), + claimed_by: None, + }) + }) + .collect::>>()?; epoll::add( &epoll_fd, wake.as_ref(), @@ -349,6 +449,10 @@ impl ReactorClient { wake: reactor_wake, commands: receiver, sockets, + sessions: HashMap::new(), + port_mappings, + private_udp_ports: [0; MAX_RETAINED_TRANSLATIONS / u64::BITS as usize], + next_private_udp_port: FIRST_PRIVATE_UDP_PORT, max_sockets, peek_cache: None, events, @@ -451,6 +555,46 @@ impl ReactorClient { let _ = receive.recv(); } + #[cfg(test)] + fn host_address(&self, kind: SocketKind, guest_port: u16) -> Option { + let (response, receive) = sync_channel(1); + self.commands + .send(ReactorCommand::HostAddress { + kind, + guest_port, + response, + }) + .unwrap(); + self.signal().unwrap(); + receive.recv().unwrap() + } + + #[cfg(test)] + fn set_next_private_udp_port(&self, port: u16) { + let (response, receive) = sync_channel(1); + self.commands + .send(ReactorCommand::SetNextPrivateUdpPort { port, response }) + .unwrap(); + self.signal().unwrap(); + receive.recv().unwrap(); + } + + fn close_session(&self, session_id: SessionId) { + let (response, receive) = sync_channel(1); + if self + .commands + .send(ReactorCommand::CloseSession { + session_id, + response, + }) + .is_err() + { + return; + } + let _ = self.signal(); + let _ = receive.recv(); + } + fn signal(&self) -> IoResult<()> { let value = 1_u64.to_ne_bytes(); loop { @@ -498,7 +642,9 @@ impl Drop for ReactorClient { enum ReactorCommand { Create { id: u64, + session_id: SessionId, request: CreateSocketRequest, + scope: PlatformSocketScope, readiness: ReadinessRegistration, snapshot: Arc>, active: Arc, @@ -512,12 +658,13 @@ enum ReactorCommand { Bind { id: u64, address: SocketAddrV4, - response: SyncSender>>, + kind: PlatformBindKind, + response: SyncSender>>, }, Listen { id: u64, backlog: u32, - response: SyncSender>>, + response: SyncSender>>, }, Accept { listener_id: u64, @@ -575,6 +722,21 @@ enum ReactorCommand { id: u64, response: SyncSender<()>, }, + CloseSession { + session_id: SessionId, + response: SyncSender<()>, + }, + #[cfg(test)] + HostAddress { + kind: SocketKind, + guest_port: u16, + response: SyncSender>, + }, + #[cfg(test)] + SetNextPrivateUdpPort { + port: u16, + response: SyncSender<()>, + }, Stop { response: SyncSender<()>, }, @@ -597,7 +759,6 @@ enum ReactorReceiveFromOutcome { } struct AcceptedEndpoints { - local_address: SocketAddrV4, remote_address: SocketAddrV4, } @@ -607,6 +768,10 @@ struct Reactor { wake: Arc, commands: Receiver, sockets: HashMap, + sessions: HashMap, + port_mappings: Vec, + private_udp_ports: [u64; MAX_RETAINED_TRANSLATIONS / u64::BITS as usize], + next_private_udp_port: u16, max_sockets: usize, peek_cache: Option, events: Vec, @@ -619,15 +784,82 @@ struct PeekCache { } /// Reactor-owned descriptor and its broker-facing readiness state. +#[allow(clippy::struct_excessive_bools)] // These are independent cached kernel attributes. struct SocketEntry { socket: OwnedFd, + session_id: SessionId, kind: SocketKind, + scope: PlatformSocketScope, readiness: ReadinessRegistration, snapshot: Arc>, read_shutdown: bool, write_shutdown: bool, peek_waitall_threshold: Option, listening: bool, + guest_local_address: Option, + port_mapping_index: Option, + retain_port_mapping_on_close: bool, + udp_allowed_peers: HashSet, + tcp_no_delay: bool, + tcp_keep_alive: bool, +} + +impl SocketEntry { + fn reserve_udp_peer(&mut self, address: SocketAddrV4) -> BrokerResult { + if self.kind != SocketKind::Udp { + return Err(BrokerError::Internal); + } + if self.port_mapping_index.is_some() { + return Ok(false); + } + if self.udp_allowed_peers.contains(&address) { + return Ok(false); + } + if self.udp_allowed_peers.len() >= MAX_UDP_PEERS_PER_SOCKET { + return Err(BrokerError::ResourceExhausted); + } + self.udp_allowed_peers + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory)?; + Ok(true) + } +} + +struct SessionSocketNamespace { + tcp: HashMap, + udp: HashMap, + // Queued datagrams and pending accepts can outlive the sending socket. + tcp_translations: HashMap<(SocketAddrV4, SocketAddrV4), SocketAddrV4>, + udp_translations: HashMap, + // Prevent a later socket in this session from acquiring a stale datagram's source port. + private_udp_ports: [u64; MAX_RETAINED_TRANSLATIONS / u64::BITS as usize], + private_udp_port_count: usize, + live_sockets: usize, + closing: bool, +} + +impl Default for SessionSocketNamespace { + fn default() -> Self { + Self { + tcp: HashMap::new(), + udp: HashMap::new(), + tcp_translations: HashMap::new(), + udp_translations: HashMap::new(), + private_udp_ports: [0; MAX_RETAINED_TRANSLATIONS / u64::BITS as usize], + private_udp_port_count: 0, + live_sockets: 0, + closing: false, + } + } +} + +#[derive(Clone, Copy)] +struct GuestPortBinding { + socket_id: u64, + guest_address: SocketAddrV4, + host_address: Option, + host_peer_address: Option, + host_mapped: bool, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -660,6 +892,287 @@ impl Default for SocketSnapshot { } } +fn private_udp_port_bit(port: u16) -> (usize, u64) { + let offset = usize::from(port - FIRST_PRIVATE_UDP_PORT); + ( + offset / u64::BITS as usize, + 1_u64 << (offset % u64::BITS as usize), + ) +} + +impl SessionSocketNamespace { + fn bindings(&self, kind: SocketKind) -> &HashMap { + match kind { + SocketKind::Tcp => &self.tcp, + SocketKind::Udp => &self.udp, + } + } + + fn bindings_mut(&mut self, kind: SocketKind) -> &mut HashMap { + match kind { + SocketKind::Tcp => &mut self.tcp, + SocketKind::Udp => &mut self.udp, + } + } + + fn retain_private_udp_port(&mut self, port: u16) { + debug_assert!(self.private_udp_port_count < MAX_PRIVATE_UDP_PORTS_PER_SESSION); + let (word, mask) = private_udp_port_bit(port); + debug_assert_eq!(self.private_udp_ports[word] & mask, 0); + self.private_udp_ports[word] |= mask; + self.private_udp_port_count += 1; + } + + fn has_private_udp_port_capacity(&self) -> bool { + self.private_udp_port_count < MAX_PRIVATE_UDP_PORTS_PER_SESSION + } + + fn insert_binding( + &mut self, + kind: SocketKind, + port: u16, + binding: GuestPortBinding, + ) -> BrokerResult<()> { + let udp_host_address = if kind == SocketKind::Udp { + let host_address = binding.host_address.ok_or(BrokerError::Internal)?; + if let Some(previous) = self.udp_translations.get(&host_address) + && (!previous.host_mapped + || !binding.host_mapped + || previous.guest_address != binding.guest_address) + { + return Err(BrokerError::Internal); + } + Some(host_address) + } else { + None + }; + let tcp_translation = if kind == SocketKind::Tcp { + binding + .host_address + .zip(binding.host_peer_address) + .map(|connection| (connection, binding.guest_address)) + } else { + None + }; + let bindings = self.bindings_mut(kind); + if bindings.insert(port, binding).is_some() { + return Err(BrokerError::Internal); + } + if let Some(host_address) = udp_host_address { + self.udp_translations.insert(host_address, binding); + } + if let Some((connection, guest_address)) = tcp_translation { + self.tcp_translations.insert(connection, guest_address); + } + Ok(()) + } + + fn reserve_binding(&mut self, kind: SocketKind) -> BrokerResult<()> { + let translation_count = match kind { + SocketKind::Tcp => self.tcp_translations.len(), + SocketKind::Udp => self.udp_translations.len(), + }; + if translation_count >= MAX_RETAINED_TRANSLATIONS { + return Err(BrokerError::ResourceExhausted); + } + self.bindings_mut(kind) + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory)?; + match kind { + SocketKind::Tcp => self.tcp_translations.try_reserve(1), + SocketKind::Udp => self.udp_translations.try_reserve(1), + } + .map_err(|_| BrokerError::OutOfMemory) + } + + fn remove_binding(&mut self, kind: SocketKind, port: u16, socket_id: u64) { + let bindings = self.bindings_mut(kind); + if bindings + .get(&port) + .is_some_and(|binding| binding.socket_id == socket_id) + { + bindings.remove(&port); + } + } + + fn guest_binding(&self, kind: SocketKind, address: SocketAddrV4) -> Option { + if !address.ip().is_loopback() { + return None; + } + let binding = self.bindings(kind).get(&address.port())?; + if binding.guest_address.ip().is_unspecified() || binding.guest_address.ip() == address.ip() + { + Some(*binding) + } else { + None + } + } + + fn set_host_address( + &mut self, + kind: SocketKind, + port: u16, + socket_id: u64, + host_address: SocketAddrV4, + host_mapped: bool, + ) -> BrokerResult<()> { + if kind == SocketKind::Udp { + let mut translation = *self + .bindings(kind) + .get(&port) + .filter(|binding| binding.socket_id == socket_id) + .ok_or(BrokerError::Internal)?; + if translation.host_address == Some(host_address) { + let translation = self + .udp_translations + .get_mut(&host_address) + .filter(|binding| binding.socket_id == socket_id) + .ok_or(BrokerError::Internal)?; + translation.host_mapped = host_mapped; + return Ok(()); + } + if let Some(existing) = self.udp_translations.get_mut(&host_address) { + if existing.socket_id != socket_id { + return Err(BrokerError::Internal); + } + existing.guest_address = translation.guest_address; + existing.host_address = Some(host_address); + existing.host_mapped = host_mapped; + return Ok(()); + } + if self.udp_translations.len() >= MAX_RETAINED_TRANSLATIONS { + return Err(BrokerError::ResourceExhausted); + } + self.udp_translations + .try_reserve(1) + .map_err(|_| BrokerError::OutOfMemory)?; + translation.host_address = Some(host_address); + translation.host_mapped = host_mapped; + self.udp_translations.insert(host_address, translation); + return Ok(()); + } + { + let binding = self + .bindings_mut(kind) + .get_mut(&port) + .ok_or(BrokerError::Internal)?; + if binding.socket_id != socket_id { + return Err(BrokerError::Internal); + } + binding.host_address = Some(host_address); + binding.host_mapped = host_mapped; + } + Ok(()) + } + + fn set_guest_address( + &mut self, + kind: SocketKind, + port: u16, + socket_id: u64, + guest_address: SocketAddrV4, + ) -> BrokerResult<()> { + { + let binding = self + .bindings_mut(kind) + .get_mut(&port) + .ok_or(BrokerError::Internal)?; + if binding.socket_id != socket_id { + return Err(BrokerError::Internal); + } + binding.guest_address = guest_address; + } + if kind == SocketKind::Udp { + let mut found = false; + for binding in self + .udp_translations + .values_mut() + .filter(|binding| binding.socket_id == socket_id) + { + binding.guest_address = guest_address; + found = true; + } + if !found { + return Err(BrokerError::Internal); + } + } + Ok(()) + } + + fn set_host_peer_address( + &mut self, + kind: SocketKind, + port: u16, + socket_id: u64, + host_peer_address: SocketAddrV4, + ) -> BrokerResult<()> { + if kind != SocketKind::Tcp { + return Err(BrokerError::Internal); + } + let (host_address, guest_address) = { + let binding = self + .bindings_mut(kind) + .get_mut(&port) + .ok_or(BrokerError::Internal)?; + if binding.socket_id != socket_id { + return Err(BrokerError::Internal); + } + binding.host_peer_address = Some(host_peer_address); + ( + binding.host_address.ok_or(BrokerError::Internal)?, + binding.guest_address, + ) + }; + self.tcp_translations + .insert((host_address, host_peer_address), guest_address); + Ok(()) + } + + fn translate_udp_source(&self, address: SocketAddrV4) -> Option { + self.udp_translations + .get(&address) + .or_else(|| { + if address.ip().is_loopback() { + self.udp_translations + .get(&SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, address.port())) + } else { + None + } + }) + .and_then(|binding| { + binding.host_address.and_then(|host_address| { + (host_address.port() == address.port() + && ((host_address.ip().is_unspecified() && address.ip().is_loopback()) + || host_address.ip() == address.ip())) + .then(|| { + if binding.guest_address.ip().is_unspecified() { + SocketAddrV4::new(*address.ip(), binding.guest_address.port()) + } else { + binding.guest_address + } + }) + }) + }) + } + + fn translate_tcp_peer( + &self, + remote_address: SocketAddrV4, + local_address: SocketAddrV4, + ) -> SocketAddrV4 { + self.tcp_translations + .get(&(remote_address, local_address)) + .or_else(|| { + self.tcp_translations.get(&( + remote_address, + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, local_address.port()), + )) + }) + .copied() + .unwrap_or(remote_address) + } +} + /// Fatal failure that terminates the socket reactor. #[derive(Debug)] enum ReactorFailure { @@ -679,55 +1192,726 @@ impl fmt::Display for ReactorFailure { } impl Reactor { - fn run(&mut self) -> core::result::Result<(), ReactorFailure> { - loop { - let mut events = core::mem::take(&mut self.events); - events.clear(); - match epoll::wait(&self.epoll, spare_capacity(&mut events), None) { - Ok(_) => {} - Err(Errno::INTR) => { - self.events = events; - continue; - } - Err(error) => return Err(ReactorFailure::Io(error)), - } + fn port_mapping_index(&self, kind: SocketKind, guest_port: u16) -> Option { + let protocol = match kind { + SocketKind::Tcp => SocketPortMappingProtocol::Tcp, + SocketKind::Udp => SocketPortMappingProtocol::Udp, + }; + self.port_mappings.iter().position(|state| { + (state.mapping.protocol, state.mapping.guest_port) == (protocol, guest_port) + }) + } - // Apply readiness observed by this wait before commands. A command - // that then reaches EAGAIN records the newer authoritative state. - let mut wake = false; - for event in events.drain(..) { - let id = event.data.u64(); - if id == WAKE_TOKEN { - wake = true; - } else if let Some(socket) = self.sockets.get_mut(&id) { - handle_socket_event(socket, event.flags).map_err(ReactorFailure::Broker)?; + fn claim_port_mapping( + &mut self, + id: u64, + mapping_index: usize, + ) -> BrokerResult> { + let mapping = self + .port_mappings + .get(mapping_index) + .ok_or(BrokerError::Internal)? + .mapping; + if self + .port_mappings + .get(mapping_index) + .is_some_and(|state| state.claimed_by.is_none() && state.reservation.is_none()) + { + match create_port_mapping_reservation( + mapping, + mapping.protocol == SocketPortMappingProtocol::Tcp, + false, + ) { + Ok(reservation) => { + self.port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)? + .reservation = Some(reservation); } - } - self.events = events; - if wake { - self.drain_wake()?; - if self.process_commands() { - return Ok(()); + Err(error) => { + return Ok(SocketOutcome::Failed(socket_operation_error_from_errno( + error, + )?)); } } } - } - - fn drain_wake(&self) -> core::result::Result<(), ReactorFailure> { - let mut value = [0_u8; size_of::()]; - loop { - match read(self.wake.as_ref(), &mut value) { - Ok(length) if length == value.len() => return Ok(()), - Ok(_) => return Err(ReactorFailure::Io(Errno::IO)), - Err(Errno::INTR) => {} - Err(Errno::AGAIN) => return Ok(()), - Err(error) => return Err(ReactorFailure::Io(error)), + let reservation = { + let state = self + .port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)?; + if state.claimed_by.is_some() { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } + let Some(reservation) = state.reservation.take() else { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + }; + reservation + }; + let preparation = match mapping.protocol { + SocketPortMappingProtocol::Tcp => (|| { + sockopt::set_socket_reuseaddr(&reservation, true) + .map_err(broker_error_from_errno)?; + sockopt::set_socket_linger(&reservation, None).map_err(broker_error_from_errno)?; + drain_tcp_listener(&reservation) + })(), + SocketPortMappingProtocol::Udp => (|| { + rustix::net::connect_unspec(&reservation).map_err(broker_error_from_errno)?; + drain_udp_socket(&reservation) + })(), + }; + let stale_state_drained = match preparation { + Ok(drained) => drained, + Err(error) => { + self.port_mappings[mapping_index].reservation = Some(reservation); + return Err(error); } + }; + if !stale_state_drained { + self.port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)? + .reservation = Some(reservation); + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); } + let replace_result = self.replace_socket_descriptor(id, reservation); + let host_address = match replace_result { + Ok(address) => address, + Err((error, reservation)) => { + self.port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)? + .reservation = Some(reservation); + return Err(error); + } + }; + self.port_mappings + .get_mut(mapping_index) + .ok_or(BrokerError::Internal)? + .claimed_by = Some(id); + Ok(SocketOutcome::Completed(host_address)) } - fn process_commands(&mut self) -> bool { - for _ in 0..MAX_QUEUED_SOCKET_COMMANDS { + fn replace_socket_descriptor( + &mut self, + id: u64, + replacement: OwnedFd, + ) -> core::result::Result { + let host_address = match local_socket_address(&replacement) { + Ok(address) => address, + Err(error) => return Err((error, replacement)), + }; + let Some(socket) = self.sockets.get_mut(&id) else { + return Err((BrokerError::Internal, replacement)); + }; + if socket.kind == SocketKind::Tcp + && let Err(error) = + apply_tcp_options(&replacement, socket.tcp_no_delay, socket.tcp_keep_alive) + { + return Err((error, replacement)); + } + let events = if socket.kind == SocketKind::Tcp { + idle_epoll_events() + } else { + active_epoll_events() + }; + if let Err(error) = epoll::add( + &self.epoll, + &replacement, + epoll::EventData::new_u64(id), + events, + ) { + return Err((broker_error_from_errno(error), replacement)); + } + if let Err(error) = epoll::delete(&self.epoll, &socket.socket) { + let _ = epoll::delete(&self.epoll, &replacement); + return Err((broker_error_from_errno(error), replacement)); + } + let old_socket = core::mem::replace(&mut socket.socket, replacement); + drop(old_socket); + Ok(host_address) + } + + fn is_private_host_endpoint( + &self, + kind: SocketKind, + address: SocketAddrV4, + ) -> BrokerResult { + for binding in self + .sessions + .values() + .flat_map(|namespace| namespace.bindings(kind).values()) + { + let Some(host_address) = binding.host_address else { + continue; + }; + if binding.host_mapped || host_address.port() != address.port() { + continue; + } + if host_address.ip() == address.ip() + || (host_address.ip().is_unspecified() + && host_ipv4_address_is_local(*address.ip())?) + { + return Ok(true); + } + } + Ok(false) + } + + fn resolve_guest_destination( + &self, + session_id: SessionId, + kind: SocketKind, + mut address: SocketAddrV4, + ) -> BrokerResult> { + if address.ip().is_unspecified() { + address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, address.port()); + } + if let Some(binding) = self + .sessions + .get(&session_id) + .and_then(|namespace| namespace.guest_binding(kind, address)) + { + return match binding.host_address { + Some(host_address) if host_address.ip().is_unspecified() => Ok( + SocketOutcome::Completed(SocketAddrV4::new(*address.ip(), host_address.port())), + ), + Some(host_address) => Ok(SocketOutcome::Completed(host_address)), + None => Ok(SocketOutcome::Failed(SocketError::ConnectionRefused)), + }; + } + if self.is_private_host_endpoint(kind, address)? { + Ok(SocketOutcome::Failed(SocketError::ConnectionRefused)) + } else { + Ok(SocketOutcome::Completed(address)) + } + } + + fn bind_private_udp_socket( + &mut self, + id: u64, + session_id: SessionId, + ip: Ipv4Addr, + ) -> BrokerResult> { + if !self + .sessions + .get(&session_id) + .ok_or(BrokerError::Internal)? + .has_private_udp_port_capacity() + { + return Err(BrokerError::ResourceExhausted); + } + for _ in 0..MAX_RETAINED_TRANSLATIONS { + let port = self.next_private_udp_port; + self.next_private_udp_port = if port == LAST_PRIVATE_UDP_PORT { + FIRST_PRIVATE_UDP_PORT + } else { + port + 1 + }; + let (word, mask) = private_udp_port_bit(port); + if self.private_udp_ports[word] & mask != 0 { + continue; + } + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + match bind_host_socket(socket, SocketAddrV4::new(ip, port))? { + SocketOutcome::Completed(address) => { + self.sessions + .get_mut(&session_id) + .ok_or(BrokerError::Internal)? + .retain_private_udp_port(port); + self.private_udp_ports[word] |= mask; + return Ok(SocketOutcome::Completed(address)); + } + SocketOutcome::Failed(SocketError::AddressInUse) => {} + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + } + } + Err(BrokerError::ResourceExhausted) + } + + fn remove_session_namespace(&mut self, session_id: SessionId) { + if let Some(namespace) = self.sessions.remove(&session_id) { + for (retained, owned) in self + .private_udp_ports + .iter_mut() + .zip(namespace.private_udp_ports) + { + *retained &= !owned; + } + } + } + + fn bind_socket( + &mut self, + id: u64, + requested_address: SocketAddrV4, + bind_kind: PlatformBindKind, + ) -> BrokerResult> { + let (session_id, kind, already_bound) = { + let socket = self.sockets.get(&id).ok_or(BrokerError::Internal)?; + ( + socket.session_id, + socket.kind, + socket.guest_local_address.is_some(), + ) + }; + if already_bound { + return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); + } + let guest_port = requested_address.port(); + if guest_port == 0 { + return Err(BrokerError::Internal); + } + if self + .sessions + .get(&session_id) + .ok_or(BrokerError::Internal)? + .bindings(kind) + .contains_key(&guest_port) + { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } + self.sessions + .get_mut(&session_id) + .ok_or(BrokerError::Internal)? + .reserve_binding(kind)?; + let guest_address = SocketAddrV4::new(*requested_address.ip(), guest_port); + let port_mapping_index = (bind_kind == PlatformBindKind::Explicit) + .then(|| self.port_mapping_index(kind, guest_port)) + .flatten(); + if let Some(mapping_index) = port_mapping_index + && kind == SocketKind::Udp + { + let host_address = self + .port_mappings + .get(mapping_index) + .ok_or(BrokerError::Internal)? + .mapping + .host_address; + if self + .sessions + .get(&session_id) + .and_then(|namespace| namespace.udp_translations.get(&host_address)) + .is_some_and(|binding| binding.guest_address != guest_address) + { + return Ok(SocketOutcome::Failed(SocketError::AddressInUse)); + } + } + let host_address = match (kind, port_mapping_index) { + (SocketKind::Udp, Some(mapping_index)) => { + match self.claim_port_mapping(id, mapping_index)? { + SocketOutcome::Completed(address) => { + self.sockets + .get_mut(&id) + .ok_or(BrokerError::Internal)? + .port_mapping_index = Some(mapping_index); + Some(address) + } + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + } + } + (SocketKind::Udp, None) => { + let private_ip = match self.sockets.get(&id).ok_or(BrokerError::Internal)?.scope { + PlatformSocketScope::LoopbackOnly => Ipv4Addr::LOCALHOST, + PlatformSocketScope::GeneralIpv4 => Ipv4Addr::UNSPECIFIED, + }; + match self.bind_private_udp_socket(id, session_id, private_ip)? { + SocketOutcome::Completed(address) => Some(address), + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + } + } + (SocketKind::Tcp, _) => None, + }; + self.sessions + .get_mut(&session_id) + .ok_or(BrokerError::Internal)? + .insert_binding( + kind, + guest_port, + GuestPortBinding { + socket_id: id, + guest_address, + host_address, + host_peer_address: None, + host_mapped: kind == SocketKind::Udp && port_mapping_index.is_some(), + }, + )?; + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + socket.guest_local_address = Some(guest_address); + socket.port_mapping_index = port_mapping_index; + socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .local_address = Some(guest_address); + Ok(SocketOutcome::Completed(())) + } + + fn listen_socket(&mut self, id: u64, backlog: u32) -> BrokerResult> { + let (session_id, kind, guest_address) = self + .sockets + .get(&id) + .map(|socket| (socket.session_id, socket.kind, socket.guest_local_address)) + .ok_or(BrokerError::Internal)?; + let guest_address = guest_address.ok_or(BrokerError::Internal)?; + let needs_host_bind = + local_socket_address(&self.sockets.get(&id).ok_or(BrokerError::Internal)?.socket)? + .port() + == 0; + if needs_host_bind { + let port_mapping_index = self + .sockets + .get(&id) + .ok_or(BrokerError::Internal)? + .port_mapping_index; + let (host_address, host_mapped) = if let Some(mapping_index) = port_mapping_index { + match self.claim_port_mapping(id, mapping_index)? { + SocketOutcome::Completed(address) => (address, true), + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + } + } else { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + let address = + match bind_host_socket(socket, SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0))? { + SocketOutcome::Completed(address) => address, + SocketOutcome::Failed(error) => { + return Ok(SocketOutcome::Failed(error)); + } + }; + (address, false) + }; + self.sessions + .get_mut(&session_id) + .ok_or(BrokerError::Internal)? + .set_host_address(kind, guest_address.port(), id, host_address, host_mapped)?; + } + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + match listen_tcp_socket(&self.epoll, id, socket, backlog)? { + SocketOutcome::Completed(_) => Ok(SocketOutcome::Completed(())), + SocketOutcome::Failed(error) => Ok(SocketOutcome::Failed(error)), + } + } + + fn ensure_guest_bound(&mut self, id: u64) -> BrokerResult> { + if self + .sockets + .get(&id) + .ok_or(BrokerError::Internal)? + .guest_local_address + .is_some() + { + return Ok(SocketOutcome::Completed(())); + } + Err(BrokerError::Internal) + } + + fn connect_guest_socket( + &mut self, + id: u64, + guest_address: SocketAddrV4, + ) -> core::result::Result { + match self + .ensure_guest_bound(id) + .map_err(PlatformConnectError::PeerUnchanged)? + { + SocketOutcome::Completed(()) => {} + SocketOutcome::Failed(error) => { + return Ok(SocketConnectionStatus::Failed(error)); + } + } + let (session_id, kind) = self + .sockets + .get(&id) + .map(|socket| (socket.session_id, socket.kind)) + .ok_or(PlatformConnectError::PeerUnchanged(BrokerError::Internal))?; + let network_address = match self + .resolve_guest_destination(session_id, kind, guest_address) + .map_err(PlatformConnectError::PeerUnchanged)? + { + SocketOutcome::Completed(address) => address, + SocketOutcome::Failed(error) => return Ok(SocketConnectionStatus::Failed(error)), + }; + let outcome = { + let socket = self + .sockets + .get_mut(&id) + .ok_or(PlatformConnectError::PeerUnchanged(BrokerError::Internal))?; + let retain_peer = if kind == SocketKind::Udp { + socket + .reserve_udp_peer(network_address) + .map_err(PlatformConnectError::PeerUnchanged)? + } else { + false + }; + let outcome = match kind { + SocketKind::Tcp => connect_tcp_socket(&self.epoll, id, socket, network_address), + SocketKind::Udp => connect_udp_socket(socket, network_address), + }?; + if retain_peer && outcome == SocketConnectionStatus::Connected { + socket.udp_allowed_peers.insert(network_address); + } + Ok(outcome) + }?; + if matches!( + outcome, + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + ) { + let (local_guest_address, host_address, host_mapped) = + { + let socket = self.sockets.get_mut(&id).ok_or( + PlatformConnectError::PeerIndeterminate(BrokerError::Internal), + )?; + let mut local_guest_address = socket.guest_local_address.ok_or( + PlatformConnectError::PeerIndeterminate(BrokerError::Internal), + )?; + if kind == SocketKind::Udp && local_guest_address.ip().is_unspecified() { + local_guest_address = + SocketAddrV4::new(Ipv4Addr::LOCALHOST, local_guest_address.port()); + socket.guest_local_address = Some(local_guest_address); + socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .local_address = Some(local_guest_address); + } + ( + local_guest_address, + local_socket_address(&socket.socket) + .map_err(PlatformConnectError::PeerIndeterminate)?, + socket.port_mapping_index.is_some(), + ) + }; + let namespace = self.sessions.get_mut(&session_id).ok_or( + PlatformConnectError::PeerIndeterminate(BrokerError::Internal), + )?; + namespace + .set_host_address( + kind, + local_guest_address.port(), + id, + host_address, + host_mapped, + ) + .map_err(PlatformConnectError::PeerIndeterminate)?; + namespace + .set_guest_address(kind, local_guest_address.port(), id, local_guest_address) + .map_err(PlatformConnectError::PeerIndeterminate)?; + if kind == SocketKind::Tcp { + namespace + .set_host_peer_address(kind, local_guest_address.port(), id, network_address) + .map_err(PlatformConnectError::PeerIndeterminate)?; + } + } + Ok(outcome) + } + + fn send_guest_datagram( + &mut self, + id: u64, + data: &[u8], + destination: Option, + ) -> BrokerResult> { + match self.ensure_guest_bound(id)? { + SocketOutcome::Completed(()) => {} + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + } + let (session_id, kind) = self + .sockets + .get(&id) + .map(|socket| (socket.session_id, socket.kind)) + .ok_or(BrokerError::Internal)?; + let network_destination = match destination { + Some(address) => match self.resolve_guest_destination(session_id, kind, address)? { + SocketOutcome::Completed(address) => Some(address), + SocketOutcome::Failed(error) => return Ok(SocketOutcome::Failed(error)), + }, + None => None, + }; + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + let retain_peer = match network_destination { + Some(address) => socket.reserve_udp_peer(address)?, + None => false, + }; + let outcome = send_udp(socket, data, network_destination)?; + if retain_peer && matches!(outcome, SocketOutcome::Completed(_)) { + socket + .udp_allowed_peers + .insert(network_destination.ok_or(BrokerError::Internal)?); + } + Ok(outcome) + } + + fn receive_guest_datagram( + &mut self, + id: u64, + length: usize, + flags: ReceiveFromFlags, + ) -> BrokerResult { + let (session_id, host_mapped) = self + .sockets + .get(&id) + .map(|socket| (socket.session_id, socket.port_mapping_index.is_some())) + .ok_or(BrokerError::Internal)?; + for _ in 0..MAX_FILTERED_UDP_DATAGRAMS { + let outcome = { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + receive_udp(socket, length, flags)? + }; + let ReactorReceiveFromOutcome::Received { + data, + datagram_length, + source_address, + } = outcome + else { + return Ok(outcome); + }; + let translated_source = self + .sessions + .get(&session_id) + .and_then(|namespace| namespace.translate_udp_source(source_address)); + let allowed = translated_source.is_some() + || host_mapped + || self + .sockets + .get(&id) + .ok_or(BrokerError::Internal)? + .udp_allowed_peers + .contains(&source_address); + if allowed { + return Ok(ReactorReceiveFromOutcome::Received { + data, + datagram_length, + source_address: translated_source.unwrap_or(source_address), + }); + } + if flags.contains(ReceiveFromFlags::PEEK) { + let discarded = { + let socket = self.sockets.get_mut(&id).ok_or(BrokerError::Internal)?; + receive_udp(socket, 0, ReceiveFromFlags::NONE)? + }; + if let ReactorReceiveFromOutcome::Failed(error) = discarded { + return Ok(ReactorReceiveFromOutcome::Failed(error)); + } + } + } + let socket = self.sockets.get(&id).ok_or(BrokerError::Internal)?; + let readiness = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + socket.readiness.republish(readiness)?; + Err(BrokerError::WouldBlock) + } + + fn remove_socket(&mut self, id: u64) { + let port_mapping = self + .sockets + .get(&id) + .and_then(|socket| socket.port_mapping_index) + .filter(|mapping_index| { + self.port_mappings + .get(*mapping_index) + .is_some_and(|state| state.claimed_by == Some(id)) + }) + .map(|mapping_index| (mapping_index, self.port_mappings[mapping_index].mapping)); + let retain_original = port_mapping.is_some_and(|_| { + self.sockets.get(&id).is_some_and(|socket| { + socket.retain_port_mapping_on_close + && delete_epoll_registration(&self.epoll, &socket.socket) + }) + }); + let Some(socket) = self.sockets.remove(&id) else { + return; + }; + let session_id = socket.session_id; + let remove_namespace = if let Some(namespace) = self.sessions.get_mut(&session_id) { + if let Some(address) = socket.guest_local_address { + namespace.remove_binding(socket.kind, address.port(), id); + } + namespace.live_sockets = namespace + .live_sockets + .checked_sub(1) + .expect("session socket count underflow"); + namespace.closing && namespace.live_sockets == 0 + } else { + false + }; + if remove_namespace { + self.remove_session_namespace(session_id); + } + let replacement_reservation = if retain_original { + let SocketEntry { socket, .. } = socket; + Some(socket) + } else { + drop(socket); + port_mapping.and_then(|(_, mapping)| { + create_port_mapping_reservation( + mapping, + mapping.protocol == SocketPortMappingProtocol::Tcp, + false, + ) + .ok() + }) + }; + if let Some((mapping_index, _)) = port_mapping + && let Some(state) = self.port_mappings.get_mut(mapping_index) + && state.claimed_by == Some(id) + { + state.claimed_by = None; + state.reservation = replacement_reservation; + } + } + + fn run(&mut self) -> core::result::Result<(), ReactorFailure> { + loop { + let mut events = core::mem::take(&mut self.events); + events.clear(); + match epoll::wait(&self.epoll, spare_capacity(&mut events), None) { + Ok(_) => {} + Err(Errno::INTR) => { + self.events = events; + continue; + } + Err(error) => return Err(ReactorFailure::Io(error)), + } + + // Apply readiness observed by this wait before commands. A command + // that then reaches EAGAIN records the newer authoritative state. + let mut wake = false; + for event in events.drain(..) { + let id = event.data.u64(); + if id == WAKE_TOKEN { + wake = true; + } else if let Some(socket) = self.sockets.get_mut(&id) { + handle_socket_event(socket, event.flags).map_err(ReactorFailure::Broker)?; + } + } + self.events = events; + if wake { + self.drain_wake()?; + if self.process_commands() { + return Ok(()); + } + } + } + } + + fn drain_wake(&self) -> core::result::Result<(), ReactorFailure> { + let mut value = [0_u8; size_of::()]; + loop { + match read(self.wake.as_ref(), &mut value) { + Ok(length) if length == value.len() => return Ok(()), + Ok(_) => return Err(ReactorFailure::Io(Errno::IO)), + Err(Errno::INTR) => {} + Err(Errno::AGAIN) => return Ok(()), + Err(error) => return Err(ReactorFailure::Io(error)), + } + } + } + + fn process_commands(&mut self) -> bool { + for _ in 0..MAX_QUEUED_SOCKET_COMMANDS { let command = match self.commands.try_recv() { Ok(command) => command, Err(TryRecvError::Empty) => return false, @@ -736,19 +1920,22 @@ impl Reactor { match command { ReactorCommand::Create { id, + session_id, request, + scope, readiness, snapshot, active, response, } => { - let outcome = self.create_socket(id, request, readiness, snapshot); + let outcome = + self.create_socket(id, session_id, request, scope, readiness, snapshot); let created = outcome.is_ok(); if created { active.store(true, Ordering::Release); } if response.send(outcome).is_err() && created { - self.sockets.remove(&id); + self.remove_socket(id); } } ReactorCommand::Connect { @@ -756,24 +1943,16 @@ impl Reactor { address, response, } => { - let outcome = match self.sockets.get_mut(&id) { - Some(socket) => connect_socket(&self.epoll, id, socket, address), - None => Err(PlatformConnectError::PeerIndeterminate( - BrokerError::Internal, - )), - }; + let outcome = self.connect_guest_socket(id, address); let _ = response.send(outcome); } ReactorCommand::Bind { id, address, + kind, response, } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| bind_socket(socket, address)); + let outcome = self.bind_socket(id, address, kind); let _ = response.send(outcome); } ReactorCommand::Listen { @@ -781,11 +1960,7 @@ impl Reactor { backlog, response, } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| listen_socket(&self.epoll, id, socket, backlog)); + let outcome = self.listen_socket(id, backlog); let _ = response.send(outcome); } ReactorCommand::Accept { @@ -796,7 +1971,8 @@ impl Reactor { active, response, } => { - let outcome = self.accept_socket(listener_id, accepted_id, readiness, snapshot); + let outcome = + self.accept_tcp_socket(listener_id, accepted_id, readiness, snapshot); let accepted = matches!( &outcome, Ok(SocketOutcome::Completed(AcceptedEndpoints { .. })) @@ -805,7 +1981,7 @@ impl Reactor { active.store(true, Ordering::Release); } if response.send(outcome).is_err() && accepted { - self.sockets.remove(&accepted_id); + self.remove_socket(accepted_id); } } ReactorCommand::Send { id, data, response } => { @@ -813,7 +1989,7 @@ impl Reactor { .sockets .get_mut(&id) .ok_or(BrokerError::Internal) - .and_then(|socket| send_socket(socket, &data)); + .and_then(|socket| send_tcp(socket, &data)); let _ = response.send(outcome); } ReactorCommand::SendTo { @@ -822,11 +1998,7 @@ impl Reactor { destination, response, } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| send_to_socket(socket, &data, destination)); + let outcome = self.send_guest_datagram(id, &data, destination); let _ = response.send(outcome); } ReactorCommand::Receive { @@ -838,7 +2010,7 @@ impl Reactor { response, } => { let outcome = match self.sockets.get_mut(&id) { - Some(socket) => receive_socket( + Some(socket) => receive_tcp( socket, &mut self.peek_cache, id, @@ -857,11 +2029,7 @@ impl Reactor { flags, response, } => { - let outcome = self - .sockets - .get_mut(&id) - .ok_or(BrokerError::Internal) - .and_then(|socket| receive_from_socket(socket, length, flags)); + let outcome = self.receive_guest_datagram(id, length, flags); let _ = response.send(outcome); } ReactorCommand::Shutdown { id, mode, response } => { @@ -886,7 +2054,7 @@ impl Reactor { } => { let outcome = self .sockets - .get(&id) + .get_mut(&id) .ok_or(BrokerError::Internal) .and_then(|socket| set_tcp_option(socket, value)); let _ = response.send(outcome); @@ -915,11 +2083,43 @@ impl Reactor { { self.peek_cache = None; } - self.sockets.remove(&id); + self.remove_socket(id); let _ = response.send(()); } - ReactorCommand::Stop { response } => { + ReactorCommand::CloseSession { + session_id, + response, + } => { + if let Some(namespace) = self.sessions.get_mut(&session_id) { + namespace.closing = true; + if namespace.live_sockets == 0 { + self.remove_session_namespace(session_id); + } + } + let _ = response.send(()); + } + #[cfg(test)] + ReactorCommand::HostAddress { + kind, + guest_port, + response, + } => { + let host_address = self.sessions.values().find_map(|namespace| { + namespace + .bindings(kind) + .get(&guest_port) + .and_then(|binding| binding.host_address) + }); + let _ = response.send(host_address); + } + #[cfg(test)] + ReactorCommand::SetNextPrivateUdpPort { port, response } => { + self.next_private_udp_port = port; + let _ = response.send(()); + } + ReactorCommand::Stop { response } => { self.sockets.clear(); + self.sessions.clear(); let _ = response.send(()); return true; } @@ -931,7 +2131,9 @@ impl Reactor { fn create_socket( &mut self, id: u64, + session_id: SessionId, request: CreateSocketRequest, + scope: PlatformSocketScope, readiness: ReadinessRegistration, snapshot: Arc>, ) -> BrokerResult<()> { @@ -942,6 +2144,10 @@ impl Reactor { if self.sockets.contains_key(&id) { return Err(BrokerError::Internal); } + let namespace = self.sessions.entry(session_id).or_default(); + if namespace.closing { + return Err(BrokerError::UnknownObject); + } let (linux_type, protocol, epoll_events, initial_readiness) = match kind { SocketKind::Tcp => ( LinuxSocketType::STREAM, @@ -981,19 +2187,31 @@ impl Reactor { id, SocketEntry { socket, + session_id, kind, + scope, readiness, snapshot, read_shutdown: false, write_shutdown: false, peek_waitall_threshold: None, listening: false, + guest_local_address: None, + port_mapping_index: None, + retain_port_mapping_on_close: true, + udp_allowed_peers: HashSet::new(), + tcp_no_delay: false, + tcp_keep_alive: false, }, ); + namespace.live_sockets = namespace + .live_sockets + .checked_add(1) + .ok_or(BrokerError::ResourceExhausted)?; Ok(()) } - fn accept_socket( + fn accept_tcp_socket( &mut self, listener_id: u64, accepted_id: u64, @@ -1013,6 +2231,11 @@ impl Reactor { if listener.kind != SocketKind::Tcp || !listener.listening { return Ok(SocketOutcome::Failed(SocketError::NotConnected)); } + let listener_session_id = listener.session_id; + let listener_scope = listener.scope; + let listener_tcp_no_delay = listener.tcp_no_delay; + let listener_tcp_keep_alive = listener.tcp_keep_alive; + let local_address = listener.guest_local_address.ok_or(BrokerError::Internal)?; let (socket, remote_address) = loop { match acceptfrom_with( &listener.socket, @@ -1049,7 +2272,13 @@ impl Reactor { } let remote_address = SocketAddrV4::try_from(remote_address.ok_or(BrokerError::Internal)?) .map_err(|_| BrokerError::Internal)?; - let local_address = local_socket_address(&socket)?; + let host_local_address = local_socket_address(&socket)?; + let remote_address = self + .sessions + .get(&listener_session_id) + .map_or(remote_address, |namespace| { + namespace.translate_tcp_peer(remote_address, host_local_address) + }); epoll::add( &self.epoll, &socket, @@ -1070,17 +2299,32 @@ impl Reactor { accepted_id, SocketEntry { socket, + session_id: listener_session_id, kind: SocketKind::Tcp, + scope: listener_scope, readiness, snapshot, read_shutdown: false, write_shutdown: false, peek_waitall_threshold: None, listening: false, + guest_local_address: Some(local_address), + port_mapping_index: None, + retain_port_mapping_on_close: true, + udp_allowed_peers: HashSet::new(), + tcp_no_delay: listener_tcp_no_delay, + tcp_keep_alive: listener_tcp_keep_alive, }, ); + let namespace = self + .sessions + .get_mut(&listener_session_id) + .ok_or(BrokerError::Internal)?; + namespace.live_sockets = namespace + .live_sockets + .checked_add(1) + .ok_or(BrokerError::ResourceExhausted)?; Ok(SocketOutcome::Completed(AcceptedEndpoints { - local_address, remote_address, })) } @@ -1095,23 +2339,21 @@ impl Reactor { snapshot.readiness = ReadinessFlags::ERROR; drop(snapshot); // The readiness path may itself be why the reactor is failing. The - // cached terminal snapshot remains authoritative if publication is + // cached terminal snapshot remains authoritative if a port mapping is // no longer available. let _ = socket.readiness.publish(ReadinessFlags::ERROR); } self.sockets.clear(); + self.sessions.clear(); } } -fn connect_socket( +fn connect_tcp_socket( epoll_fd: &OwnedFd, id: u64, socket: &mut SocketEntry, address: SocketAddrV4, ) -> core::result::Result { - if socket.kind == SocketKind::Udp { - return connect_datagram_socket(socket, address); - } if let Err(error) = epoll::modify( epoll_fd, &socket.socket, @@ -1148,13 +2390,11 @@ fn connect_socket( }; let readiness = match status { SocketConnectionStatus::Connected | SocketConnectionStatus::Connecting => { - let local_address = local_socket_address(&socket.socket) - .map_err(PlatformConnectError::PeerIndeterminate)?; - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(local_address); + if socket.guest_local_address.is_none() { + return Err(PlatformConnectError::PeerIndeterminate( + BrokerError::Internal, + )); + } if status == SocketConnectionStatus::Connected { ReadinessFlags::WRITE } else { @@ -1174,20 +2414,18 @@ fn connect_socket( Ok(status) } -fn connect_datagram_socket( +fn connect_udp_socket( socket: &mut SocketEntry, address: SocketAddrV4, ) -> core::result::Result { loop { match connect(&socket.socket, &address) { Ok(()) | Err(Errno::ISCONN) => { - let local_address = local_socket_address(&socket.socket) - .map_err(PlatformConnectError::PeerIndeterminate)?; - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(local_address); + if socket.guest_local_address.is_none() { + return Err(PlatformConnectError::PeerIndeterminate( + BrokerError::Internal, + )); + } let readiness = socket .snapshot .lock() @@ -1213,7 +2451,7 @@ fn connect_datagram_socket( } } -fn bind_socket( +fn bind_host_socket( socket: &mut SocketEntry, address: SocketAddrV4, ) -> BrokerResult> { @@ -1221,11 +2459,6 @@ fn bind_socket( match bind(&socket.socket, &address) { Ok(()) => { let local_address = local_socket_address(&socket.socket)?; - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(local_address); return Ok(SocketOutcome::Completed(local_address)); } Err(Errno::INTR) => {} @@ -1238,7 +2471,59 @@ fn bind_socket( } } -fn listen_socket( +fn delete_epoll_registration(epoll_fd: &OwnedFd, socket: &OwnedFd) -> bool { + loop { + match epoll::delete(epoll_fd, socket) { + Ok(()) | Err(Errno::NOENT) => return true, + Err(Errno::INTR) => {} + Err(_) => return false, + } + } +} + +fn host_ipv4_address_is_local(address: Ipv4Addr) -> BrokerResult { + let socket = socket_with( + LinuxAddressFamily::INET, + LinuxSocketType::DGRAM, + LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, + Some(ipproto::UDP), + ) + .map_err(broker_error_from_errno)?; + match bind(&socket, &SocketAddrV4::new(address, 0)) { + Ok(()) => Ok(true), + Err(Errno::ADDRNOTAVAIL) => Ok(false), + Err(error) => Err(broker_error_from_errno(error)), + } +} + +fn drain_udp_socket(socket: &OwnedFd) -> BrokerResult { + let mut byte = [0_u8; 1]; + for _ in 0..MAX_STALE_PUBLICATION_ITEMS { + match recvfrom(socket, &mut byte, LinuxRecvFlags::TRUNC) { + Ok(_) | Err(Errno::INTR) => {} + Err(Errno::AGAIN) => return Ok(true), + Err(error) => return Err(broker_error_from_errno(error)), + } + } + Ok(false) +} + +fn drain_tcp_listener(socket: &OwnedFd) -> BrokerResult { + for _ in 0..MAX_STALE_PUBLICATION_ITEMS { + match acceptfrom_with( + socket, + LinuxSocketFlags::CLOEXEC | LinuxSocketFlags::NONBLOCK, + ) { + Ok((accepted, _)) => drop(accepted), + Err(Errno::INTR) => {} + Err(Errno::AGAIN | Errno::INVAL) => return Ok(true), + Err(error) => return Err(broker_error_from_errno(error)), + } + } + Ok(false) +} + +fn listen_tcp_socket( epoll_fd: &OwnedFd, id: u64, socket: &mut SocketEntry, @@ -1248,7 +2533,6 @@ fn listen_socket( return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); } let backlog = i32::try_from(backlog).map_err(|_| BrokerError::UnsupportedOperation)?; - let local_address = local_socket_address(&socket.socket)?; let was_listening = socket.listening; if !was_listening { epoll::modify( @@ -1280,18 +2564,18 @@ fn listen_socket( } } socket.listening = true; + let local_address = local_socket_address(&socket.socket)?; let mut snapshot = socket .snapshot .lock() .expect("Linux socket snapshot mutex poisoned"); - snapshot.local_address = Some(local_address); if !was_listening { snapshot.readiness = ReadinessFlags::default(); } Ok(SocketOutcome::Completed(local_address)) } -fn send_socket(socket: &mut SocketEntry, data: &[u8]) -> BrokerResult> { +fn send_tcp(socket: &mut SocketEntry, data: &[u8]) -> BrokerResult> { if socket.kind != SocketKind::Tcp { return Ok(SocketOutcome::Failed(SocketError::InvalidArgument)); } @@ -1315,7 +2599,7 @@ fn send_socket(socket: &mut SocketEntry, data: &[u8]) -> BrokerResult, @@ -1356,7 +2640,7 @@ fn send_to_socket( } } -fn receive_socket( +fn receive_tcp( socket: &mut SocketEntry, peek_cache: &mut Option, socket_id: u64, @@ -1379,7 +2663,7 @@ fn receive_socket( { *peek_cache = None; } - return receive_socket_once(socket, zeroed_vec(length)?, LinuxRecvFlags::empty()); + return receive_tcp_once(socket, zeroed_vec(length)?, LinuxRecvFlags::empty()); } let peek_end = peek_offset @@ -1428,7 +2712,7 @@ fn receive_socket( } else { LinuxRecvFlags::PEEK }; - match receive_socket_once(socket, zeroed_vec(peek_length)?, flags)? { + match receive_tcp_once(socket, zeroed_vec(peek_length)?, flags)? { ReactorReceiveOutcome::Received(data) => { *peek_cache = Some(PeekCache { socket_id, @@ -1468,7 +2752,7 @@ fn receive_socket( Ok(ReactorReceiveOutcome::Received(data)) } -fn receive_from_socket( +fn receive_udp( socket: &mut SocketEntry, length: usize, flags: ReceiveFromFlags, @@ -1524,7 +2808,7 @@ fn zeroed_vec(length: usize) -> BrokerResult> { Ok(data) } -fn receive_socket_once( +fn receive_tcp_once( socket: &mut SocketEntry, mut data: Vec, flags: LinuxRecvFlags, @@ -1571,16 +2855,28 @@ fn receive_socket_once( } } -fn set_tcp_option(socket: &SocketEntry, value: TcpOptionValue) -> BrokerResult<()> { +fn set_tcp_option(socket: &mut SocketEntry, value: TcpOptionValue) -> BrokerResult<()> { if socket.kind != SocketKind::Tcp { return Err(BrokerError::UnsupportedOperation); } match value { - TcpOptionValue::NoDelay(value) => sockopt::set_tcp_nodelay(&socket.socket, value), - TcpOptionValue::KeepAlive(value) => sockopt::set_socket_keepalive(&socket.socket, value), + TcpOptionValue::NoDelay(value) => { + sockopt::set_tcp_nodelay(&socket.socket, value).map_err(broker_error_from_errno)?; + socket.tcp_no_delay = value; + } + TcpOptionValue::KeepAlive(value) => { + sockopt::set_socket_keepalive(&socket.socket, value) + .map_err(broker_error_from_errno)?; + socket.tcp_keep_alive = value; + } _ => return Err(BrokerError::UnsupportedOperation), } - .map_err(broker_error_from_errno) + Ok(()) +} + +fn apply_tcp_options(socket: &OwnedFd, no_delay: bool, keep_alive: bool) -> BrokerResult<()> { + sockopt::set_tcp_nodelay(socket, no_delay).map_err(broker_error_from_errno)?; + sockopt::set_socket_keepalive(socket, keep_alive).map_err(broker_error_from_errno) } fn get_tcp_option(socket: &SocketEntry, name: TcpOptionName) -> BrokerResult { @@ -1651,51 +2947,53 @@ fn shutdown_socket( ), _ => return Err(BrokerError::UnsupportedOperation), }; - loop { - match shutdown(&socket.socket, mode) { - Ok(()) => {} - Err(Errno::INTR) => continue, - // Linux applies directional shutdown to unconnected UDP sockets even - // though it reports ENOTCONN. - Err(Errno::NOTCONN) if socket.kind == SocketKind::Udp => {} - Err(Errno::NOTCONN) => { - return Ok(SocketOutcome::Failed(SocketError::NotConnected)); - } - Err(error) => { - let error = socket_operation_error_from_errno(error)?; - return Ok(SocketOutcome::Failed(error)); + if socket.kind != SocketKind::Udp { + loop { + match shutdown(&socket.socket, mode) { + Ok(()) => break, + Err(Errno::INTR) => {} + Err(Errno::NOTCONN) => { + return Ok(SocketOutcome::Failed(SocketError::NotConnected)); + } + Err(error) => { + let error = socket_operation_error_from_errno(error)?; + return Ok(SocketOutcome::Failed(error)); + } } } - if stop_listening { - socket.listening = false; - socket.read_shutdown = true; - socket.peek_waitall_threshold = None; - update_snapshot( - socket, - Some(SocketConnectionStatus::Failed(SocketError::NotConnected)), - ReadinessFlags::WRITE | ReadinessFlags::HANGUP, - )?; - return Ok(SocketOutcome::Completed(())); - } - socket.read_shutdown |= shuts_down_read; - socket.write_shutdown |= shuts_down_write; - let republish_readiness = shuts_down_read && socket.peek_waitall_threshold.take().is_some(); - if clear.0 != 0 { - clear_readiness(socket, clear)?; - } - if add.0 != 0 { - add_readiness(socket, add)?; - } - if republish_readiness { - let readiness = socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .readiness; - socket.readiness.republish(readiness)?; + } + if stop_listening { + socket.listening = false; + if socket.port_mapping_index.is_some() { + socket.retain_port_mapping_on_close = false; } + socket.read_shutdown = true; + socket.peek_waitall_threshold = None; + update_snapshot( + socket, + Some(SocketConnectionStatus::Failed(SocketError::NotConnected)), + ReadinessFlags::WRITE | ReadinessFlags::HANGUP, + )?; return Ok(SocketOutcome::Completed(())); } + socket.read_shutdown |= shuts_down_read; + socket.write_shutdown |= shuts_down_write; + let republish_readiness = shuts_down_read && socket.peek_waitall_threshold.take().is_some(); + if clear.0 != 0 { + clear_readiness(socket, clear)?; + } + if add.0 != 0 { + add_readiness(socket, add)?; + } + if republish_readiness { + let readiness = socket + .snapshot + .lock() + .expect("Linux socket snapshot mutex poisoned") + .readiness; + socket.readiness.republish(readiness)?; + } + Ok(SocketOutcome::Completed(())) } fn handle_socket_event(socket: &mut SocketEntry, events: epoll::EventFlags) -> BrokerResult<()> { @@ -1751,15 +3049,7 @@ fn handle_socket_event(socket: &mut SocketEntry, events: epoll::EventFlags) -> B fn complete_connect(socket: &mut SocketEntry, events: epoll::EventFlags) -> BrokerResult<()> { let status = match sockopt::socket_error(&socket.socket) { Ok(Ok(())) => match getpeername(&socket.socket) { - Ok(Some(_)) => { - let local_address = local_socket_address(&socket.socket)?; - socket - .snapshot - .lock() - .expect("Linux socket snapshot mutex poisoned") - .local_address = Some(local_address); - SocketConnectionStatus::Connected - } + Ok(Some(_)) => SocketConnectionStatus::Connected, Ok(None) | Err(Errno::NOTCONN) => SocketConnectionStatus::Connecting, Err(error) => SocketConnectionStatus::Failed(socket_error_from_errno(error)), }, @@ -2016,81 +3306,1181 @@ const fn socket_error_from_errno(error: Errno) -> SocketError { } } -const fn socket_operation_error_from_errno(error: Errno) -> BrokerResult { - match broker_resource_error_from_errno(error) { - Some(error) => Err(error), - None => Ok(socket_error_from_errno(error)), - } -} +const fn socket_operation_error_from_errno(error: Errno) -> BrokerResult { + match broker_resource_error_from_errno(error) { + Some(error) => Err(error), + None => Ok(socket_error_from_errno(error)), + } +} + +const fn broker_error_from_errno(error: Errno) -> BrokerError { + match broker_resource_error_from_errno(error) { + Some(error) => error, + None => BrokerError::Internal, + } +} + +const fn broker_resource_error_from_errno(error: Errno) -> Option { + match error { + Errno::NOMEM => Some(BrokerError::OutOfMemory), + Errno::MFILE | Errno::NFILE | Errno::NOBUFS | Errno::NOSPC => { + Some(BrokerError::ResourceExhausted) + } + _ => None, + } +} + +#[cfg(test)] +mod tests { + use std::io::{Read as _, Write as _}; + use std::net::Ipv4Addr; + use std::net::{Shutdown, TcpListener, TcpStream, UdpSocket}; + use std::sync::mpsc::{Receiver, Sender, channel}; + use std::time::{Duration, Instant}; + + use super::*; + use litebox_broker_core::readiness::ReadinessSink; + use litebox_broker_core::socket::ReceivedPlatformDatagram; + use litebox_broker_core::{ + BrokerCore, BrokerCoreLimits, CallerCredential, DestinationPortRange, DestinationRule, + Ipv4Cidr, ObjectRights, PolicyEngine, SocketPolicy, + }; + use litebox_broker_protocol::ObjectHandle; + use litebox_broker_protocol::socket::{Ipv4Address, Port, ReceiveSocketResponse}; + + const TEST_TIMEOUT: Duration = Duration::from_secs(5); + const FIRST_GUEST_EPHEMERAL_PORT: u16 = 49152; + + #[test] + fn cached_socket_error_precedes_a_new_kernel_error() { + assert_eq!( + shift_pending_error( + Some(SocketError::ConnectionRefused), + Some(SocketError::NetworkUnreachable), + ), + ( + Some(SocketError::ConnectionRefused), + Some(SocketError::NetworkUnreachable), + ) + ); + assert_eq!( + shift_pending_error(None, Some(SocketError::NetworkUnreachable)), + (Some(SocketError::NetworkUnreachable), None) + ); + } + + #[test] + fn synchronous_errors_do_not_consume_tcp_connect_status() { + assert!(!can_consume_synchronous_error( + SocketKind::Tcp, + SocketConnectionStatus::Connecting, + )); + assert!(can_consume_synchronous_error( + SocketKind::Tcp, + SocketConnectionStatus::Connected, + )); + assert!(can_consume_synchronous_error( + SocketKind::Udp, + SocketConnectionStatus::Unconnected, + )); + } + + #[test] + fn wildcard_bindings_translate_only_loopback_peers() { + let mut namespace = SessionSocketNamespace::default(); + namespace + .insert_binding( + SocketKind::Udp, + 53, + GuestPortBinding { + socket_id: 1, + guest_address: SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53), + host_address: Some(SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 8053)), + host_peer_address: None, + host_mapped: true, + }, + ) + .unwrap(); + + assert_eq!( + namespace.translate_udp_source(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 8053)), + Some(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53)) + ); + let external_peer = SocketAddrV4::new(Ipv4Addr::new(192, 0, 2, 1), 8053); + assert_eq!(namespace.translate_udp_source(external_peer), None); + } + + #[test] + fn connected_udp_adds_a_translation_without_losing_its_wildcard_binding() { + let mut namespace = SessionSocketNamespace::default(); + let guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53); + let wildcard_address = SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 8053); + let connected_address = SocketAddrV4::new(Ipv4Addr::new(192, 0, 2, 2), 8053); + namespace + .insert_binding( + SocketKind::Udp, + guest_address.port(), + GuestPortBinding { + socket_id: 1, + guest_address, + host_address: Some(wildcard_address), + host_peer_address: None, + host_mapped: false, + }, + ) + .unwrap(); + + namespace + .set_host_address( + SocketKind::Udp, + guest_address.port(), + 1, + connected_address, + false, + ) + .unwrap(); + + assert_eq!( + namespace + .bindings(SocketKind::Udp) + .get(&guest_address.port()) + .and_then(|binding| binding.host_address), + Some(wildcard_address) + ); + assert_eq!( + namespace.translate_udp_source(connected_address), + Some(guest_address) + ); + assert_eq!( + namespace.translate_udp_source(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 8053)), + Some(guest_address) + ); + } + + #[test] + fn private_udp_retention_is_bounded_per_session() { + let mut namespace = SessionSocketNamespace::default(); + for offset in 0..MAX_PRIVATE_UDP_PORTS_PER_SESSION { + namespace + .retain_private_udp_port(FIRST_PRIVATE_UDP_PORT + u16::try_from(offset).unwrap()); + } + + assert!(!namespace.has_private_udp_port_capacity()); + } + + #[test] + fn tcp_peer_translation_uses_the_complete_connection_tuple() { + let mut namespace = SessionSocketNamespace::default(); + let shared_host_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 40000); + for (guest_port, host_peer_port) in [(1000, 5000), (1001, 5001)] { + namespace + .insert_binding( + SocketKind::Tcp, + guest_port, + GuestPortBinding { + socket_id: u64::from(guest_port), + guest_address: SocketAddrV4::new(Ipv4Addr::LOCALHOST, guest_port), + host_address: Some(shared_host_address), + host_peer_address: Some(SocketAddrV4::new( + Ipv4Addr::LOCALHOST, + host_peer_port, + )), + host_mapped: false, + }, + ) + .unwrap(); + } + + assert_eq!( + namespace.translate_tcp_peer( + shared_host_address, + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 5001), + ), + SocketAddrV4::new(Ipv4Addr::LOCALHOST, 1001) + ); + } + + #[test] + fn port_mapping_reservation_is_close_on_exec() { + let host_reservation = UdpSocket::bind("127.0.0.1:0").unwrap(); + let host_address = socket_address_v4(host_reservation.local_addr().unwrap()); + drop(host_reservation); + let retained = create_port_mapping_reservation( + SocketPortMapping { + protocol: SocketPortMappingProtocol::Udp, + guest_port: 53, + host_address, + }, + false, + false, + ) + .unwrap(); + + assert!( + rustix::io::fcntl_getfd(&retained) + .unwrap() + .contains(rustix::io::FdFlags::CLOEXEC) + ); + } + + #[test] + fn unavailable_publish_endpoint_rejects_provider_startup() { + let occupied = TcpListener::bind("127.0.0.1:0").unwrap(); + let host_address = socket_address_v4(occupied.local_addr().unwrap()); + let error = LinuxSocketProvider::new_with_port_mappings( + 1, + &[SocketPortMapping { + protocol: SocketPortMappingProtocol::Tcp, + guest_port: 80, + host_address, + }], + ) + .err() + .unwrap(); + + assert_eq!(error.kind(), ErrorKind::AddrInUse); + } + + #[test] + fn udp_port_mappings_can_share_a_host_port_on_distinct_addresses() { + let first_reservation = UdpSocket::bind("127.0.0.1:0").unwrap(); + let port = first_reservation.local_addr().unwrap().port(); + let second_reservation = + UdpSocket::bind(SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 2), port)).unwrap(); + drop((first_reservation, second_reservation)); + let provider = Arc::new( + LinuxSocketProvider::new_with_port_mappings( + 2, + &[ + SocketPortMapping { + protocol: SocketPortMappingProtocol::Udp, + guest_port: 53, + host_address: SocketAddrV4::new(Ipv4Addr::LOCALHOST, port), + }, + SocketPortMapping { + protocol: SocketPortMappingProtocol::Udp, + guest_port: 54, + host_address: SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 2), port), + }, + ], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(2, 0, 2, 2), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + + for guest_port in [53, 54] { + let handle = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + let guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, guest_port); + assert_eq!( + litebox_broker_core::socket::bind(&session, handle, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + } + } + + #[test] + fn guest_tcp_ports_are_session_scoped_and_do_not_bind_host_ports() { + let occupied_host_listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let occupied_host_address = socket_address_v4(occupied_host_listener.local_addr().unwrap()); + let provider = Arc::new(LinuxSocketProvider::new(4).unwrap()); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(8, 0, 4, 4), + provider, + ) + .unwrap(); + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let first = create_socket(&first_session, readiness.clone()); + let second = create_socket(&second_session, readiness.clone()); + let occupied_host_port = create_socket(&first_session, readiness); + let guest_port_80 = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + + assert_eq!( + litebox_broker_core::socket::bind(&first_session, first, guest_port_80), + Ok(SocketOutcome::Completed(guest_port_80)) + ); + assert_eq!( + litebox_broker_core::socket::bind(&second_session, second, guest_port_80), + Ok(SocketOutcome::Completed(guest_port_80)) + ); + assert_eq!( + litebox_broker_core::socket::bind( + &first_session, + occupied_host_port, + occupied_host_address, + ), + Ok(SocketOutcome::Completed(occupied_host_address)) + ); + } + + #[test] + fn private_backend_endpoints_are_not_guest_destinations() { + let provider = Arc::new(LinuxSocketProvider::new(4).unwrap()); + let all_destinations = DestinationRule::new( + CallerCredential::Unauthenticated, + Ipv4Cidr::new(Ipv4Address([0, 0, 0, 0]), 0).unwrap(), + DestinationPortRange::new(Port(1), Port(u16::MAX)).unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()).with_socket_policy( + SocketPolicy::from_tcp_udp_destination_rules( + &[all_destinations], + &[all_destinations], + ) + .unwrap(), + ), + BrokerCoreLimits::new_with_all_limits(8, 0, 4, 4), + provider.clone(), + ) + .unwrap(); + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + + let listener = create_socket(&first_session, readiness.clone()); + let guest_listener_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + assert_eq!( + litebox_broker_core::socket::bind(&first_session, listener, guest_listener_address,), + Ok(SocketOutcome::Completed(guest_listener_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&first_session, listener, 1), + Ok(SocketOutcome::Completed(guest_listener_address)) + ); + let private_tcp_address = provider + .reactor + .host_address(SocketKind::Tcp, guest_listener_address.port()) + .unwrap(); + let tcp_client = create_socket(&second_session, readiness.clone()); + let private_tcp_alias = + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, private_tcp_address.port()); + assert_eq!( + litebox_broker_core::socket::connect(&second_session, tcp_client, private_tcp_alias,), + Ok(SocketOutcome::Completed(SocketConnectionStatus::Failed( + SocketError::ConnectionRefused, + ))) + ); + + let udp_server = litebox_broker_core::socket::create( + &first_session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + let guest_udp_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53); + assert_eq!( + litebox_broker_core::socket::bind(&first_session, udp_server, guest_udp_address), + Ok(SocketOutcome::Completed(guest_udp_address)) + ); + let private_udp_address = provider + .reactor + .host_address(SocketKind::Udp, guest_udp_address.port()) + .unwrap(); + let udp_client = litebox_broker_core::socket::create( + &second_session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + let private_udp_alias = + SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, private_udp_address.port()); + assert_eq!( + litebox_broker_core::socket::send_to( + &second_session, + udp_client, + b"x", + SendFlags::NONE, + Some(private_udp_alias), + ), + Ok(SocketOutcome::Failed(SocketError::ConnectionRefused)) + ); + } + + #[test] + fn authorized_nonloopback_udp_uses_a_general_backend_binding() { + let route = UdpSocket::bind("0.0.0.0:0").unwrap(); + route.connect("192.0.2.1:9").unwrap(); + let local_ip = match route.local_addr().unwrap().ip() { + std::net::IpAddr::V4(address) if !address.is_loopback() => address, + _ => return, + }; + let native = UdpSocket::bind(SocketAddrV4::new(local_ip, 0)).unwrap(); + native.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); + let native_address = socket_address_v4(native.local_addr().unwrap()); + let rule = DestinationRule::new( + CallerCredential::Unauthenticated, + Ipv4Cidr::new(Ipv4Address(local_ip.octets()), 32).unwrap(), + DestinationPortRange::new(Port(native_address.port()), Port(native_address.port())) + .unwrap(), + ); + let provider = Arc::new(LinuxSocketProvider::new(1).unwrap()); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::from_udp_destination_rules(&[rule]).unwrap()), + BrokerCoreLimits::new_with_all_limits(2, 0, 1, 1), + provider.clone(), + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, _retirements) = channel(); + let handle = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + Arc::new(TestReadinessSink { published, retired }), + ) + .unwrap(); + let guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53); + assert_eq!( + litebox_broker_core::socket::bind(&session, handle, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + let private_address = provider + .reactor + .host_address(SocketKind::Udp, guest_address.port()) + .unwrap(); + native + .send_to( + b"blocked", + SocketAddrV4::new(local_ip, private_address.port()), + ) + .unwrap(); + wait_for_readiness(&publications, handle, ReadinessFlags::READ); + assert_eq!( + litebox_broker_core::socket::receive_from( + &session, + handle, + &mut [0_u8; 7], + ReceiveFromFlags::PEEK, + ), + Err(BrokerError::WouldBlock) + ); + let (published_handle, published_readiness) = + publications.recv_timeout(TEST_TIMEOUT).unwrap(); + assert_eq!(published_handle, handle); + assert!(!published_readiness.contains(ReadinessFlags::READ)); + + assert_eq!( + litebox_broker_core::socket::send_to( + &session, + handle, + b"routed", + SendFlags::NONE, + Some(native_address), + ), + Ok(SocketOutcome::Completed(6)) + ); + let mut data = [0_u8; 6]; + let (received, source) = native.recv_from(&mut data).unwrap(); + assert_eq!(received, data.len()); + assert_eq!(data, *b"routed"); + + let unsolicited = UdpSocket::bind(SocketAddrV4::new(local_ip, 0)).unwrap(); + for _ in 0..=MAX_FILTERED_UDP_DATAGRAMS { + unsolicited + .send_to( + b"blocked", + SocketAddrV4::new(local_ip, private_address.port()), + ) + .unwrap(); + } + native.send_to(b"reply", source).unwrap(); + wait_for_readiness(&publications, handle, ReadinessFlags::READ); + assert_eq!( + litebox_broker_core::socket::receive_from( + &session, + handle, + &mut [0_u8; 5], + ReceiveFromFlags::NONE, + ), + Err(BrokerError::WouldBlock) + ); + wait_for_readiness(&publications, handle, ReadinessFlags::READ); + let mut reply = [0_u8; 5]; + assert_eq!( + litebox_broker_core::socket::receive_from( + &session, + handle, + &mut reply, + ReceiveFromFlags::NONE, + ), + Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { + received: 5, + datagram_length: 5, + source_address: native_address, + })) + ); + assert_eq!(reply, *b"reply"); + } + + #[test] + fn udp_port_mapping_exposes_a_distinct_host_endpoint() { + let host_reservation = UdpSocket::bind("127.0.0.1:0").unwrap(); + let host_address = socket_address_v4(host_reservation.local_addr().unwrap()); + drop(host_reservation); + let guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53); + let provider = Arc::new( + LinuxSocketProvider::new_with_port_mappings( + 1, + &[SocketPortMapping { + protocol: SocketPortMappingProtocol::Udp, + guest_port: guest_address.port(), + host_address, + }], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(2, 0, 1, 1), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let handle = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + wait_for_readiness(&publications, handle, ReadinessFlags::WRITE); + assert_eq!( + litebox_broker_core::socket::bind(&session, handle, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + create_port_mapping_reservation( + SocketPortMapping { + protocol: SocketPortMappingProtocol::Udp, + guest_port: guest_address.port(), + host_address, + }, + true, + true, + ) + .unwrap_err(), + Errno::ADDRINUSE + ); + + let native = UdpSocket::bind("127.0.0.1:0").unwrap(); + native.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); + native.send_to(b"ping", host_address).unwrap(); + wait_for_readiness(&publications, handle, ReadinessFlags::READ); + let mut request = [0; 4]; + let native_address = socket_address_v4(native.local_addr().unwrap()); + assert_eq!( + litebox_broker_core::socket::receive_from( + &session, + handle, + &mut request, + ReceiveFromFlags::NONE, + ), + Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { + received: 4, + datagram_length: 4, + source_address: native_address, + })) + ); + assert_eq!(request, *b"ping"); + assert_eq!( + litebox_broker_core::socket::status(&session, handle) + .unwrap() + .local_address, + Some(guest_address) + ); + assert_eq!( + litebox_broker_core::socket::send_to( + &session, + handle, + b"pong", + SendFlags::NONE, + Some(native_address), + ), + Ok(SocketOutcome::Completed(4)) + ); + let mut response = [0; 4]; + let (received, source) = native.recv_from(&mut response).unwrap(); + assert_eq!(received, response.len()); + assert_eq!(response, *b"pong"); + assert_eq!(socket_address_v4(source), host_address); + + let mapped_peers = (0..MAX_UDP_PEERS_PER_SOCKET) + .map(|_| UdpSocket::bind("127.0.0.1:0").unwrap()) + .collect::>(); + for peer in &mapped_peers { + let peer_address = socket_address_v4(peer.local_addr().unwrap()); + assert_eq!( + litebox_broker_core::socket::send_to( + &session, + handle, + b"x", + SendFlags::NONE, + Some(peer_address), + ), + Ok(SocketOutcome::Completed(1)) + ); + } + assert_eq!( + litebox_broker_core::socket::connect(&session, handle, native_address), + Ok(SocketOutcome::Completed(SocketConnectionStatus::Connected)) + ); + assert_eq!( + litebox_broker_core::socket::shutdown(&session, handle, ShutdownMode::Write), + Ok(SocketOutcome::Completed(())) + ); + + session.close_object_reference(handle).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), handle); + let conflicting = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + let conflicting_guest_address = + SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 2), guest_address.port()); + assert_eq!( + litebox_broker_core::socket::bind(&session, conflicting, conflicting_guest_address,), + Ok(SocketOutcome::Failed(SocketError::AddressInUse)) + ); + session.close_object_reference(conflicting).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), conflicting); + native.send_to(b"stale", host_address).unwrap(); + let replacement = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + assert_eq!( + litebox_broker_core::socket::bind(&session, replacement, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::receive_from( + &session, + replacement, + &mut [0_u8; 5], + ReceiveFromFlags::NONE, + ), + Err(BrokerError::WouldBlock) + ); + let other_native = UdpSocket::bind("127.0.0.1:0").unwrap(); + other_native.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); + let other_native_address = socket_address_v4(other_native.local_addr().unwrap()); + other_native.send_to(b"fresh", host_address).unwrap(); + let mut fresh = [0_u8; 5]; + assert_eq!( + litebox_broker_core::socket::receive_from( + &session, + replacement, + &mut fresh, + ReceiveFromFlags::NONE, + ), + Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { + received: 5, + datagram_length: 5, + source_address: other_native_address, + })) + ); + assert_eq!(fresh, *b"fresh"); + assert_eq!( + litebox_broker_core::socket::send_to( + &session, + replacement, + b"clean", + SendFlags::NONE, + Some(other_native_address), + ), + Ok(SocketOutcome::Completed(5)) + ); + let mut clean = [0_u8; 5]; + let (received, _) = other_native.recv_from(&mut clean).unwrap(); + assert_eq!(received, clean.len()); + assert_eq!(clean, *b"clean"); + session.close_object_reference(replacement).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), replacement); + let third = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness, + ) + .unwrap(); + assert_eq!( + litebox_broker_core::socket::bind(&session, third, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + } + + #[test] + fn implicit_udp_binding_does_not_claim_a_port_mapping() { + let host_reservation = UdpSocket::bind("127.0.0.1:0").unwrap(); + let host_address = socket_address_v4(host_reservation.local_addr().unwrap()); + drop(host_reservation); + let provider = Arc::new( + LinuxSocketProvider::new_with_port_mappings( + 2, + &[SocketPortMapping { + protocol: SocketPortMappingProtocol::Udp, + guest_port: FIRST_GUEST_EPHEMERAL_PORT, + host_address, + }], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 2, 2), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let handle = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + + assert_eq!( + litebox_broker_core::socket::receive_from( + &session, + handle, + &mut [0], + ReceiveFromFlags::NONE, + ), + Err(BrokerError::WouldBlock) + ); + assert_eq!( + litebox_broker_core::socket::status(&session, handle) + .unwrap() + .local_address, + Some(SocketAddrV4::new( + Ipv4Addr::UNSPECIFIED, + FIRST_GUEST_EPHEMERAL_PORT + 1, + )) + ); + assert_eq!( + UdpSocket::bind(host_address).unwrap_err().kind(), + ErrorKind::AddrInUse + ); + let mapped = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness, + ) + .unwrap(); + let mapped_guest_address = + SocketAddrV4::new(Ipv4Addr::LOCALHOST, FIRST_GUEST_EPHEMERAL_PORT); + assert_eq!( + litebox_broker_core::socket::bind(&session, mapped, mapped_guest_address), + Ok(SocketOutcome::Completed(mapped_guest_address)) + ); + } + + #[test] + fn guest_tcp_loopback_routes_within_the_session_namespace() { + let provider = Arc::new(LinuxSocketProvider::new(3).unwrap()); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 3, 3), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let listener = create_socket(&session, readiness.clone()); + let guest_listener_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + assert_eq!( + litebox_broker_core::socket::bind(&session, listener, guest_listener_address), + Ok(SocketOutcome::Completed(guest_listener_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, listener, 2), + Ok(SocketOutcome::Completed(guest_listener_address)) + ); + + let client = create_socket(&session, readiness.clone()); + let connect = + litebox_broker_core::socket::connect(&session, client, guest_listener_address).unwrap(); + assert!(matches!( + connect, + SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + ) + )); + wait_until_connected(&session, client, &publications); + let client_address = litebox_broker_core::socket::status(&session, client) + .unwrap() + .local_address + .expect("connected client must have a guest-local address"); + session.close_object_reference(client).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), client); + + let replacement = create_socket(&session, readiness.clone()); + let replacement_address = + SocketAddrV4::new(Ipv4Addr::LOCALHOST, FIRST_GUEST_EPHEMERAL_PORT + 1); + assert_eq!( + litebox_broker_core::socket::bind(&session, replacement, replacement_address), + Ok(SocketOutcome::Completed(replacement_address)) + ); + let connect = + litebox_broker_core::socket::connect(&session, replacement, guest_listener_address) + .unwrap(); + assert!(matches!( + connect, + SocketOutcome::Completed( + SocketConnectionStatus::Connecting | SocketConnectionStatus::Connected + ) + )); + wait_until_connected(&session, replacement, &publications); + if !session + .check_readiness(listener) + .unwrap() + .contains(ReadinessFlags::READ) + { + wait_for_readiness(&publications, listener, ReadinessFlags::READ); + } + let accepted = + match litebox_broker_core::socket::accept(&session, listener, readiness.clone()) + .unwrap() + { + SocketOutcome::Completed(accepted) => accepted, + SocketOutcome::Failed(error) => panic!("accept failed: {error:?}"), + }; + assert_eq!(accepted.local_address, guest_listener_address); + assert_eq!(accepted.remote_address, client_address); + session.close_object_reference(accepted.handle).unwrap(); + assert_eq!( + retirements.recv_timeout(TEST_TIMEOUT).unwrap(), + accepted.handle + ); + + let accepted = + match litebox_broker_core::socket::accept(&session, listener, readiness).unwrap() { + SocketOutcome::Completed(accepted) => accepted, + SocketOutcome::Failed(error) => panic!("replacement accept failed: {error:?}"), + }; + assert_eq!(accepted.local_address, guest_listener_address); + assert_eq!(accepted.remote_address, replacement_address); + assert_eq!( + litebox_broker_core::socket::status(&session, replacement) + .unwrap() + .local_address, + Some(replacement_address) + ); + assert_eq!( + litebox_broker_core::socket::send(&session, replacement, b"x", SendFlags::NONE), + Ok(SocketOutcome::Completed(1)) + ); + wait_for_readiness(&publications, accepted.handle, ReadinessFlags::READ); + let mut byte = [0]; + assert_eq!( + litebox_broker_core::socket::receive( + &session, + accepted.handle, + &mut byte, + ReceiveFlags::NONE, + 0, + 0, + ), + Ok(SocketOutcome::Completed(ReceiveSocketResponse::Received(1))) + ); + assert_eq!(byte, *b"x"); + } + + #[test] + fn queued_udp_datagram_retains_closed_sender_guest_address() { + let provider = Arc::new(LinuxSocketProvider::new(3).unwrap()); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 3, 3), + provider.clone(), + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let receiver = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + let receiver_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53); + assert_eq!( + litebox_broker_core::socket::bind(&session, receiver, receiver_address), + Ok(SocketOutcome::Completed(receiver_address)) + ); + + let sender = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + let sender_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 54); + assert_eq!( + litebox_broker_core::socket::bind(&session, sender, sender_address), + Ok(SocketOutcome::Completed(sender_address)) + ); + let old_host_address = provider + .reactor + .host_address(SocketKind::Udp, sender_address.port()) + .unwrap(); + assert_eq!( + litebox_broker_core::socket::send_to( + &session, + sender, + b"queued", + SendFlags::NONE, + Some(receiver_address), + ), + Ok(SocketOutcome::Completed(6)) + ); + session.close_object_reference(sender).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), sender); -const fn broker_error_from_errno(error: Errno) -> BrokerError { - match broker_resource_error_from_errno(error) { - Some(error) => error, - None => BrokerError::Internal, - } -} + let replacement = litebox_broker_core::socket::create( + &session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness, + ) + .unwrap(); + assert_eq!( + litebox_broker_core::socket::bind(&session, replacement, sender_address), + Ok(SocketOutcome::Completed(sender_address)) + ); + let replacement_host_address = provider + .reactor + .host_address(SocketKind::Udp, sender_address.port()) + .unwrap(); + assert_ne!(replacement_host_address, old_host_address); -const fn broker_resource_error_from_errno(error: Errno) -> Option { - match error { - Errno::NOMEM => Some(BrokerError::OutOfMemory), - Errno::MFILE | Errno::NFILE | Errno::NOBUFS | Errno::NOSPC => { - Some(BrokerError::ResourceExhausted) + if !session + .check_readiness(receiver) + .unwrap() + .contains(ReadinessFlags::READ) + { + wait_for_readiness(&publications, receiver, ReadinessFlags::READ); } - _ => None, + let mut payload = [0; 6]; + assert_eq!( + litebox_broker_core::socket::receive_from( + &session, + receiver, + &mut payload, + ReceiveFromFlags::NONE, + ), + Ok(SocketOutcome::Completed(ReceivedPlatformDatagram { + received: 6, + datagram_length: 6, + source_address: sender_address, + })) + ); + assert_eq!(payload, *b"queued"); } -} - -#[cfg(test)] -mod tests { - use std::io::{Read as _, Write as _}; - use std::net::Ipv4Addr; - use std::net::{Shutdown, TcpListener, TcpStream, UdpSocket}; - use std::sync::mpsc::{Receiver, Sender, channel}; - use std::time::{Duration, Instant}; - - use super::*; - use litebox_broker_core::readiness::ReadinessSink; - use litebox_broker_core::{ - BrokerCore, BrokerCoreLimits, CallerCredential, DestinationPortRange, DestinationRule, - Ipv4Cidr, ObjectRights, PolicyEngine, SocketPolicy, - }; - use litebox_broker_protocol::ObjectHandle; - use litebox_broker_protocol::socket::{Ipv4Address, Port}; - - const TEST_TIMEOUT: Duration = Duration::from_secs(5); #[test] - fn cached_socket_error_precedes_a_new_kernel_error() { + fn private_udp_host_ports_are_retained_until_the_owning_session_closes() { + let provider = Arc::new(LinuxSocketProvider::new(3).unwrap()); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(3, 0, 3, 3), + provider.clone(), + ) + .unwrap(); + let first_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let second_session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let first = litebox_broker_core::socket::create( + &first_session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + let first_guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 53); assert_eq!( - shift_pending_error( - Some(SocketError::ConnectionRefused), - Some(SocketError::NetworkUnreachable), - ), - ( - Some(SocketError::ConnectionRefused), - Some(SocketError::NetworkUnreachable), - ) + litebox_broker_core::socket::bind(&first_session, first, first_guest_address), + Ok(SocketOutcome::Completed(first_guest_address)) ); + let retained_host_address = provider + .reactor + .host_address(SocketKind::Udp, first_guest_address.port()) + .unwrap(); + first_session.close_object_reference(first).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), first); + + provider + .reactor + .set_next_private_udp_port(retained_host_address.port()); + let second = litebox_broker_core::socket::create( + &second_session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness.clone(), + ) + .unwrap(); + let second_guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 54); assert_eq!( - shift_pending_error(None, Some(SocketError::NetworkUnreachable)), - (Some(SocketError::NetworkUnreachable), None) + litebox_broker_core::socket::bind(&second_session, second, second_guest_address), + Ok(SocketOutcome::Completed(second_guest_address)) + ); + assert_ne!( + provider + .reactor + .host_address(SocketKind::Udp, second_guest_address.port()) + .unwrap(), + retained_host_address ); - } - #[test] - fn synchronous_errors_do_not_consume_tcp_connect_status() { - assert!(!can_consume_synchronous_error( - SocketKind::Tcp, - SocketConnectionStatus::Connecting, - )); - assert!(can_consume_synchronous_error( - SocketKind::Tcp, - SocketConnectionStatus::Connected, - )); - assert!(can_consume_synchronous_error( - SocketKind::Udp, - SocketConnectionStatus::Unconnected, - )); + drop(first_session); + provider + .reactor + .set_next_private_udp_port(retained_host_address.port()); + let third = litebox_broker_core::socket::create( + &second_session, + CreateSocketRequest { + address_family: AddressFamily::Ipv4, + socket_type: SocketType::Datagram, + protocol: IpProtocol::Udp, + }, + readiness, + ) + .unwrap(); + let third_guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 55); + assert_eq!( + litebox_broker_core::socket::bind(&second_session, third, third_guest_address), + Ok(SocketOutcome::Completed(third_guest_address)) + ); + assert_eq!( + provider + .reactor + .host_address(SocketKind::Udp, third_guest_address.port()) + .unwrap(), + retained_host_address + ); } struct TestReadinessSink { @@ -2203,7 +4593,7 @@ mod tests { .local_address .expect("connected socket must expose its local address"); assert_eq!(*local_address.ip(), Ipv4Addr::LOCALHOST); - assert_ne!(local_address.port(), 0); + assert_eq!(local_address.port(), FIRST_GUEST_EPHEMERAL_PORT); assert_eq!(status.pending_error, None); assert_eq!( litebox_broker_core::socket::get_tcp_option(&session, handle, TcpOptionName::NoDelay,), @@ -2522,9 +4912,192 @@ mod tests { server.join().unwrap(); } + #[test] + fn reactor_assigns_a_port_to_an_unbound_tcp_listener() { + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_port_mappings( + 2, + &[SocketPortMapping { + protocol: SocketPortMappingProtocol::Tcp, + guest_port: FIRST_GUEST_EPHEMERAL_PORT, + host_address, + }], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 2, 2), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, _publications) = channel(); + let (retired, _retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let listener = create_socket(&session, readiness.clone()); + + let local_address = match litebox_broker_core::socket::listen(&session, listener, 1) + .expect("listen request must succeed") + { + SocketOutcome::Completed(address) => address, + SocketOutcome::Failed(error) => panic!("listen failed: {error:?}"), + }; + assert_eq!(local_address.port(), FIRST_GUEST_EPHEMERAL_PORT + 1); + assert_eq!( + TcpListener::bind(host_address).unwrap_err().kind(), + ErrorKind::AddrInUse + ); + let mapped_listener = create_socket(&session, readiness); + let mapped_guest_address = + SocketAddrV4::new(Ipv4Addr::LOCALHOST, FIRST_GUEST_EPHEMERAL_PORT); + assert_eq!( + litebox_broker_core::socket::set_tcp_option( + &session, + mapped_listener, + TcpOptionValue::NoDelay(true), + ), + Ok(()) + ); + assert_eq!( + litebox_broker_core::socket::set_tcp_option( + &session, + mapped_listener, + TcpOptionValue::KeepAlive(true), + ), + Ok(()) + ); + assert_eq!( + litebox_broker_core::socket::bind(&session, mapped_listener, mapped_guest_address,), + Ok(SocketOutcome::Completed(mapped_guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::get_tcp_option( + &session, + mapped_listener, + TcpOptionName::NoDelay, + ), + Ok(TcpOptionValue::NoDelay(true)) + ); + assert_eq!( + litebox_broker_core::socket::get_tcp_option( + &session, + mapped_listener, + TcpOptionName::KeepAlive, + ), + Ok(TcpOptionValue::KeepAlive(true)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, mapped_listener, 1), + Ok(SocketOutcome::Completed(mapped_guest_address)) + ); + assert_eq!( + TcpListener::bind(host_address).unwrap_err().kind(), + ErrorKind::AddrInUse + ); + assert_eq!( + create_port_mapping_reservation( + SocketPortMapping { + protocol: SocketPortMappingProtocol::Tcp, + guest_port: FIRST_GUEST_EPHEMERAL_PORT, + host_address, + }, + true, + true, + ) + .unwrap_err(), + Errno::ADDRINUSE + ); + TcpStream::connect(host_address).unwrap(); + } + + #[test] + fn mapped_tcp_listener_can_be_replaced_while_an_accepted_socket_remains() { + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_port_mappings( + 2, + &[SocketPortMapping { + protocol: SocketPortMappingProtocol::Tcp, + guest_port: 80, + host_address, + }], + ) + .unwrap(), + ); + let broker = BrokerCore::new_with_limits( + PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) + .with_socket_policy(SocketPolicy::Ipv4Loopback), + BrokerCoreLimits::new_with_all_limits(4, 0, 2, 2), + provider, + ) + .unwrap(); + let session = broker + .create_session(CallerCredential::Unauthenticated) + .unwrap(); + let (published, publications) = channel(); + let (retired, retirements) = channel(); + let readiness = Arc::new(TestReadinessSink { published, retired }); + let guest_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 80); + let listener = create_socket(&session, readiness.clone()); + assert_eq!( + litebox_broker_core::socket::bind(&session, listener, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, listener, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + + let first_client = TcpStream::connect(host_address).unwrap(); + wait_for_readiness(&publications, listener, ReadinessFlags::READ); + let accepted = + match litebox_broker_core::socket::accept(&session, listener, readiness.clone()) + .unwrap() + { + SocketOutcome::Completed(accepted) => accepted, + SocketOutcome::Failed(error) => panic!("accept failed: {error:?}"), + }; + assert_eq!( + litebox_broker_core::socket::shutdown(&session, listener, ShutdownMode::StopListening,), + Ok(SocketOutcome::Completed(())) + ); + session.close_object_reference(listener).unwrap(); + assert_eq!(retirements.recv_timeout(TEST_TIMEOUT).unwrap(), listener); + + let replacement = create_socket(&session, readiness); + assert_eq!( + litebox_broker_core::socket::bind(&session, replacement, guest_address), + Ok(SocketOutcome::Completed(guest_address)) + ); + assert_eq!( + litebox_broker_core::socket::listen(&session, replacement, 1), + Ok(SocketOutcome::Completed(guest_address)) + ); + let second_client = TcpStream::connect(host_address).unwrap(); + + drop((first_client, second_client)); + session.close_object_reference(accepted.handle).unwrap(); + } + #[test] fn reactor_drives_a_loopback_tcp_listener() { - let provider = Arc::new(LinuxSocketProvider::new(4).unwrap()); + let host_address = unused_tcp_address(); + let provider = Arc::new( + LinuxSocketProvider::new_with_port_mappings( + 4, + &[SocketPortMapping { + protocol: SocketPortMappingProtocol::Tcp, + guest_port: FIRST_GUEST_EPHEMERAL_PORT, + host_address, + }], + ) + .unwrap(), + ); let broker = BrokerCore::new_with_limits( PolicyEngine::with_unauthenticated_rights(ObjectRights::all()) .with_socket_policy(SocketPolicy::Ipv4Loopback), @@ -2539,7 +5112,7 @@ mod tests { let (retired, retirements) = channel(); let readiness = Arc::new(TestReadinessSink { published, retired }); let listener = create_socket(&session, readiness.clone()); - let requested_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0); + let requested_address = SocketAddrV4::new(Ipv4Addr::LOCALHOST, FIRST_GUEST_EPHEMERAL_PORT); let local_address = match litebox_broker_core::socket::bind(&session, listener, requested_address).unwrap() { @@ -2547,7 +5120,7 @@ mod tests { SocketOutcome::Failed(error) => panic!("bind failed: {error:?}"), }; assert_eq!(local_address.ip(), requested_address.ip()); - assert_ne!(local_address.port(), 0); + assert_eq!(local_address, requested_address); assert_eq!( litebox_broker_core::socket::listen(&session, listener, 8), Ok(SocketOutcome::Completed(local_address)) @@ -2569,8 +5142,8 @@ mod tests { .contains(ReadinessFlags::READ) ); - let mut first_client = TcpStream::connect(local_address).unwrap(); - let second_client = TcpStream::connect(local_address).unwrap(); + let mut first_client = TcpStream::connect(host_address).unwrap(); + let second_client = TcpStream::connect(host_address).unwrap(); first_client.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); first_client.set_write_timeout(Some(TEST_TIMEOUT)).unwrap(); wait_for_readiness(&publications, listener, ReadinessFlags::READ); @@ -2824,7 +5397,7 @@ mod tests { assert_eq!(status.status, SocketConnectionStatus::Unconnected); let local_address = status.local_address.unwrap(); assert!(local_address.ip().is_unspecified()); - assert_eq!(local_address.port(), source.port()); + assert_eq!(local_address, implicitly_bound); server.send_to(&[], source).unwrap(); wait_for_readiness(&publications, handle, ReadinessFlags::READ); @@ -2900,7 +5473,13 @@ mod tests { ); let connected_status = litebox_broker_core::socket::status(&session, handle).unwrap(); assert_eq!(connected_status.status, SocketConnectionStatus::Connected); - assert_eq!(connected_status.local_address, Some(source)); + assert_eq!( + connected_status.local_address, + Some(SocketAddrV4::new( + Ipv4Addr::LOCALHOST, + implicitly_bound.port(), + )) + ); let maximum = vec![0x5a; MAX_UDP_DATAGRAM_SIZE as usize]; assert_eq!( litebox_broker_core::socket::send_to(&session, handle, &maximum, SendFlags::NONE, None,), @@ -3021,6 +5600,13 @@ mod tests { address } + fn unused_tcp_address() -> SocketAddrV4 { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let address = socket_address_v4(listener.local_addr().unwrap()); + drop(listener); + address + } + fn create_socket( session: &litebox_broker_core::BrokerSession, readiness: Arc, diff --git a/litebox_broker_protocol/src/socket.rs b/litebox_broker_protocol/src/socket.rs index 36867db2f4..f81c77edd8 100644 --- a/litebox_broker_protocol/src/socket.rs +++ b/litebox_broker_protocol/src/socket.rs @@ -305,7 +305,7 @@ impl SocketError { pub enum SocketOutcome { /// The operation completed successfully. Completed(T), - /// The host network stack reported an ordinary socket failure. + /// The broker socket backend reported an ordinary socket failure. Failed(SocketError), } @@ -355,7 +355,7 @@ pub struct BindSocketRequest { /// Response to a socket bind request. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub struct BindSocketResponse { - /// Local address assigned by the host network stack. + /// Local address reserved in the guest's broker-managed namespace. pub local_address: SocketAddrV4, } @@ -371,7 +371,7 @@ pub struct ListenSocketRequest { /// Response to a socket listen request. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub struct ListenSocketResponse { - /// Local address assigned by the host network stack. + /// Local address reserved in the guest's broker-managed namespace. pub local_address: SocketAddrV4, } @@ -553,7 +553,7 @@ pub struct SocketStatusRequest { pub struct SocketStatusResponse { /// Current connection state. pub status: SocketConnectionStatus, - /// Local endpoint assigned by the host network stack, if any. + /// Local endpoint reserved in the guest's broker-managed namespace, if any. pub local_address: Option, /// Pending asynchronous socket error, consumed by this status query. pub pending_error: Option, diff --git a/litebox_broker_userland/src/main.rs b/litebox_broker_userland/src/main.rs index 61db89ec2b..5c29ecd228 100644 --- a/litebox_broker_userland/src/main.rs +++ b/litebox_broker_userland/src/main.rs @@ -4,7 +4,7 @@ use std::error::Error; use std::ffi::OsString; use std::io::{Error as IoError, ErrorKind, Result as IoResult}; -use std::net::Ipv4Addr; +use std::net::{Ipv4Addr, SocketAddrV4}; use std::os::unix::net::{UnixListener, UnixStream}; use std::path::PathBuf; use std::process::{Child, Command}; @@ -22,7 +22,9 @@ use litebox_broker_core::{ Ipv4Cidr, ObjectRights, PolicyEngine, SocketPolicy, SocketPolicyError, }; use litebox_broker_host::{BrokerHostAssociation, ConnectionTermination, setup_connection}; -use litebox_broker_platform_linux_userland::LinuxSocketProvider; +use litebox_broker_platform_linux_userland::{ + LinuxSocketProvider, SocketPortMapping, SocketPortMappingProtocol, +}; use litebox_broker_protocol::message::BrokerRequest; use litebox_broker_protocol::shared_buffer::{SHARED_BUFFER_LAYOUT, SHARED_BUFFER_POOL_SIZE}; use litebox_broker_protocol::socket::{Ipv4Address, Port}; @@ -85,6 +87,35 @@ impl FromStr for AllowedTcpDestination { } } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct PortMappingArgument { + host_address: SocketAddrV4, + guest_port: u16, +} + +impl FromStr for PortMappingArgument { + type Err = String; + + fn from_str(value: &str) -> Result { + let (host_address, guest_port) = value + .rsplit_once(':') + .ok_or_else(|| "expected HOST_IP:HOST_PORT:GUEST_PORT".to_owned())?; + let host_address = host_address + .parse::() + .map_err(|error| format!("invalid host IPv4 endpoint: {error}"))?; + let guest_port = guest_port + .parse::() + .map_err(|error| format!("invalid guest port: {error}"))?; + if host_address.port() == 0 || guest_port == 0 { + return Err("mapped host and guest ports must be nonzero".to_owned()); + } + Ok(Self { + host_address, + guest_port, + }) + } +} + #[derive(Parser, Debug)] struct CliArgs { /// Permit outbound TCP connections to a destination CIDR and port range. @@ -93,6 +124,12 @@ struct CliArgs { /// `0.0.0.0/0:1-65535` permits every nonzero IPv4 TCP destination. #[arg(long, value_name = "CIDR:PORT[-PORT]")] allow_tcp_destination: Vec, + /// Publish a host IPv4 TCP endpoint to a guest-local TCP port. + #[arg(long, value_name = "HOST_IP:HOST_PORT:GUEST_PORT")] + publish_tcp: Vec, + /// Publish a host IPv4 UDP endpoint to a guest-local UDP port. + #[arg(long, value_name = "HOST_IP:HOST_PORT:GUEST_PORT")] + publish_udp: Vec, /// Local runner executable to launch. #[arg(long, value_name = "PATH", value_hint = clap::ValueHint::ExecutablePath)] runner: PathBuf, @@ -110,11 +147,15 @@ fn main() -> Result<(), Box> { let control_listener = UnixListener::bind(&control_socket_path)?; control_listener.set_nonblocking(true)?; let limits = BrokerCoreLimits::DEFAULT; + let port_mappings = configured_port_mappings(&args.publish_tcp, &args.publish_udp); let broker = BrokerCore::new_with_limits( PolicyEngine::with_host_guaranteed_rights(ObjectRights::all()) .with_socket_policy(configured_socket_policy(&args.allow_tcp_destination)?), limits, - Arc::new(LinuxSocketProvider::new(limits.max_sockets)?), + Arc::new(LinuxSocketProvider::new_with_port_mappings( + limits.max_sockets, + &port_mappings, + )?), )?; let mut runner_command = Command::new(&args.runner); @@ -139,6 +180,24 @@ fn main() -> Result<(), Box> { Ok(()) } +fn configured_port_mappings( + tcp: &[PortMappingArgument], + udp: &[PortMappingArgument], +) -> Vec { + tcp.iter() + .map(|mapping| SocketPortMapping { + protocol: SocketPortMappingProtocol::Tcp, + guest_port: mapping.guest_port, + host_address: mapping.host_address, + }) + .chain(udp.iter().map(|mapping| SocketPortMapping { + protocol: SocketPortMappingProtocol::Udp, + guest_port: mapping.guest_port, + host_address: mapping.host_address, + })) + .collect() +} + fn configured_socket_policy( allowed_destinations: &[AllowedTcpDestination], ) -> Result { @@ -563,6 +622,30 @@ mod tests { assert!("203.0.113.0/24:0".parse::().is_err()); } + #[test] + fn socket_port_mapping_arguments_name_distinct_host_and_guest_ports() { + let mapping = "127.0.0.1:8080:80".parse::().unwrap(); + + assert_eq!( + mapping, + PortMappingArgument { + host_address: "127.0.0.1:8080".parse().unwrap(), + guest_port: 80, + } + ); + assert!("127.0.0.1:0:80".parse::().is_err()); + assert!("127.0.0.1:8080:0".parse::().is_err()); + assert!("8080:80".parse::().is_err()); + assert_eq!( + configured_port_mappings(&[mapping], &[]), + vec![SocketPortMapping { + protocol: SocketPortMappingProtocol::Tcp, + guest_port: 80, + host_address: "127.0.0.1:8080".parse().unwrap(), + }] + ); + } + #[test] fn tcp_destination_arguments_replace_the_loopback_default() { assert_eq!( diff --git a/litebox_common_linux/src/errno/mod.rs b/litebox_common_linux/src/errno/mod.rs index cf2127f120..780bd88944 100644 --- a/litebox_common_linux/src/errno/mod.rs +++ b/litebox_common_linux/src/errno/mod.rs @@ -388,7 +388,6 @@ impl From for Errno { match value { litebox::net::errors::ConnectError::InvalidFd => Errno::EBADF, litebox::net::errors::ConnectError::UnsupportedAddress(_) => Errno::EAFNOSUPPORT, - litebox::net::errors::ConnectError::PortAllocationFailure(_) => Errno::EADDRINUSE, litebox::net::errors::ConnectError::Unaddressable => Errno::ECONNREFUSED, litebox::net::errors::ConnectError::InProgress => Errno::EINPROGRESS, litebox::net::errors::ConnectError::InvalidState => Errno::ECONNREFUSED, @@ -482,7 +481,6 @@ impl From for Errno { litebox::net::errors::ListenError::InvalidFd => Errno::EBADF, litebox::net::errors::ListenError::InvalidAddress => Errno::EINVAL, litebox::net::errors::ListenError::InvalidState => Errno::EINVAL, - litebox::net::errors::ListenError::NoAvailableFreeEphemeralPorts => Errno::ENOSPC, litebox::net::errors::ListenError::UnsupportedOperation => Errno::EOPNOTSUPP, litebox::net::errors::ListenError::OperationFailed(error) => error.into(), @@ -491,19 +489,6 @@ impl From for Errno { } } -impl From for Errno { - fn from(value: litebox::net::local_ports::LocalPortAllocationError) -> Self { - match value { - litebox::net::local_ports::LocalPortAllocationError::AlreadyInUse(_) => { - Errno::EADDRINUSE - } - litebox::net::local_ports::LocalPortAllocationError::NoAvailableFreePorts => { - Errno::EAGAIN - } - } - } -} - impl From for Errno { fn from(value: litebox::net::errors::SendError) -> Self { match value { @@ -512,7 +497,6 @@ impl From for Errno { litebox::net::errors::SendError::Unaddressable => Errno::EINVAL, litebox::net::errors::SendError::BufferFull => Errno::EAGAIN, litebox::net::errors::SendError::MessageTooLong => Errno::EMSGSIZE, - litebox::net::errors::SendError::PortAllocationFailure(e) => e.into(), litebox::net::errors::SendError::UnnecessaryDestinationAddress => Errno::EISCONN, litebox::net::errors::SendError::DestinationAddressRequired => Errno::EDESTADDRREQ, _ => unimplemented!(), diff --git a/litebox_platform_linux_kernel/src/host/mock.rs b/litebox_platform_linux_kernel/src/host/mock.rs index 0849edaef0..1d2fe2279b 100644 --- a/litebox_platform_linux_kernel/src/host/mock.rs +++ b/litebox_platform_linux_kernel/src/host/mock.rs @@ -58,14 +58,6 @@ impl HostInterface for MockHostInterface { todo!() } - fn send_ip_packet(_packet: &[u8]) -> Result { - todo!() - } - - fn receive_ip_packet(_packet: &mut [u8]) -> Result { - todo!() - } - fn log(msg: &str) { let _ = unsafe { syscalls::syscall3( diff --git a/litebox_platform_linux_kernel/src/host/snp/snp-sandbox.h b/litebox_platform_linux_kernel/src/host/snp/snp-sandbox.h index 83d6d84f23..1c49466d1a 100644 --- a/litebox_platform_linux_kernel/src/host/snp/snp-sandbox.h +++ b/litebox_platform_linux_kernel/src/host/snp/snp-sandbox.h @@ -17,8 +17,6 @@ #define SNP_VMPL_ALLOC_REQ 0x4 #define SNP_VMPL_KPTI_REQ 0x5 #define SNP_VMPL_PRINT_REQ 0x6 -#define SNP_VMPL_TUN_READ_REQ 0x7 -#define SNP_VMPL_TUN_WRITE_REQ 0x8 #define SNP_VMPL_SLEEP_REQ 0x9 #define SNP_VMPL_ALLOC_FUTEX_REQ 0xa #define SNP_VMPL_FILEMAP_READ_REQ 0xb diff --git a/litebox_platform_linux_kernel/src/host/snp/snp_impl.rs b/litebox_platform_linux_kernel/src/host/snp/snp_impl.rs index cc994a2202..5c6173f7c0 100644 --- a/litebox_platform_linux_kernel/src/host/snp/snp_impl.rs +++ b/litebox_platform_linux_kernel/src/host/snp/snp_impl.rs @@ -432,26 +432,6 @@ impl HostSnpInterface { } impl HostInterface for HostSnpInterface { - fn send_ip_packet(packet: &[u8]) -> Result { - let mut req = bindings::SnpVmplRequestArgs::new_request( - bindings::SNP_VMPL_TUN_WRITE_REQ, - 3, - [packet.as_ptr() as u64, packet.len() as u64, 0, 0, 0, 0], - ); - Self::request(&mut req); - Self::parse_result(req.ret) - } - - fn receive_ip_packet(packet: &mut [u8]) -> Result { - let mut req = bindings::SnpVmplRequestArgs::new_request( - bindings::SNP_VMPL_TUN_READ_REQ, - 3, - [packet.as_ptr() as u64, packet.len() as u64, 0, 0, 0, 0], - ); - Self::request(&mut req); - Self::parse_result(req.ret) - } - fn log(msg: &str) { ghcb_prints(msg); } diff --git a/litebox_platform_linux_kernel/src/lib.rs b/litebox_platform_linux_kernel/src/lib.rs index bfc2fc01ce..7934813309 100644 --- a/litebox_platform_linux_kernel/src/lib.rs +++ b/litebox_platform_linux_kernel/src/lib.rs @@ -13,9 +13,8 @@ use litebox::mm::linux::PageRange; use litebox::platform::RawPointerProvider; use litebox::platform::page_mgmt::FixedAddressBehavior; use litebox::platform::{ - ArchSpecificError, ArchSpecificProvider, ArchSpecificRegister, IPInterfaceProvider, - ImmediatelyWokenUp, PageManagementProvider, Provider, RawMutexProvider, TimeProvider, - UnblockedOrTimedOut, + ArchSpecificError, ArchSpecificProvider, ArchSpecificRegister, ImmediatelyWokenUp, + PageManagementProvider, Provider, RawMutexProvider, TimeProvider, UnblockedOrTimedOut, }; use litebox_common_linux::errno::Errno; @@ -285,47 +284,6 @@ impl litebox::platform::SystemTime for SystemTime { } } -impl IPInterfaceProvider for LinuxKernel { - fn send_ip_packet(&self, packet: &[u8]) -> Result<(), litebox::platform::SendError> { - match Host::send_ip_packet(packet) { - Ok(n) => { - if n != packet.len() { - unimplemented!() - } - Ok(()) - } - Err(e) => { - // Avoid allocation for error message - crate::print_str_and_int!( - "Error sending IP packet: ", - u64::from(e.as_neg().unsigned_abs()), - 16 - ); - unimplemented!() - } - } - } - - fn receive_ip_packet( - &self, - packet: &mut [u8], - ) -> Result { - match Host::receive_ip_packet(packet) { - Ok(n) => Ok(n), - Err(Errno::EAGAIN) => Err(litebox::platform::ReceiveError::WouldBlock), - Err(e) => { - // Avoid allocation for error message - crate::print_str_and_int!( - "Error receiving IP packet: ", - u64::from(e.as_neg().unsigned_abs()), - 16 - ); - unimplemented!() - } - } - } -} - impl litebox::platform::StdioProvider for LinuxKernel { fn read_from_stdin(&self, buf: &mut [u8]) -> Result { Host::read_from_stdin(buf).map_err(|err| match err { @@ -385,11 +343,6 @@ pub trait HostInterface: 'static { /// Terminate the current process. fn terminate_process(code: i32) -> !; - /// For Network - fn send_ip_packet(packet: &[u8]) -> Result; - - fn receive_ip_packet(packet: &mut [u8]) -> Result; - // For Stdio fn read_from_stdin(buf: &mut [u8]) -> Result; diff --git a/litebox_platform_linux_userland/README.md b/litebox_platform_linux_userland/README.md deleted file mode 100644 index 3e84b314af..0000000000 --- a/litebox_platform_linux_userland/README.md +++ /dev/null @@ -1,9 +0,0 @@ -# A LiteBox Platform for running LiteBox on userland Linux - -This crate provides an instantiation of the LiteBox `platform::Provider`, with -parameterized punchthrough. - -It requires access to a TUN device that has been initialized. For convenience, -[`./scripts/tun-setup.sh`](./scripts/tun-setup.sh) will help you initialize a -TUN device and ready it for usage with for this platform. Passing `-h` as an -argument will show a help message for how to use this helper script. diff --git a/litebox_platform_linux_userland/scripts/_common.sh b/litebox_platform_linux_userland/scripts/_common.sh deleted file mode 100644 index 6708004527..0000000000 --- a/litebox_platform_linux_userland/scripts/_common.sh +++ /dev/null @@ -1,50 +0,0 @@ -#! /bin/bash - -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -# Script intended to be sourced as a common helper for other scripts. - -set -eo pipefail - -SCRIPT_DIR=$( cd "$( dirname "$0" )" && pwd -P ) - -RED="\033[0;31m" -YELLOW="\033[0;33m" -GREEN="\033[0;32m" -BOLD="\033[1m" -RESET="\033[0m" - -fatal() { - echo -e "${RED}${BOLD}[!]${RESET} $1" 1>&2 - exit 1 -} - -warn() { - echo -e "${YELLOW}${BOLD}[!]${RESET} $1" 1>&2 -} - -info() { - echo -e "${BOLD}[i]${RESET} $1" 1>&2 -} -info2() { - echo -e " $1" 1>&2 -} - -success() { - echo -e "${GREEN}${BOLD}[+]${RESET} $1" 1>&2 -} - -check_for_tools() { - missing_tools=0 - while [ $# -gt 0 ]; do - if ! command -v "$1" &> /dev/null; then - warn "Required tool ${BOLD}$1${RESET} not found" - missing_tools=1 - fi - shift - done - if [ $missing_tools -ne 0 ]; then - fatal "Please install the missing tools and try again" - fi -} diff --git a/litebox_platform_linux_userland/scripts/tun-setup.sh b/litebox_platform_linux_userland/scripts/tun-setup.sh deleted file mode 100755 index 2334583857..0000000000 --- a/litebox_platform_linux_userland/scripts/tun-setup.sh +++ /dev/null @@ -1,94 +0,0 @@ -#! /bin/bash - -# Copyright (c) Microsoft Corporation. -# Licensed under the MIT license. - -source "$(dirname "$0")/_common.sh" - -# If not root, exit -if [ "$(id -u)" -ne 0 ]; then - fatal "Requires running as root" -fi - -# Parse arguments -REMOVE_TUN=0 -FORCE_RECREATE=0 -TUN_DEV="tun99" -TUN_IP="10.0.0.1" -while getopts ":hdft:i:" opt; do - case $opt in - h) - echo "Usage: $0 [-h] [-d | -f] [-t TUN] [-i IP]" 1>&2 - echo " -h Show this help message" 1>&2 - echo " -d Remove the TUN device" 1>&2 - echo " -f Force re-create the TUN device" 1>&2 - echo " -t TUN device name (default: $TUN_DEV)" 1>&2 - echo " -i TUN IP address (default: $TUN_IP)" 1>&2 - exit 0 - ;; - d) - REMOVE_TUN=1 - ;; - f) - FORCE_RECREATE=1 - ;; - t) - TUN_DEV="$OPTARG" - ;; - i) - TUN_IP="$OPTARG" - ;; - :) - fatal "Option -$OPTARG requires an argument" - ;; - \?) - fatal "Invalid option: -$OPTARG" - ;; - esac -done -if [ $REMOVE_TUN -eq 1 ] && [ $FORCE_RECREATE -eq 1 ]; then - fatal "Cannot remove and force re-create at the same time" -fi - -check_for_tools ip - -info "Script parameters:" -info2 "TUN device: ${BOLD}${TUN_DEV}${RESET}" -if [ $REMOVE_TUN -eq 1 ]; then - info2 "Remove TUN device: ${BOLD}yes${RESET}" -else - info2 "TUN IPs: ${BOLD}${TUN_IP}/24${RESET}" -fi -if [ $FORCE_RECREATE -eq 1 ]; then - info2 "Force recreate TUN device: ${BOLD}yes${RESET}" -fi - -if ip tuntap show | grep "^${TUN_DEV}:" &> /dev/null; then - info "TUN device already exists" - if [ $REMOVE_TUN -eq 1 ] || [ $FORCE_RECREATE -eq 1 ]; then - info "Bringing down TUN device" - ip link set dev "$TUN_DEV" down - info "Removing TUN device" - ip tuntap del dev "$TUN_DEV" mode tun - success "Removed TUN device" - if [ $REMOVE_TUN -eq 1 ]; then - exit 0 - fi - else - fatal "Use ${BOLD}-d${RESET} to remove the TUN device or ${BOLD}-f${RESET} to force re-create it" - fi -elif [ $REMOVE_TUN -eq 1 ]; then - warn "TUN device does not exist" - fatal "Nothing to remove" -fi - -info "Creating TUN device" -ip tuntap add dev "$TUN_DEV" mode tun - -info "Assigning IP addresses" -ip addr add "$TUN_IP"/24 dev "$TUN_DEV" - -info "Bringing up TUN device " -ip link set dev "$TUN_DEV" up - -success "Created TUN device" diff --git a/litebox_platform_linux_userland/src/lib.rs b/litebox_platform_linux_userland/src/lib.rs index 9c295267cd..31fe233189 100644 --- a/litebox_platform_linux_userland/src/lib.rs +++ b/litebox_platform_linux_userland/src/lib.rs @@ -9,7 +9,6 @@ use std::cell::Cell; use std::io::IsTerminal as _; -use std::os::fd::{AsRawFd as _, FromRawFd as _}; use std::path::PathBuf; use std::sync::atomic::{AtomicI32, AtomicU32, Ordering}; use std::time::Duration; @@ -94,7 +93,6 @@ macro_rules! saved_tls { /// This implements the main [`litebox::platform::Provider`] trait, i.e., implements all platform /// traits. pub struct LinuxUserland { - tun_socket_fd: std::sync::RwLock>, /// Reserved pages that are not available for guest programs to use. reserved_pages: Vec>, /// CoW-eligible memory regions. Maps start address of the static slice, to the info needed to @@ -122,115 +120,13 @@ struct CowRegionInfo { file_length: usize, } -const IF_NAMESIZE: usize = 16; -/// Use TUN device -const IFF_TUN: i32 = 0x0001; -/// Do not provide packet information -const IFF_NO_PI: i32 = 0x1000; -/// libc `ifreq` structure, used for TUN/TAP devices. -#[repr(C)] -struct Ifreq { - /// interface name, e.g. "en0" - pub ifr_name: [i8; IF_NAMESIZE], - pub ifr_ifru: Ifru, -} - -#[repr(C)] -#[derive(Clone, Copy)] -struct Ifmap { - mem_start: usize, - mem_end: usize, - base_addr: u16, - irq: u8, - dma: u8, - port: u8, -} - -/// libc `ifreq.ifr_ifru` union, used for TUN/TAP devices. -/// -/// We only need `ifru_flags` for now; `ifru_map` is to ensure the size of the union -/// matches libc. -#[repr(C)] -pub union Ifru { - // pub ifru_addr: crate::sockaddr, - // pub ifru_dstaddr: crate::sockaddr, - // pub ifru_broadaddr: crate::sockaddr, - // pub ifru_netmask: crate::sockaddr, - // pub ifru_hwaddr: crate::sockaddr, - ifru_flags: i16, - // pub ifru_ifindex: i32, - // pub ifru_metric: i32, - // pub ifru_mtu: i32, - ifru_map: Ifmap, - // pub ifru_slave: [i8; IF_NAMESIZE], - // pub ifru_newname: [i8; IF_NAMESIZE], - // pub ifru_data: *mut i8, -} - impl LinuxUserland { /// Create a new userland-Linux platform for use in `LiteBox`. - /// - /// Takes an optional tun device name (such as `"tun0"` or `"tun99"`) to connect networking (if - /// not specified, networking is disabled). - /// - /// # Panics - /// - /// Panics if the tun device could not be successfully opened. - pub fn new(tun_device_name: Option<&str>) -> &'static Self { + pub fn new() -> &'static Self { register_exception_handlers(); - let tun_socket_fd = tun_device_name - .map(|tun_device_name| { - let tun_path = b"/dev/net/tun\0"; - let tun_fd = unsafe { - syscalls::syscall3( - syscalls::Sysno::open, - tun_path.as_ptr() as usize, - (litebox::fs::OFlags::RDWR - | litebox::fs::OFlags::CLOEXEC - | litebox::fs::OFlags::NONBLOCK) - .bits() as usize, - litebox::fs::Mode::empty().bits() as usize, - ) - } - .expect("failed to open tun device"); - - let tunsetiff = |fd: usize, ifreq: *const Ifreq| { - let cmd = - litebox_common_linux::iow!(b'T', 202, size_of::<::core::ffi::c_int>()); - unsafe { - syscalls::syscall3(syscalls::Sysno::ioctl, fd, cmd as usize, ifreq as usize) - } - .expect("failed to set TUN interface flags"); - }; - let ifreq = Ifreq { - ifr_name: { - let mut name = [0i8; 16]; - assert!(tun_device_name.len() < 16); // Note: strictly-less-than 16, to ensure it fits - for (i, b) in tun_device_name.char_indices() { - let b = b as u32; - assert!(b < 128); - name[i] = i8::try_from(b).unwrap(); - } - name - }, - ifr_ifru: Ifru { - // IFF_NO_PI: no tun header - // IFF_TUN: create tun (i.e., IP) - ifru_flags: i16::try_from(IFF_TUN | IFF_NO_PI).unwrap(), - }, - }; - tunsetiff(tun_fd, &raw const ifreq); - - // By taking ownership, we are letting the drop handler automatically run `libc::close` - // when necessary. - unsafe { std::os::fd::OwnedFd::from_raw_fd(tun_fd.reinterpret_as_signed().trunc()) } - }) - .into(); - let reserved_pages = Self::read_maps(); let platform = Self { - tun_socket_fd, reserved_pages, cow_regions: std::sync::RwLock::new(std::collections::BTreeMap::new()), boot_id: std::sync::OnceLock::new(), @@ -393,30 +289,6 @@ impl LinuxUserland { } } - /// Wait until there is data available on the TUN device. - /// - /// # Panics - /// - /// Panics if the TUN device is not initialized. - pub fn wait_on_tun(&self, timeout: Option) { - let tun_fd = self.tun_socket_fd.read().unwrap(); - let mut pfd = libc::pollfd { - fd: tun_fd.as_ref().unwrap().as_raw_fd(), - events: libc::POLLIN, - revents: 0, - }; - let _ = unsafe { - libc::poll( - &raw mut pfd, - 1, - timeout.map_or(-1, |t| { - let ms = t.as_millis(); - i32::try_from(ms).unwrap_or(i32::MAX) - }), - ) - }; - } - #[cfg(target_arch = "x86_64")] #[allow( clippy::missing_panics_doc, @@ -435,7 +307,7 @@ impl LinuxUserland { }; let mut rules = vec![ - // TUN and terminal + // Terminal and broker I/O (libc::SYS_read, vec![]), (libc::SYS_write, vec![]), (libc::SYS_poll, vec![]), @@ -1264,58 +1136,6 @@ impl litebox::platform::RawMutex for RawMutex { } } -impl litebox::platform::IPInterfaceProvider for LinuxUserland { - fn send_ip_packet(&self, packet: &[u8]) -> Result<(), litebox::platform::SendError> { - let tun_fd = self.tun_socket_fd.read().unwrap(); - let Some(tun_socket_fd) = tun_fd.as_ref() else { - unimplemented!("networking without tun is unimplemented") - }; - match unsafe { - syscalls::syscall3( - syscalls::Sysno::write, - usize::try_from(tun_socket_fd.as_raw_fd()).unwrap(), - packet.as_ptr() as usize, - packet.len(), - ) - } { - Ok(n) => { - if n != packet.len() { - unimplemented!("unexpected size {n}") - } - Ok(()) - } - Err(errno) => { - unimplemented!("unexpected error {errno}") - } - } - } - - fn receive_ip_packet( - &self, - packet: &mut [u8], - ) -> Result { - let tun_fd = self.tun_socket_fd.read().unwrap(); - let Some(tun_socket_fd) = tun_fd.as_ref() else { - unimplemented!("networking without tun is unimplemented") - }; - unsafe { - syscalls::syscall3( - syscalls::Sysno::read, - usize::try_from(tun_socket_fd.as_raw_fd()).unwrap(), - packet.as_mut_ptr() as usize, - packet.len(), - ) - } - .map_err(|errno| match errno { - #[allow(unreachable_patterns, reason = "EAGAIN == EWOULDBLOCK")] - syscalls::Errno::EWOULDBLOCK | syscalls::Errno::EAGAIN => { - litebox::platform::ReceiveError::WouldBlock - } - _ => unimplemented!("unexpected error {errno}"), - }) - } -} - impl litebox::platform::TimeProvider for LinuxUserland { type Instant = Instant; type SystemTime = SystemTime; @@ -2541,7 +2361,7 @@ mod tests { #[test] fn test_reserved_pages() { - let platform = LinuxUserland::new(None); + let platform = LinuxUserland::new(); let reserved_pages: Vec<_> = >::reserved_pages(platform).collect(); @@ -2565,7 +2385,7 @@ mod tests { unsafe { OwnedFd::from_raw_fd(fd) } } - let _platform: &LinuxUserland = LinuxUserland::new(None); + let _platform: &LinuxUserland = LinuxUserland::new(); let allowed = test_memfd(c"seccomp-allowed-positional-io"); let denied = test_memfd(c"seccomp-denied-positional-io"); let (allowed_shutdown, _allowed_peer) = UnixStream::pair().unwrap(); diff --git a/litebox_platform_lvbs/src/host/lvbs_impl.rs b/litebox_platform_lvbs/src/host/lvbs_impl.rs index 5c0bc70c5d..669af9299f 100644 --- a/litebox_platform_lvbs/src/host/lvbs_impl.rs +++ b/litebox_platform_lvbs/src/host/lvbs_impl.rs @@ -263,14 +263,6 @@ pub struct HostLvbsInterface; impl HostLvbsInterface {} impl HostInterface for HostLvbsInterface { - fn send_ip_packet(_packet: &[u8]) -> Result { - unimplemented!() - } - - fn receive_ip_packet(_packet: &mut [u8]) -> Result { - unimplemented!() - } - fn log(msg: &str) { serial_print_string(msg); } diff --git a/litebox_platform_lvbs/src/host/mock.rs b/litebox_platform_lvbs/src/host/mock.rs index c5f2f3bc0a..f9d485a297 100644 --- a/litebox_platform_lvbs/src/host/mock.rs +++ b/litebox_platform_lvbs/src/host/mock.rs @@ -32,14 +32,6 @@ impl HostInterface for MockHostInterface { todo!() } - fn send_ip_packet(_packet: &[u8]) -> Result { - todo!() - } - - fn receive_ip_packet(_packet: &mut [u8]) -> Result { - todo!() - } - fn log(msg: &str) { unsafe { libc::write(libc::STDOUT_FILENO, msg.as_ptr().cast(), msg.len()) }; } diff --git a/litebox_platform_lvbs/src/lib.rs b/litebox_platform_lvbs/src/lib.rs index a1d17f1148..33d8128400 100644 --- a/litebox_platform_lvbs/src/lib.rs +++ b/litebox_platform_lvbs/src/lib.rs @@ -10,10 +10,9 @@ use crate::{host::per_cpu_variables::PerCpuVariablesAsm, mshv::vsm::Vtl0KernelIn use core::sync::atomic::AtomicU32; use hashbrown::HashMap; use litebox::platform::{ - ArchSpecificError, ArchSpecificProvider, ArchSpecificRegister, IPInterfaceProvider, - ImmediatelyWokenUp, PageManagementProvider, RawMutex as _, RawMutexProvider, - RawPointerProvider, StdioProvider, TimeProvider, UnblockedOrTimedOut, - page_mgmt::DeallocationError, + ArchSpecificError, ArchSpecificProvider, ArchSpecificRegister, ImmediatelyWokenUp, + PageManagementProvider, RawMutex as _, RawMutexProvider, RawPointerProvider, StdioProvider, + TimeProvider, UnblockedOrTimedOut, page_mgmt::DeallocationError, }; use litebox::{ mm::linux::{PAGE_SIZE, PageRange}, @@ -881,34 +880,6 @@ impl litebox::platform::SystemTime for SystemTime { } } -impl IPInterfaceProvider for LinuxKernel { - fn send_ip_packet(&self, packet: &[u8]) -> Result<(), litebox::platform::SendError> { - match Host::send_ip_packet(packet) { - Ok(n) => { - if n != packet.len() { - unimplemented!() - } - Ok(()) - } - Err(e) => { - unimplemented!("Error: {:?}", e) - } - } - } - - fn receive_ip_packet( - &self, - packet: &mut [u8], - ) -> Result { - match Host::receive_ip_packet(packet) { - Ok(n) => Ok(n), - Err(e) => { - unimplemented!("Error: {:?}", e) - } - } - } -} - /// Platform-Host Interface pub trait HostInterface: 'static { /// Page allocation from host. @@ -954,11 +925,6 @@ pub trait HostInterface: 'static { timeout: Option, ) -> Result<(), Errno>; - /// For Network - fn send_ip_packet(packet: &[u8]) -> Result; - - fn receive_ip_packet(packet: &mut [u8]) -> Result; - /// For Debugging fn log(msg: &str); diff --git a/litebox_platform_windows_userland/src/lib.rs b/litebox_platform_windows_userland/src/lib.rs index 2cc2740210..3f4c090fba 100644 --- a/litebox_platform_windows_userland/src/lib.rs +++ b/litebox_platform_windows_userland/src/lib.rs @@ -1375,25 +1375,6 @@ impl litebox::platform::RawMutex for RawMutex { } } -impl litebox::platform::IPInterfaceProvider for WindowsUserland { - fn send_ip_packet(&self, packet: &[u8]) -> Result<(), litebox::platform::SendError> { - unimplemented!( - "send_ip_packet is not implemented for Windows yet. packet length: {}", - packet.len() - ); - } - - fn receive_ip_packet( - &self, - packet: &mut [u8], - ) -> Result { - unimplemented!( - "receive_ip_packet is not implemented for Windows yet. packet length: {}", - packet.len() - ); - } -} - impl litebox::platform::TimeProvider for WindowsUserland { type Instant = Instant; type SystemTime = SystemTime; diff --git a/litebox_runner_linux_userland/src/lib.rs b/litebox_runner_linux_userland/src/lib.rs index a985c1353b..37f78ba1df 100644 --- a/litebox_runner_linux_userland/src/lib.rs +++ b/litebox_runner_linux_userland/src/lib.rs @@ -60,13 +60,6 @@ pub struct CliArgs { help_heading = "Unstable Options" )] pub rewrite_syscalls: bool, - /// Connect to a TUN device with this name - #[arg( - long = "tun-device-name", - requires = "unstable", - help_heading = "Unstable Options" - )] - pub tun_device_name: Option, /// Load the program binary from the tar file instead of from the host filesystem. /// /// When set, the program path refers to a path inside the tar filesystem. @@ -211,7 +204,7 @@ pub fn run(cli_args: CliArgs) -> Result<()> { }; // TODO(jb): Clean up platform initialization once we have https://github.com/MSRSSP/litebox/issues/24 - let platform = Platform::new(cli_args.tun_device_name.as_deref()); + let platform = Platform::new(); for file in cow_eligible_regions { platform.register_cow_region(file.data, file.abs_path); @@ -356,38 +349,6 @@ pub fn run(cli_args: CliArgs) -> Result<()> { let shim = shim_builder.build(); - let shutdown = std::sync::Arc::new(core::sync::atomic::AtomicBool::new(false)); - let net_worker = if cli_args.tun_device_name.is_some() { - let shim = shim.clone(); - let shutdown_clone = shutdown.clone(); - let child = litebox_platform_linux_userland::spawn_host_thread(move || { - const DEFAULT_TIMEOUT: core::time::Duration = core::time::Duration::from_micros(100); - const MAX_TIMEOUT: core::time::Duration = core::time::Duration::from_millis(1); - pin_thread_to_cpu(0); - - while !shutdown_clone.load(core::sync::atomic::Ordering::Relaxed) { - let timeout = loop { - match shim.perform_network_interaction() { - litebox::net::PlatformInteractionReinvocationAdvice::CallAgainImmediately => {} - litebox::net::PlatformInteractionReinvocationAdvice::WaitOnDeviceOrSocketInteraction{ timeout } => { - break timeout; - } - } - }; - // TODO: We only wait for ingress packets on the TUN device and thus may block processing egress packets for up to `timeout`. - // Set a maximum timeout to ensure we don't wait too long. Alternatively, shim could notify us when there are egress packets to process, - // but that would require more invasive changes. - platform.wait_on_tun(Some(timeout.unwrap_or(DEFAULT_TIMEOUT).min(MAX_TIMEOUT))); - } - // Final flush - // TODO: keep running until all sockets are closed? - while shim.perform_network_interaction().call_again_immediately() {} - }); - Some(child) - } else { - None - }; - let argv = cli_args .program_and_arguments .iter() @@ -441,22 +402,5 @@ pub fn run(cli_args: CliArgs) -> Result<()> { } } - if let Some(net_worker) = net_worker { - shutdown.store(true, core::sync::atomic::Ordering::Relaxed); - net_worker.join().unwrap(); - } std::process::exit(program.process.wait()) } - -/// Pin the current thread to a specific CPU core -fn pin_thread_to_cpu(cpu: usize) { - unsafe { - let mut set = std::mem::zeroed(); - libc::CPU_ZERO(&mut set); - libc::CPU_SET(cpu, &mut set); - - if libc::sched_setaffinity(0, std::mem::size_of::(), &raw const set) != 0 { - eprintln!("Warning: Failed to pin thread to CPU core {cpu}"); - } - } -} diff --git a/litebox_runner_linux_userland/tests/loader.rs b/litebox_runner_linux_userland/tests/loader.rs index c9408d6456..6082fbdc0e 100644 --- a/litebox_runner_linux_userland/tests/loader.rs +++ b/litebox_runner_linux_userland/tests/loader.rs @@ -16,12 +16,8 @@ struct TestLauncher { } impl TestLauncher { - fn init_platform( - tar_data: &'static [u8], - initial_files: &[&str], - tun_device_name: Option<&str>, - ) -> Self { - let platform = Platform::new(tun_device_name); + fn init_platform(tar_data: &'static [u8], initial_files: &[&str]) -> Self { + let platform = Platform::new(); let shim_builder = litebox_shim_linux::LinuxShimBuilder::new(platform); let litebox = shim_builder.litebox(); @@ -129,7 +125,6 @@ fn test_load_exec_dynamic() { .iter() .map(std::string::String::as_str) .collect::>(), - None, ); launcher.install_file(executable_data, executable_path); launcher.test_load_exec_common(executable_path); @@ -143,7 +138,7 @@ fn test_load_exec_static() { let executable_path = "/hello_exec"; let executable_data = std::fs::read(path).unwrap(); - let mut launcher = TestLauncher::init_platform(&[], &[], None); + let mut launcher = TestLauncher::init_platform(&[], &[]); launcher.install_file(executable_data, executable_path); @@ -256,7 +251,7 @@ fn test_syscall_rewriter() { let executable_path = "/hello_exec_nolibc.hooked"; let executable_data = std::fs::read(hooked_path).unwrap(); - let mut launcher = TestLauncher::init_platform(&[], &[], None); + let mut launcher = TestLauncher::init_platform(&[], &[]); launcher.install_file(executable_data, executable_path); launcher.test_load_exec_common(executable_path); } diff --git a/litebox_runner_linux_userland/tests/run.rs b/litebox_runner_linux_userland/tests/run.rs index 16da35fd45..937a6154e7 100644 --- a/litebox_runner_linux_userland/tests/run.rs +++ b/litebox_runner_linux_userland/tests/run.rs @@ -341,6 +341,16 @@ fn spawn_test_broker( control_socket_path: &Path, policy: litebox_broker_core::PolicyEngine, connection_count: usize, +) -> TestBroker { + spawn_test_broker_with_port_mappings(control_socket_path, policy, connection_count, Vec::new()) +} + +#[cfg(all(target_arch = "x86_64", target_os = "linux"))] +fn spawn_test_broker_with_port_mappings( + control_socket_path: &Path, + policy: litebox_broker_core::PolicyEngine, + connection_count: usize, + port_mappings: Vec, ) -> TestBroker { let _ = std::fs::remove_file(control_socket_path); @@ -359,8 +369,9 @@ fn spawn_test_broker( policy, limits, std::sync::Arc::new( - litebox_broker_platform_linux_userland::LinuxSocketProvider::new( + litebox_broker_platform_linux_userland::LinuxSocketProvider::new_with_port_mappings( limits.max_sockets, + &port_mappings, ) .expect("failed to create broker test socket provider"), ), @@ -739,23 +750,37 @@ fn test_runner_broker_tcp_server_with_rewriter() { use std::net::{Ipv4Addr, TcpStream}; use std::process::Stdio; + const GUEST_PORT: u16 = 18_080; + let target = common::compile( "./tests/tcp_broker_server.c", "broker_tcp_server_rewriter", false, false, ); + let host_listener = std::net::TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let host_address = match host_listener.local_addr().unwrap() { + std::net::SocketAddr::V4(address) => address, + std::net::SocketAddr::V6(_) => unreachable!("IPv4 bind returned an IPv6 address"), + }; + drop(host_listener); let control_socket_path = unique_test_socket_path("runner-broker-tcp-server-control"); - let broker = spawn_test_broker( + let broker = spawn_test_broker_with_port_mappings( &control_socket_path, litebox_broker_core::PolicyEngine::with_host_guaranteed_rights( litebox_broker_core::ObjectRights::all(), ) .with_socket_policy(litebox_broker_core::SocketPolicy::Ipv4Loopback), 1, + vec![litebox_broker_platform_linux_userland::SocketPortMapping { + protocol: litebox_broker_platform_linux_userland::SocketPortMappingProtocol::Tcp, + guest_port: GUEST_PORT, + host_address, + }], ); let mut child = Runner::new(&target, "broker_tcp_server_rewriter") .broker_socket(&control_socket_path) + .arg(GUEST_PORT.to_string()) .spawn_with_stdio(Stdio::null(), Stdio::piped(), Stdio::inherit()); let stdout = child.stdout.take().unwrap(); let (line_sender, line_receiver) = std::sync::mpsc::channel(); @@ -786,8 +811,9 @@ fn test_runner_broker_tcp_server_with_rewriter() { }; let listen = next_marker("LISTEN "); - let port = listen.split_whitespace().nth(1).unwrap().parse().unwrap(); - let mut first = TcpStream::connect((Ipv4Addr::LOCALHOST, port)).unwrap(); + let port: u16 = listen.split_whitespace().nth(1).unwrap().parse().unwrap(); + assert_eq!(port, GUEST_PORT); + let mut first = TcpStream::connect(host_address).unwrap(); first.set_read_timeout(Some(BROKER_HELPER_TIMEOUT)).unwrap(); first .set_write_timeout(Some(BROKER_HELPER_TIMEOUT)) @@ -813,7 +839,7 @@ fn test_runner_broker_tcp_server_with_rewriter() { child.try_wait().unwrap().is_none(), "blocking accept returned before a client connected" ); - let mut second = TcpStream::connect((Ipv4Addr::LOCALHOST, port)).unwrap(); + let mut second = TcpStream::connect(host_address).unwrap(); second .set_read_timeout(Some(BROKER_HELPER_TIMEOUT)) .unwrap(); diff --git a/litebox_runner_linux_userland/tests/tcp_broker_server.c b/litebox_runner_linux_userland/tests/tcp_broker_server.c index 2d053e9c27..1e07b7dff9 100644 --- a/litebox_runner_linux_userland/tests/tcp_broker_server.c +++ b/litebox_runner_linux_userland/tests/tcp_broker_server.c @@ -11,6 +11,7 @@ #include #include #include +#include #include #include @@ -58,8 +59,13 @@ static void *blocking_accept_thread(void *argument) { return NULL; } -int main(void) { +int main(int argc, char **argv) { setvbuf(stdout, NULL, _IONBF, 0); + assert(argc == 2); + char *end = NULL; + unsigned long guest_port = strtoul(argv[1], &end, 10); + assert(end != argv[1] && *end == '\0' && guest_port > 0 && + guest_port <= UINT16_MAX); int listener = socket(AF_INET, SOCK_STREAM | SOCK_NONBLOCK | SOCK_CLOEXEC, 0); @@ -67,7 +73,7 @@ int main(void) { struct sockaddr_in local = { .sin_family = AF_INET, .sin_addr.s_addr = htonl(INADDR_LOOPBACK), - .sin_port = 0, + .sin_port = htons((uint16_t)guest_port), }; assert(bind(listener, (const struct sockaddr *)&local, sizeof(local)) == 0); assert(listen(listener, 8) == 0); diff --git a/litebox_runner_optee_on_linux_userland/src/lib.rs b/litebox_runner_optee_on_linux_userland/src/lib.rs index e2f0ea7d0b..9c0f0c6ff9 100644 --- a/litebox_runner_optee_on_linux_userland/src/lib.rs +++ b/litebox_runner_optee_on_linux_userland/src/lib.rs @@ -79,7 +79,7 @@ pub fn run(cli_args: CliArgs) -> Result<()> { }; // TODO(jb): Clean up platform initialization once we have https://github.com/MSRSSP/litebox/issues/24 - let platform = Platform::new(None); + let platform = Platform::new(); litebox_platform_multiplex::set_platform(platform); let shim_builder = litebox_shim_optee::OpteeShimBuilder::new(); let _litebox = shim_builder.litebox(); diff --git a/litebox_runner_snp/src/entry.S b/litebox_runner_snp/src/entry.S index 8a2e79f6fd..cdc7040e20 100644 --- a/litebox_runner_snp/src/entry.S +++ b/litebox_runner_snp/src/entry.S @@ -55,13 +55,6 @@ _exit: hlt jmp 1b -sandbox_tun_read_write_start: - and rsp, -16 /* align stack to 16 bytes */ - call sandbox_tun_read_write -1: /* hlt the machine */ - hlt - jmp 1b - sandbox_mount_start: sandbox_signal_ret: rust_begin_unwind: diff --git a/litebox_runner_snp/src/main.rs b/litebox_runner_snp/src/main.rs index d857480ce5..cfb793002b 100644 --- a/litebox_runner_snp/src/main.rs +++ b/litebox_runner_snp/src/main.rs @@ -300,31 +300,6 @@ pub extern "C" fn do_syscall_64(pt_regs: &mut litebox_common_linux::PtRegs) -> ! litebox_platform_linux_kernel::host::snp::snp_impl::handle_syscall(pt_regs); } -#[unsafe(no_mangle)] -pub extern "C" fn sandbox_tun_read_write() { - // wait until shim is initialized - let shim = loop { - if let Some(shim) = SHIM.get() { - break shim; - } - core::hint::spin_loop(); - }; - #[cfg(debug_assertions)] - litebox_util_log::debug!("sandbox_tun_read_write started"); - while !litebox_platform_linux_kernel::host::snp::snp_impl::all_threads_exited() { - let _timeout = loop { - match shim - .perform_network_interaction() { - litebox::net::PlatformInteractionReinvocationAdvice::CallAgainImmediately => {}, - litebox::net::PlatformInteractionReinvocationAdvice::WaitOnDeviceOrSocketInteraction { timeout } => break timeout, - } - }; - // TODO: use timeout to wait on host events - } - - litebox_platform_linux_kernel::host::snp::snp_impl::HostSnpInterface::return_to_host(); -} - /// This function is called on panic. #[panic_handler] fn panic(info: &core::panic::PanicInfo) -> ! { diff --git a/litebox_runner_windows_on_linux_userland/src/lib.rs b/litebox_runner_windows_on_linux_userland/src/lib.rs index a0e55e2695..fadadc4c3e 100644 --- a/litebox_runner_windows_on_linux_userland/src/lib.rs +++ b/litebox_runner_windows_on_linux_userland/src/lib.rs @@ -67,7 +67,7 @@ pub fn run(cli_args: CliArgs) -> Result<()> { let tar_data = std::fs::read(tar_file) .with_context(|| format!("Could not read tar file at {}", tar_file.display()))?; - let platform = LinuxUserland::new(None); + let platform = LinuxUserland::new(); let shim_builder = litebox_shim_windows::WindowsShimBuilder::new(platform); let litebox = shim_builder.litebox(); diff --git a/litebox_shim_linux/src/lib.rs b/litebox_shim_linux/src/lib.rs index 34821cc854..8b3fa4ca77 100644 --- a/litebox_shim_linux/src/lib.rs +++ b/litebox_shim_linux/src/lib.rs @@ -90,7 +90,6 @@ pub trait ShimPlatform: + litebox::platform::ThreadProvider + litebox::platform::TimerProvider + litebox::platform::SignalProvider - + litebox::platform::IPInterfaceProvider + 'static { } @@ -109,7 +108,6 @@ impl ShimPlatform for T where + litebox::platform::ThreadProvider + litebox::platform::TimerProvider + litebox::platform::SignalProvider - + litebox::platform::IPInterfaceProvider + 'static { } @@ -230,8 +228,7 @@ impl LinuxShimBuilder { /// Build the shim. pub fn build(self) -> LinuxShim { - let mut net = Network::new(&self.litebox); - net.set_platform_interaction(litebox::net::PlatformInteraction::Manual); + let net = Network::new(&self.litebox); let global = Arc::new(GlobalState { platform: self.platform, pm: PageManager::new(&self.litebox), @@ -325,15 +322,6 @@ impl LinuxShim { &self.0.pm } - /// Perform queued network interactions with the outside world. - /// - /// This function should be invoked in a loop, based on the returned advice. - pub fn perform_network_interaction( - &self, - ) -> litebox::net::PlatformInteractionReinvocationAdvice { - self.0.net.lock().perform_platform_interaction() - } - /// Establish a TCP connection to the given address. /// /// Returns a [`transport::ShimTransport`] that can be used as a diff --git a/litebox_shim_linux/src/loader/elf.rs b/litebox_shim_linux/src/loader/elf.rs index b0449c25b6..2c5b052031 100644 --- a/litebox_shim_linux/src/loader/elf.rs +++ b/litebox_shim_linux/src/loader/elf.rs @@ -493,7 +493,7 @@ mod tests { #[test] fn et_exec_interpreter_loads_top_down_above_low_heap() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); write_file(&task, "/main", &minimal_elf(ET_EXEC, Some(INTERP_PATH))); write_file(&task, "/ld.so", &minimal_elf(ET_DYN, None)); diff --git a/litebox_shim_linux/src/stdio.rs b/litebox_shim_linux/src/stdio.rs index cc6c79314f..0b209ceddf 100644 --- a/litebox_shim_linux/src/stdio.rs +++ b/litebox_shim_linux/src/stdio.rs @@ -14,7 +14,7 @@ mod tests { #[test] fn test_stdio() { - let task = init_platform(None); + let task = init_platform(); // Check that the stdio streams are in the file table let stdin_stat = task.sys_fstat(0).unwrap(); @@ -60,7 +60,7 @@ mod tests { #[test] fn test_stdio_flags_with_dup() { - let task = init_platform(None); + let task = init_platform(); let stdin = 0; let flags = task.sys_fcntl(stdin, FcntlArg::GETFL).unwrap(); diff --git a/litebox_shim_linux/src/syscalls/epoll.rs b/litebox_shim_linux/src/syscalls/epoll.rs index 102476bb96..4bbe790a89 100644 --- a/litebox_shim_linux/src/syscalls/epoll.rs +++ b/litebox_shim_linux/src/syscalls/epoll.rs @@ -652,14 +652,14 @@ mod test { extern crate std; fn platform() -> &'static TestPlatform { - crate::syscalls::tests::test_platform(None) + crate::syscalls::tests::test_platform() } fn setup_epoll() -> ( crate::Task>, EpollFile>, ) { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); let epoll = EpollFile::new(); (task, epoll) @@ -712,7 +712,7 @@ mod test { #[test] fn test_poll() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); let mut set = super::PollSet::with_capacity(0); let (rfd_u, wfd_u) = task @@ -767,7 +767,7 @@ mod test { #[test] fn test_pselect() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); let (rfd_u, wfd_u) = task .sys_pipe2(litebox::fs::OFlags::empty()) @@ -806,7 +806,7 @@ mod test { #[test] fn test_pselect_read_hup() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); let (rfd_u, wfd_u) = task .sys_pipe2(litebox::fs::OFlags::empty()) @@ -847,7 +847,7 @@ mod test { #[test] fn test_pselect_invalid_fd() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); let invalid_fd_u = 100u32; diff --git a/litebox_shim_linux/src/syscalls/eventfd.rs b/litebox_shim_linux/src/syscalls/eventfd.rs index 4b6e1845ed..d1f48f8d8d 100644 --- a/litebox_shim_linux/src/syscalls/eventfd.rs +++ b/litebox_shim_linux/src/syscalls/eventfd.rs @@ -106,7 +106,7 @@ mod tests { #[test] fn test_eventfd_requires_broker_control() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); assert!(matches!( task.global.create_linux_eventfd(0, EfdFlags::NONBLOCK), diff --git a/litebox_shim_linux/src/syscalls/file.rs b/litebox_shim_linux/src/syscalls/file.rs index 41fa348849..a8243de225 100644 --- a/litebox_shim_linux/src/syscalls/file.rs +++ b/litebox_shim_linux/src/syscalls/file.rs @@ -2920,7 +2920,7 @@ mod tests { #[test] fn getcwd_and_chdir() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); // Default CWD is root. let mut buf = [0u8; 256]; @@ -2963,7 +2963,7 @@ mod tests { #[test] fn chdir_relative_path() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); // Create nested dirs: /rel_parent/rel_child task.sys_mkdirat(litebox_common_linux::AT_FDCWD, "/rel_parent", 0o777) @@ -2995,7 +2995,7 @@ mod tests { fn mknodat_regular_file_does_not_consume_fd_limit() { use litebox_common_linux::{Rlimit, RlimitResource}; - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); let old_limit = task.do_prlimit(RlimitResource::NOFILE, None).unwrap(); task.do_prlimit( RlimitResource::NOFILE, @@ -3023,7 +3023,7 @@ mod tests { #[test] fn empty_pathnames_return_enoent() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); assert_eq!( task.sys_open("", OFlags::RDONLY, Mode::empty()) @@ -3069,7 +3069,7 @@ mod tests { fn all_path_syscalls_respect_chdir() { use litebox_common_linux::{AccessFlags, AtFlags}; - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); // Set up: mkdir + chdir into /cwd_test/. task.sys_mkdirat(litebox_common_linux::AT_FDCWD, "/cwd_test", 0o777) diff --git a/litebox_shim_linux/src/syscalls/misc.rs b/litebox_shim_linux/src/syscalls/misc.rs index ee546e53eb..2dd745bafd 100644 --- a/litebox_shim_linux/src/syscalls/misc.rs +++ b/litebox_shim_linux/src/syscalls/misc.rs @@ -171,7 +171,7 @@ mod tests { fn test_getrandom() { use litebox_common_linux::RngFlags; - let task = init_platform(None); + let task = init_platform(); let mut buf = [0u8; 16]; let ptr = UserPtrMut::from_ptr(buf.as_mut_ptr()); @@ -188,7 +188,7 @@ mod tests { #[test] fn test_uname() { - let task = init_platform(None); + let task = init_platform(); let mut utsname = litebox_common_linux::Utsname::new_zeroed(); let ptr = UserPtrMut::from_ptr(&raw mut utsname); diff --git a/litebox_shim_linux/src/syscalls/mm.rs b/litebox_shim_linux/src/syscalls/mm.rs index d62999b85a..4daa808921 100644 --- a/litebox_shim_linux/src/syscalls/mm.rs +++ b/litebox_shim_linux/src/syscalls/mm.rs @@ -1136,7 +1136,7 @@ mod tests { #[test] fn test_anonymous_mmap() { - let task = init_platform(None); + let task = init_platform(); let addr = task .sys_mmap( @@ -1156,7 +1156,7 @@ mod tests { #[test] fn test_file_backed_mmap() { - let task = init_platform(None); + let task = init_platform(); let content = b"Hello, world!"; let fd = task @@ -1186,7 +1186,7 @@ mod tests { #[test] fn test_mremap() { - let task = init_platform(None); + let task = init_platform(); let addr = task .sys_mmap( @@ -1224,7 +1224,7 @@ mod tests { #[test] fn test_mmap_fixed_noreplace() { - let task = init_platform(None); + let task = init_platform(); // First, create an initial mapping at a specific address away from boundaries let base_addr = 0x1000_0000usize; // 256 MiB - safe middle ground @@ -1334,7 +1334,7 @@ mod tests { #[cfg(any(target_os = "linux", target_os = "windows"))] #[test] fn test_collision_with_global_allocator() { - let task = init_platform(None); + let task = init_platform(); let platform = task.global.platform; let mut data = alloc::vec::Vec::new(); // Find an address that is allocated to the global allocator but not in reserved regions. @@ -1428,7 +1428,7 @@ mod tests { #[test] fn test_map_shared_anonymous() { - let task = init_platform(None); + let task = init_platform(); // MAP_SHARED | MAP_ANON with PROT_READ should succeed let addr = task @@ -1463,7 +1463,7 @@ mod tests { #[test] fn test_map_shared_anonymous_writable() { - let task = init_platform(None); + let task = init_platform(); // MAP_SHARED | MAP_ANON with PROT_WRITE should succeed let addr = task @@ -1486,7 +1486,7 @@ mod tests { #[test] fn test_map_shared_readonly_file() { - let task = init_platform(None); + let task = init_platform(); let content = b"Hello, shared!"; let fd = task @@ -1520,7 +1520,7 @@ mod tests { #[test] fn test_madvise() { - let task = init_platform(None); + let task = init_platform(); let addr = task .sys_mmap( @@ -1566,7 +1566,7 @@ mod tests { #[cfg(not(target_os = "windows"))] #[test] fn test_fallible_read() { - let _ = init_platform(None); + let _ = init_platform(); let ptr = UserPtrMut::::from_usize(0xdeadbeef); let result = ptr.read_at_offset::(0); diff --git a/litebox_shim_linux/src/syscalls/net.rs b/litebox_shim_linux/src/syscalls/net.rs index e365dc110d..850e185030 100644 --- a/litebox_shim_linux/src/syscalls/net.rs +++ b/litebox_shim_linux/src/syscalls/net.rs @@ -635,10 +635,12 @@ impl GlobalState { self.net.lock().set_tcp_option( fd, match name { - "reno" | "cubic" => { - log_unsupported!("enable {} for smoltcp?", name); - return Err(Errno::EINVAL); - } + "reno" => litebox::net::TcpOptionData::CONGESTION( + litebox::net::CongestionControl::Reno, + ), + "cubic" => litebox::net::TcpOptionData::CONGESTION( + litebox::net::CongestionControl::Cubic, + ), "none" => litebox::net::TcpOptionData::CONGESTION( litebox::net::CongestionControl::None, ), @@ -651,31 +653,17 @@ impl GlobalState { } TcpOption::NODELAY | TcpOption::CORK => { let proxy = self.get_proxy(fd)?; - let is_broker_stream = match proxy.as_ref() { - NetworkProxy::Stream(_) => false, - NetworkProxy::BrokerStream(_) => true, - NetworkProxy::Datagram(_) - | NetworkProxy::Raw - | NetworkProxy::BrokerDatagram(_) => return Err(Errno::ENOPROTOOPT), + let NetworkProxy::BrokerStream(_) = proxy.as_ref() else { + return Err(Errno::ENOPROTOOPT); }; drop(proxy); let val: u32 = super::read_from_user::<_, Platform>(optval, optlen)?; - if matches!(to, TcpOption::CORK) && is_broker_stream { + if matches!(to, TcpOption::CORK) { return Err(Errno::EOPNOTSUPP); } - // Some applications use Nagle's Algorithm (via the TCP_NODELAY option) for a similar effect. - // However, TCP_CORK offers more fine-grained control, as it's designed for applications that - // send variable-length chunks of data that don't necessarily fit nicely into a full TCP segment. - // Because smoltcp does not support TCP_CORK, we emulate it by enabling/disabling Nagle's Algorithm. - let on = if let TcpOption::NODELAY = to { - val != 0 - } else { - // CORK is the opposite of NODELAY - val == 0 - }; self.net .lock() - .set_tcp_option(fd, litebox::net::TcpOptionData::NODELAY(on))?; + .set_tcp_option(fd, litebox::net::TcpOptionData::NODELAY(val != 0))?; } TcpOption::KEEPINTVL => { const MAX_TCP_KEEPINTVL: u32 = 32767; @@ -964,38 +952,6 @@ impl GlobalState { ) -> Result { let proxy = self.get_proxy(fd)?; - // Auto-bind UDP sockets if not already bound (Linux behavior: sendto() on an unbound - // UDP socket implicitly binds it to an ephemeral port before sending). - // This is mostly lock-free: we only take the network lock if we need to allocate a port. - if let NetworkProxy::Datagram(proxy) = proxy.as_ref() - && proxy.local_port() == 0 - { - // UDP socket is unbound - bind to an ephemeral port - let mut net = self.net.lock(); - // Bind with port 0 to get an ephemeral port - if let Err(err) = net.bind( - fd, - &SocketAddr::V4(core::net::SocketAddrV4::new( - core::net::Ipv4Addr::UNSPECIFIED, - 0, - )), - ) { - match err { - litebox::net::errors::BindError::AlreadyBound => { - // Another thread bound it in the meantime - that's fine - } - litebox::net::errors::BindError::InvalidFd => return Err(Errno::EBADF), - litebox::net::errors::BindError::UnsupportedAddress(_) - | litebox::net::errors::BindError::PortAlreadyInUse(_) => unreachable!(), - _ => unimplemented!(), - } - } - // Get the assigned port - let local_addr = net.get_local_addr(fd).map_err(Errno::from)?; - // If another thread already set a port, that's fine - we'll use theirs - let _ = proxy.set_local_port(local_addr.port()); - } - // Convert `SendFlags` to `litebox::net::SendFlags` // `DONTWAIT` is handled in this function and `NOSIGNAL` should be handled by caller, // so we don't convert them. @@ -1013,11 +969,8 @@ impl GlobalState { let timeout = self.with_socket_options(fd, |opt| opt.send_timeout); let is_nonblock = self.get_status(fd).contains(OFlags::NONBLOCK) || flags.contains(SendFlags::DONTWAIT); - let is_empty_stream = buf.is_empty() - && matches!( - proxy.as_ref(), - NetworkProxy::Stream(_) | NetworkProxy::BrokerStream(_) - ); + let is_empty_stream = + buf.is_empty() && matches!(proxy.as_ref(), NetworkProxy::BrokerStream(_)); cx.with_timeout(timeout) .wait_on_events( @@ -1151,12 +1104,7 @@ impl GlobalState { let wait_all = flags.contains(ReceiveFlags::WAITALL) && matches!(socket_type, SockType::Stream) && !is_nonblock; - if buf.is_empty() - && matches!( - socket.proxy.as_ref(), - NetworkProxy::Stream(_) | NetworkProxy::BrokerStream(_) - ) - { + if buf.is_empty() && matches!(socket.proxy.as_ref(), NetworkProxy::BrokerStream(_)) { return Ok(0); } @@ -2914,7 +2862,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn dropping_inet_socket_pin_reaps_deferred_close() { - let task = init_platform(None); + let task = init_platform(); let fd = task .do_socket(AddressFamily::INET, SockType::Stream, SockFlags::empty(), 0) .unwrap(); @@ -2941,7 +2889,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn socket_io_pin_keeps_backend_alive_for_send_after_close() { - let task = init_platform(None); + let task = init_platform(); let fd = task .do_socket( AddressFamily::INET, @@ -2983,7 +2931,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn raw_inet_socket_pin_does_not_follow_dup2_replacement() { - let task = init_platform(None); + let task = init_platform(); let old_fd = task .do_socket( AddressFamily::INET, @@ -3053,7 +3001,7 @@ mod tests { #[test] fn recvmmsg_lock_wakes_contended_waiter() { - let task = init_platform(None); + let task = init_platform(); let lock = alloc::sync::Arc::new(super::RecvmmsgLock::new()); let guard = lock.try_lock().unwrap(); let (started_tx, started_rx) = std::sync::mpsc::channel(); @@ -3081,7 +3029,7 @@ mod tests { #[test] fn recvmmsg_lock_honors_wait_deadline() { - let task = init_platform(None); + let task = init_platform(); let lock = super::RecvmmsgLock::new(); let guard = lock.try_lock().unwrap(); let wait_cx = task @@ -3098,7 +3046,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn inet_socket_returns_emfile_at_raw_fd_limit() { - let task = init_platform(None); + let task = init_platform(); let fd = task .do_socket(AddressFamily::INET, SockType::Stream, SockFlags::empty(), 0) .unwrap(); @@ -3356,7 +3304,7 @@ mod tests { } fn test_tcp_socket_with_external_client(is_nonblocking: bool, test_trunc: bool) { - let task = init_platform(None); + let task = init_platform(); test_tcp_socket_as_server( &task, LOOPBACK_IP_ADDR, @@ -3376,7 +3324,7 @@ mod tests { } fn test_tcp_socket_send(is_nonblocking: bool, test_trunc: bool) { - let task = init_platform(None); + let task = init_platform(); test_tcp_socket_as_server( &task, LOOPBACK_IP_ADDR, @@ -3428,7 +3376,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_tcp_connection_refused() { - let task = init_platform(None); + let task = init_platform(); let port = find_free_tcp_port(); let socket_fd = task .do_socket(AddressFamily::INET, SockType::Stream, SockFlags::empty(), 0) @@ -3459,7 +3407,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_tcp_socket_as_client() { - let task = init_platform(None); + let task = init_platform(); let port = find_free_tcp_port(); let child_handle = std::thread::spawn(move || { @@ -3656,7 +3604,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_blocking_udp_server_socket() { - let task = init_platform(None); + let task = init_platform(); blocking_udp_server_socket(&task, false, false, false, "recvfrom"); blocking_udp_server_socket(&task, false, false, false, "recvmsg"); } @@ -3664,7 +3612,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_nonblocking_udp_server_socket() { - let task = init_platform(None); + let task = init_platform(); blocking_udp_server_socket(&task, false, false, true, "recvfrom"); blocking_udp_server_socket(&task, false, false, true, "recvmsg"); } @@ -3672,7 +3620,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_blocking_udp_server_socket_with_truncation() { - let task = init_platform(None); + let task = init_platform(); blocking_udp_server_socket(&task, true, true, false, "recvfrom"); blocking_udp_server_socket(&task, true, true, false, "recvmsg"); blocking_udp_server_socket(&task, true, false, false, "recvmsg"); @@ -3681,7 +3629,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_udp_client_socket_without_server() { - let task = init_platform(None); + let task = init_platform(); let server_port = find_free_udp_port(); // Client socket and explicit bind @@ -3732,7 +3680,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_tcp_keepalive_sockopt() { - let task = init_platform(None); + let task = init_platform(); let sockfd = task .do_socket(AddressFamily::INET, SockType::Stream, SockFlags::empty(), 0) .expect("failed to create socket"); @@ -3766,7 +3714,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_socket_dup_and_close() { - let task = init_platform(None); + let task = init_platform(); let socket_fd = task .do_socket( litebox_common_linux::AddressFamily::INET, @@ -3852,7 +3800,7 @@ mod unix_tests { #[test] fn test_unix_datagram_socket() { - let task = init_platform(None); + let task = init_platform(); for _ in 0..10 { let server_path = "/unix_stream_socket_server.sock"; @@ -3929,7 +3877,7 @@ mod unix_tests { #[test] fn test_unix_stream_socket() { - let task = init_platform(None); + let task = init_platform(); for _ in 0..10 { let addr = "/unix_stream_socket.sock"; @@ -3985,7 +3933,7 @@ mod unix_tests { #[test] fn test_unix_stream_socket_refused() { - let task = init_platform(None); + let task = init_platform(); let client_fd = create_unix_socket(&task, SockType::Stream, SockFlags::empty()); let addr = "/unix_stream_socket_refused.sock"; let result = task.do_connect( @@ -4033,7 +3981,7 @@ mod unix_tests { } fn test_multiple_unix_stream_connections(is_nonblocking: bool) { - let task = init_platform(None); + let task = init_platform(); let addr = "/unix_multi_stream_socket.sock"; let server_fd = create_unix_server_socket( &task, @@ -4132,7 +4080,7 @@ mod unix_tests { #[test] fn test_unix_stream_socket_on_same_addr() { - let task = init_platform(None); + let task = init_platform(); for _ in 0..10 { let addr = "/unix_stream_socket_server.sock"; let server1_fd = create_unix_server_socket(&task, addr, SockFlags::NONBLOCK).unwrap(); @@ -4183,7 +4131,7 @@ mod unix_tests { #[test] fn test_unix_datagram_socket_on_same_addr() { - let task = init_platform(None); + let task = init_platform(); for _ in 0..10 { let addr = "/unix_datagram_socket_server.sock"; let server_fd = create_unix_socket(&task, SockType::Datagram, SockFlags::empty()); @@ -4217,7 +4165,7 @@ mod unix_tests { } fn unix_socketpair_bidirectional(ty: SockType, is_nonblocking: bool) { - let task = init_platform(None); + let task = init_platform(); let mut sv_ptr = alloc::vec![0u32; 2]; let sv_mut_ptr = UserPtrMut::from_usize(sv_ptr.as_mut_ptr() as usize); @@ -4284,7 +4232,7 @@ mod unix_tests { #[test] fn pinned_receive_does_not_follow_dup2_replacement() { - let task = init_platform(None); + let task = init_platform(); let (old_sender, old_receiver) = task .do_socketpair(AddressFamily::UNIX, SockType::Stream, SockFlags::empty(), 0) .unwrap(); @@ -4333,7 +4281,7 @@ mod unix_tests { } fn unix_socket_recv_timeout(ty: SockType) { - let task = init_platform(None); + let task = init_platform(); let (sock1, _sock2) = task .do_socketpair(AddressFamily::UNIX, ty, SockFlags::empty(), 0) .expect("socketpair failed"); @@ -4370,7 +4318,7 @@ mod unix_tests { #[test] fn test_unix_stream_addr() { - let task = init_platform(None); + let task = init_platform(); let server_path = "/unix_stream_sockname.sock"; let server_fd = create_unix_server_socket(&task, server_path, SockFlags::empty()).unwrap(); @@ -4438,7 +4386,7 @@ mod unix_tests { #[test] fn test_unix_datagram_addr() { - let task = init_platform(None); + let task = init_platform(); let server_path = "/unix_datagram_sockname_server.sock"; let client_path = "/unix_datagram_sockname_client.sock"; diff --git a/litebox_shim_linux/src/syscalls/process.rs b/litebox_shim_linux/src/syscalls/process.rs index 6024e091a4..af8531ab3a 100644 --- a/litebox_shim_linux/src/syscalls/process.rs +++ b/litebox_shim_linux/src/syscalls/process.rs @@ -1654,7 +1654,7 @@ mod tests { use crate::syscalls::tests::init_platform; use litebox_common_linux::ArchPrctlArg; - let task = init_platform(None); + let task = init_platform(); // Save old FS base let mut old_fs_base: usize = 0; @@ -1683,7 +1683,7 @@ mod tests { #[test] fn test_sched_getaffinity() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); let cpuset = task.sys_sched_getaffinity(None); assert_eq!(cpuset.bits.len(), super::NR_CPUS); @@ -1698,7 +1698,7 @@ mod tests { #[test] fn test_prctl_set_get_name() { - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); // Prepare a null-terminated name to set let name: &[u8] = b"litebox-test\0"; @@ -1754,7 +1754,7 @@ mod tests { use litebox_common_linux::{ClockId, TimerFlags, Timespec}; let callback_addr = 0x1000usize; // dummy non-null address for the callback - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); ::run_test_thread(|| { let act = SigAction { sigaction: callback_addr, @@ -1823,7 +1823,7 @@ mod tests { use litebox::platform::{Instant as _, TimeProvider}; use litebox_common_linux::{ClockId, TimerFlags, Timespec}; - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); ::run_test_thread(|| { let platform = task.global.platform; @@ -1882,7 +1882,7 @@ mod tests { fn test_alarm_cancel_prevents_signal() { use litebox_common_linux::{ClockId, TimerFlags, Timespec}; - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); ::run_test_thread(|| { assert_eq!(task.sys_alarm(1).unwrap(), 0); // Cancel before it fires. @@ -1918,7 +1918,7 @@ mod tests { signal::{SigSet, SigmaskHow, Signal}, }; - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); ::run_test_thread(|| { let block_set = SigSet::empty().with(Signal::SIGUSR1); task.sys_rt_sigprocmask( @@ -1965,7 +1965,7 @@ mod tests { use litebox_common_linux::signal::{SIG_IGN, SaFlags, SigAction, SigSet, Signal}; use litebox_common_linux::{ClockId, TimerFlags, Timespec}; - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); ::run_test_thread(|| { // Install SIG_IGN for SIGALRM. let act = SigAction { @@ -2021,7 +2021,7 @@ mod tests { use litebox_common_linux::signal::Signal; use litebox_common_linux::{ClockId, TimerFlags, Timespec}; - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); ::run_test_thread(|| { let platform = task.global.platform; diff --git a/litebox_shim_linux/src/syscalls/tests.rs b/litebox_shim_linux/src/syscalls/tests.rs index 0f0813250a..78fd8be1f0 100644 --- a/litebox_shim_linux/src/syscalls/tests.rs +++ b/litebox_shim_linux/src/syscalls/tests.rs @@ -22,27 +22,23 @@ pub(crate) use litebox_platform_linux_userland::LinuxUserland as TestPlatform; pub(crate) use litebox_platform_windows_userland::WindowsUserland as TestPlatform; /// Returns the process-wide test platform, initializing it once. -pub(crate) fn test_platform(tun_device_name: Option<&str>) -> &'static TestPlatform { +pub(crate) fn test_platform() -> &'static TestPlatform { static PLATFORM: std::sync::OnceLock<&'static TestPlatform> = std::sync::OnceLock::new(); PLATFORM.get_or_init(|| { - // Only the Linux userland platform takes a tun device name. #[cfg(target_os = "linux")] { - TestPlatform::new(tun_device_name) + TestPlatform::new() } #[cfg(target_os = "windows")] { - let _ = tun_device_name; TestPlatform::new() } }) } #[must_use] -pub(crate) fn init_platform( - tun_device_name: Option<&str>, -) -> crate::Task> { - let platform = test_platform(tun_device_name); +pub(crate) fn init_platform() -> crate::Task> { + let platform = test_platform(); let shim_builder = crate::LinuxShimBuilder::new(platform); let litebox = shim_builder.litebox(); @@ -52,30 +48,12 @@ pub(crate) fn init_platform( .expect("Failed to set permissions on root"); }); let fs = alloc::sync::Arc::new(shim_builder.default_fs(in_mem_fs, TEST_TAR_FILE.into())); - let task = shim_builder.build().0.new_test_task(fs); - - if tun_device_name.is_some() { - let global = task.global.clone(); - // Start a background thread to perform network interaction - // Naive implementation for testing purpose only - std::thread::spawn(move || { - loop { - while global - .net - .lock() - .perform_platform_interaction() - .call_again_immediately() - {} - core::hint::spin_loop(); - } - }); - } - task + shim_builder.build().0.new_test_task(fs) } #[test] fn test_fcntl() { - let task = init_platform(None); + let task = init_platform(); let check = |fd: i32, flags1: OFlags, flags2: OFlags| { assert_eq!( @@ -135,7 +113,7 @@ fn test_fcntl() { #[test] fn test_dup() { - let task = init_platform(None); + let task = init_platform(); let fd = task .sys_open("/dev/stdin", OFlags::RDONLY, Mode::empty()) @@ -166,7 +144,7 @@ fn test_dup() { // Note the test was generated by copilot with minor fixes. #[test] fn test_getdent64() { - let task = init_platform(None); + let task = init_platform(); // Create test files in root directory for testing let file1_fd = task @@ -437,7 +415,7 @@ fn test_getdent64() { #[test] fn test_umask_behavior() { - let task = init_platform(None); + let task = init_platform(); // 1. Capture original mask without changing final state. let orig = task.sys_umask(0).bits(); // sets mask to 0, returns previous @@ -498,7 +476,7 @@ fn test_umask_behavior() { fn test_rlimit_nofile() { use litebox_common_linux::{Rlimit, RlimitResource, errno::Errno}; - let task = crate::syscalls::tests::init_platform(None); + let task = crate::syscalls::tests::init_platform(); // 1. Get the current NOFILE limit. let cur_lim = task @@ -548,7 +526,7 @@ fn test_rlimit_nofile() { #[test] fn test_unlinkat() { - let task = init_platform(None); + let task = init_platform(); // 1. Create a regular file and unlink it. let file_path = "/unlink_test_file.txt"; @@ -655,7 +633,7 @@ fn test_rwlock_readers_not_starved_after_writer_handoff() { } // Initialize the platform (reuses the global Once-based init). - let _task = init_platform(None); + let _task = init_platform(); let join_timeout = std::time::Duration::from_secs(5); // We run the test many times to increase the probability of hitting the diff --git a/litebox_shim_linux/src/transport.rs b/litebox_shim_linux/src/transport.rs index 7be8c91da6..e7c168e479 100644 --- a/litebox_shim_linux/src/transport.rs +++ b/litebox_shim_linux/src/transport.rs @@ -285,7 +285,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_nine_p_create_and_read_file() { - let task = init_platform(None); + let task = init_platform(); let server = DiodServer::start(); let fs = connect_9p(&task, &server); @@ -320,7 +320,7 @@ mod tests { #[test] #[ignore = "requires broker-backed socket test setup"] fn test_nine_p_host_files_visible() { - let task = init_platform(None); + let task = init_platform(); let server = DiodServer::start(); diff --git a/litebox_shim_optee/src/syscalls/tests.rs b/litebox_shim_optee/src/syscalls/tests.rs index 4412289188..99cd2ed075 100644 --- a/litebox_shim_optee/src/syscalls/tests.rs +++ b/litebox_shim_optee/src/syscalls/tests.rs @@ -7,14 +7,10 @@ use litebox_platform_multiplex::{Platform, set_platform}; static INIT_FUNC: spin::Once = spin::Once::new(); #[must_use] -#[cfg_attr( - not(target_os = "linux"), - expect(unused_variables, reason = "ignored parameter on non-linux platforms") -)] pub(crate) fn init_platform() -> crate::Task { INIT_FUNC.call_once(|| { #[cfg(target_os = "linux")] - let platform = Platform::new(None); + let platform = Platform::new(); #[cfg(not(target_os = "linux"))] let platform = Platform::new(); diff --git a/litebox_shim_windows/src/tests.rs b/litebox_shim_windows/src/tests.rs index 065468c5be..c654bc8224 100644 --- a/litebox_shim_windows/src/tests.rs +++ b/litebox_shim_windows/src/tests.rs @@ -73,7 +73,7 @@ pub(crate) fn test_platform() -> &'static TestPlatform { static PLATFORM: std::sync::OnceLock<&'static TestPlatform> = std::sync::OnceLock::new(); PLATFORM.get_or_init(|| { #[cfg(target_os = "linux")] - let platform = TestPlatform::new(None); + let platform = TestPlatform::new(); #[cfg(target_os = "windows")] let platform = TestPlatform::new();