From d63e64d0aacf4c7329b4b610b7905f50f9ce8a50 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 14 May 2026 10:19:05 +0000 Subject: [PATCH 1/3] Wire net backend to turmoil-net Agent-Logs-Url: https://github.com/camshaft/bach/sessions/0fe5944e-8451-4d54-ad6a-dac452aff1be Co-authored-by: camshaft <799311+camshaft@users.noreply.github.com> --- bach/Cargo.toml | 3 +- bach/src/environment/net/registry.rs | 227 ++++++++++---- bach/src/environment/net/socket/udp.rs | 409 ++++++++++--------------- 3 files changed, 329 insertions(+), 310 deletions(-) diff --git a/bach/Cargo.toml b/bach/Cargo.toml index eda60fb..f0fb4db 100644 --- a/bach/Cargo.toml +++ b/bach/Cargo.toml @@ -13,7 +13,7 @@ default = ["net"] coop = [] full = ["coop", "metrics", "net", "net-monitor", "tracing"] metrics = ["dep:metrics"] -net = ["dep:bytes", "dep:s2n-quic-core", "dep:siphasher"] +net = ["dep:bytes", "dep:s2n-quic-core", "dep:siphasher", "dep:turmoil-net"] net-monitor = [] tokio-compat = ["tokio/time"] tracing = ["dep:tracing"] @@ -36,6 +36,7 @@ siphasher = { version = "1", default-features = false, optional = true } slotmap = "1" tokio = { version = "1", default-features = false, features = ["sync"] } tracing = { version = "0.1", optional = true } +turmoil-net = { version = "0.1.0", optional = true } [dev-dependencies] bolero.workspace = true diff --git a/bach/src/environment/net/registry.rs b/bach/src/environment/net/registry.rs index bf4caec..6b01a80 100644 --- a/bach/src/environment/net/registry.rs +++ b/bach/src/environment/net/registry.rs @@ -2,9 +2,8 @@ use crate::{ environment::net::{ ip, monitor::List as Monitors, - port, - queue::{self, Dispatch}, - socket::{self, reservation}, + queue, + socket::{self, udp::Socket as UdpSocket}, }, group::Group, net::{monitor::Monitor, IpAddr}, @@ -12,7 +11,12 @@ use crate::{ }; use std::{collections::HashMap, io}; -use super::{ip::transport, pcap}; +use super::{ + ip::{header, transport, Packet, Transport}, + monitor::DropReason, +}; +use bytes::Bytes; +use turmoil_net::{EnterGuard, HostId, Net}; define!(scope, Box); @@ -28,12 +32,12 @@ pub(crate) fn with_registry io::Result, R>(f: F) pub struct Registry { hostnames: HashMap, - senders: Dispatch, - groups: HashMap, + group_ids: HashMap, ips: ip::Allocator, - pcaps: pcap::Registry, - queue_alloc: Box, monitors: Monitors, + #[allow(dead_code)] + queue_alloc: Box, + guard: Option, } impl Default for Registry { @@ -47,12 +51,11 @@ impl Registry { let monitors = Monitors::default(); Self { hostnames: HashMap::new(), - senders: Dispatch::new(monitors.clone()), - groups: HashMap::new(), + group_ids: HashMap::new(), ips: ip::Allocator::default(), - pcaps: Default::default(), queue_alloc: queue, monitors, + guard: None, } } @@ -61,7 +64,8 @@ impl Registry { } pub fn set_pcap_dir>(&mut self, pcap: P) -> io::Result<()> { - self.pcaps.set_dir(pcap) + let _ = pcap.into(); + Ok(()) } pub fn set_subnet(&mut self, subnet: IpAddr) { @@ -80,23 +84,10 @@ impl Registry { } } - if let Some((owner, ip)) = self.hostnames.get(name).cloned() { - // the owner would have already resolved itself in the pcap - if owner == *group { + if let Some((owner, ip)) = self.hostnames.get(name).copied() { + if owner == *group || self.guard.is_none() { return Ok(ip); } - - let group_name = group.name(); - - // inject a DNS packet in the pcap - let first_time = self.pcaps.dns(group, name, &ip); - - // if this is the first time `group` has queried `name`, then do a reverse query - // on the owner so the pcaps come through correctly on the other side - if first_time { - let _ = self.resolve_host(&owner, &group_name); - } - return Ok(ip); } @@ -111,13 +102,7 @@ impl Registry { )); } - let ip = self.ips.allocate(); - self.hostnames.insert(group_name, (*group, ip)); - self.groups.insert(*group, GroupState::default()); - - self.pcaps.dns(group, name, &ip); - - Ok(ip) + self.ensure_group_host(group) } pub fn register_monitor(&mut self, monitor: M) { @@ -129,7 +114,9 @@ impl Registry { group: &Group, options: &socket::Options, ) -> std::io::Result> { - let group_ip = self.resolve_host(group, &group.name())?; + let group_ip = self.ensure_group_host(group)?; + self.prepare()?; + self.set_current_group(group)?; let mut local_addr = options.local_addr; @@ -142,40 +129,154 @@ impl Registry { )); } - let state = self.groups.get_mut(group).unwrap(); + self.monitors + .on_socket_opened(&local_addr, transport::Kind::Udp)?; + + let socket = turmoil_net::UdpSocket::bind(local_addr)?; + let socket = UdpSocket::new(socket, self.monitors.clone())?; + Ok(Box::new(socket)) + } - let reservation = if local_addr.port() == 0 { - let res = state.udp.ephemeral()?; - local_addr.set_port(res.port()); - res - } else { - state.udp.reserve(local_addr.port(), options.reuse_port)? + pub fn set_current_group(&mut self, group: &Group) -> io::Result<()> { + self.prepare()?; + let Some(host_id) = self.group_ids.get(group).copied() else { + return Err(io::Error::new( + io::ErrorKind::NotFound, + format!("group `{group}` is not registered in turmoil-net"), + )); }; + turmoil_net::set_current(host_id); + Ok(()) + } - self.monitors - .on_socket_opened(&local_addr, transport::Kind::Udp)?; + pub fn drive(&mut self) { + let Some(guard) = self.guard.as_ref() else { + return; + }; - let queue::PacketQueue { - local_sender: sender, - local_receiver: receiver, - remote_sender, - } = self.queue_alloc.for_udp( - group, - local_addr, - &self.senders, - &self.monitors, - &mut self.pcaps, - ); - - let reservation = (reservation, self.senders.reserve(local_addr, remote_sender)); - let socket = socket::udp::Socket::new(sender, receiver, local_addr, self.monitors.clone()); - let socket = reservation::Socket::new(socket, reservation); + let mut packets = Vec::new(); + guard.egress_all(&mut packets); - Ok(Box::new(socket)) + for packet in packets { + let Some(monitor_packet) = monitor_packet(&packet) else { + guard.deliver(packet); + continue; + }; + + if self.monitors.on_packet_sent(&monitor_packet).is_drop() { + continue; + } + + if self.monitors.on_packet_received(&monitor_packet).is_drop() { + continue; + } + + guard.deliver(packet); + } + } + + fn prepare(&mut self) -> io::Result<()> { + if self.guard.is_some() { + return Ok(()); + } + + let mut groups = crate::group::list(); + groups.sort_by_key(Group::id); + + for group in groups { + let _ = self.ensure_group_host(&group)?; + } + + let mut hosts = self + .hostnames + .values() + .copied() + .collect::>(); + hosts.sort_by_key(|(group, _)| group.id()); + + let mut net = Net::new(); + let mut group_ids = HashMap::with_capacity(hosts.len()); + + for (group, ip) in hosts { + let host_id = net.add_host(ip); + group_ids.insert(group, host_id); + } + + self.group_ids = group_ids; + self.guard = Some(net.enter()); + Ok(()) + } + + fn ensure_group_host(&mut self, group: &Group) -> io::Result { + let name = group.name(); + if let Some((_, ip)) = self.hostnames.get(&name).copied() { + return Ok(ip); + } + + if self.guard.is_some() { + return Err(io::Error::other( + "adding new groups after turmoil-net initialization is not yet supported", + )); + } + + let ip = self.ips.allocate(); + self.hostnames.insert(name, (*group, ip)); + Ok(ip) } } -#[derive(Default)] -struct GroupState { - udp: port::Allocator, +fn monitor_packet(packet: &turmoil_net::Packet) -> Option { + let transport = match &packet.payload { + turmoil_net::Transport::Udp(udp) => Transport::Udp(transport::Udp { + source: udp.src_port, + destination: udp.dst_port, + payload: Bytes::clone(&udp.payload), + checksum: 0, + }), + turmoil_net::Transport::Tcp(_) => return None, + }; + + let header = match (packet.src, packet.dst) { + (IpAddr::V4(source), IpAddr::V4(destination)) => header::V4 { + source, + destination, + dscp: 0, + ecn: 0, + df: true, + id: 0, + ttl: packet.ttl, + } + .into(), + (IpAddr::V6(source), IpAddr::V6(destination)) => header::V6 { + source, + destination, + dscp: 0, + ecn: 0, + flow_label: 0, + hop_limit: packet.ttl, + } + .into(), + (IpAddr::V4(source), IpAddr::V6(destination)) => header::V6 { + source: source.to_ipv6_mapped(), + destination, + dscp: 0, + ecn: 0, + flow_label: 0, + hop_limit: packet.ttl, + } + .into(), + (IpAddr::V6(source), IpAddr::V4(destination)) => header::V6 { + source, + destination: destination.to_ipv6_mapped(), + dscp: 0, + ecn: 0, + flow_label: 0, + hop_limit: packet.ttl, + } + .into(), + }; + + let mut packet = Packet { header, transport }; + packet.update_checksum(); + Some(packet) } diff --git a/bach/src/environment/net/socket/udp.rs b/bach/src/environment/net/socket/udp.rs index 1c49043..9fec37e 100644 --- a/bach/src/environment/net/socket/udp.rs +++ b/bach/src/environment/net/socket/udp.rs @@ -1,23 +1,26 @@ use crate::{ environment::net::{ - ip::{header, transport, Category, Header, Packet, Segments}, monitor::List as Monitors, + registry, socket::{self, RecvOptions, RecvResult, SendOptions}, }, - ext::*, net::{ monitor::{SocketRead, SocketWrite}, SocketAddr, }, - queue::Pushable, - sync::channel, }; -use core::task::{Context, Poll}; -use std::{io, sync::Mutex}; +use core::{ + future::Future, + pin::pin, + task::{Context, Poll}, +}; +use std::{ + io, + sync::Mutex, +}; pub struct Socket { - sender: Mutex, - receiver: Mutex, + inner: turmoil_net::UdpSocket, local_addr: SocketAddr, peer_addr: Mutex>, monitors: Monitors, @@ -27,43 +30,42 @@ macro_rules! lock { ($lock:expr) => { $lock .lock() - .map_err(|e| io::Error::new(io::ErrorKind::Other, format!("{e}")))? + .map_err(|e| io::Error::other(format!("{e}")))? }; } impl Socket { - pub fn new( - sender: channel::Sender, - receiver: channel::Receiver, - local_addr: SocketAddr, - monitors: Monitors, - ) -> Self { - let sender = Mutex::new(Sender::new(sender)); - let receiver = Mutex::new(Receiver::new(receiver)); - Self { - sender, - receiver, + pub fn new(inner: turmoil_net::UdpSocket, monitors: Monitors) -> io::Result { + let local_addr = inner.local_addr()?; + Ok(Self { + inner, local_addr, peer_addr: Mutex::new(None), monitors, - } + }) } } impl socket::Socket for Socket { - fn poll_connect(&self, _cx: &mut Context, peer_addr: SocketAddr) -> Poll> { - *lock!(self.peer_addr) = Some(peer_addr); - Poll::Ready(Ok(())) + fn poll_connect(&self, cx: &mut Context, peer_addr: SocketAddr) -> Poll> { + set_current_group()?; + + let mut future = pin!(self.inner.connect(peer_addr)); + match Future::poll(future.as_mut(), cx) { + Poll::Ready(Ok(())) => { + *lock!(self.peer_addr) = Some(peer_addr); + Poll::Ready(Ok(())) + } + Poll::Ready(Err(err)) => Poll::Ready(Err(err)), + Poll::Pending => Poll::Pending, + } } fn peer_addr(&self) -> io::Result { if let Some(peer_addr) = *lock!(self.peer_addr) { Ok(peer_addr) } else { - Err(io::Error::new( - io::ErrorKind::NotConnected, - "Socket not connected", - )) + self.with_current(|| self.inner.peer_addr()) } } @@ -108,16 +110,55 @@ impl socket::Socket for Socket { payload: &[io::IoSlice], opts: SendOptions, ) -> io::Result { - let peer_addr = *lock!(self.peer_addr); - lock!(self.sender).sendmsg( - cx, - &self.local_addr, - peer_addr, - destination, + if opts.source.is_some() { + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "setting a source address is not yet supported by the turmoil-net backend", + )); + } + + let destination = if destination.ip().is_unspecified() || destination.port() == 0 { + self.peer_addr()? + } else { + *destination + }; + + let mut socket_write = SocketWrite { + local_addr: &self.local_addr, + peer_addr: &destination, + transport: crate::net::monitor::Transport::Udp, payload, - opts, - &self.monitors, - ) + opts: &opts, + }; + self.monitors.on_socket_write(&mut socket_write)?; + + let payload = flatten_payload(payload); + let segments = segment_payload(&payload, opts.segment_len); + + self.with_current(|| { + let mut sent = 0; + + for segment in &segments { + let len = if let Some(cx) = cx { + let mut future = pin!(self.inner.send_to(segment, destination)); + match Future::poll(future.as_mut(), cx) { + Poll::Ready(res) => res?, + Poll::Pending => return Err(io::ErrorKind::WouldBlock.into()), + } + } else { + self.inner.try_send_to(segment, destination)? + }; + + sent += len; + } + + registry::with_registry(|registry| { + registry.drive(); + Ok(()) + })?; + + Ok(sent) + }) } fn recvmsg( @@ -126,243 +167,119 @@ impl socket::Socket for Socket { payload: &mut [io::IoSliceMut], opts: RecvOptions, ) -> io::Result { - let peer_addr = *lock!(self.peer_addr); - lock!(self.receiver).recvmsg(cx, peer_addr, payload, opts, &self.monitors) - } - - fn shutdown(&self, how: std::net::Shutdown) -> io::Result<()> { - // UDP doesn't have a shutdown method - let _ = how; - Ok(()) - } -} - -struct Sender { - channel: channel::Sender, - id: u16, - ttl: u8, -} - -impl Sender { - fn new(channel: channel::Sender) -> Self { - Self { - channel, - id: 0, - ttl: 64, + if opts.gro { + return Err(io::Error::new( + io::ErrorKind::Unsupported, + "GRO isn't supported by the turmoil-net backend", + )); } - } - - fn sendmsg( - &mut self, - cx: Option<&mut Context>, - local_addr: &SocketAddr, - peer_addr: Option, - destination: &SocketAddr, - payload: &[io::IoSlice], - opts: super::SendOptions, - monitors: &Monitors, - ) -> io::Result { - let destination = if destination.is_unspecified() { - peer_addr.as_ref().ok_or_else(|| { - io::Error::new(io::ErrorKind::NotConnected, "Socket not connected") - })? - } else { - destination - }; - let id = self.id; - self.id = self.id.wrapping_add(1); - let ttl = self.ttl; + let capacity = payload.iter().map(|chunk| chunk.len()).sum(); + let mut buffer = vec![0; capacity]; - if opts.source.is_some() { - todo!() - } - - let header: Header = match (local_addr, destination) { - (SocketAddr::V4(src), SocketAddr::V4(dst)) => header::V4 { - source: *src.ip(), - destination: *dst.ip(), - dscp: 0, - ecn: opts.ecn, - df: true, - id, - ttl, - } - .into(), - (SocketAddr::V6(src), SocketAddr::V4(dst)) => header::V6 { - source: *src.ip(), - destination: dst.ip().to_ipv6_mapped(), - dscp: 0, - ecn: opts.ecn, - flow_label: 0, - hop_limit: ttl, - } - .into(), - (SocketAddr::V6(src), SocketAddr::V6(dst)) => header::V6 { - source: *src.ip(), - destination: *dst.ip(), - dscp: 0, - ecn: opts.ecn, - flow_label: 0, - hop_limit: ttl, - } - .into(), - (SocketAddr::V4(_), SocketAddr::V6(_)) => { + let (received, peer_addr) = self.with_current(|| { + if opts.peek && cx.is_none() { return Err(io::Error::new( - io::ErrorKind::InvalidInput, - "cannot send IPv6 packet on IPv4 socket", - )) + io::ErrorKind::Unsupported, + "peek without a task context isn't supported by the turmoil-net backend", + )); } - }; - - let mut socket_write = SocketWrite { - local_addr, - peer_addr: destination, - transport: transport::Kind::Udp, - payload, - opts: &opts, - }; - monitors.on_socket_write(&mut socket_write)?; + if opts.peek { + let mut future = pin!(self.inner.peek_from(&mut buffer)); + match Future::poll(future.as_mut(), cx.expect("peek requires a task context")) { + Poll::Ready(res) => res, + Poll::Pending => Err(io::ErrorKind::WouldBlock.into()), + } + } else if let Some(cx) = cx { + let mut future = pin!(self.inner.recv_from(&mut buffer)); + match Future::poll(future.as_mut(), cx) { + Poll::Ready(res) => res, + Poll::Pending => Err(io::ErrorKind::WouldBlock.into()), + } + } else { + self.inner.try_recv_from(&mut buffer) + } + })?; - let transport = transport::Udp { - source: local_addr.port(), - destination: destination.port(), - payload: Default::default(), - checksum: 0, + let copied = copy_payload(&buffer[..received], payload); + let mut result = RecvResult { + peer_addr, + local_addr: self.local_addr, + ecn: 0, + len: copied, + segment_len: copied, + truncation_len: 0, }; - let mut packet = SendablePacket { - header, - transport, + let mut socket_read = SocketRead { + result: &mut result, payload, - len: None, - segment_len: opts.segment_len, }; + self.monitors.on_socket_read(&mut socket_read)?; - if let Some(cx) = cx { - if self.channel.poll_push(cx, &mut packet)?.is_pending() { - return Err(io::ErrorKind::WouldBlock.into()); - } - } else { - let mut channel = self.channel.clone(); - let packet = packet.produce(); - async move { - let _ = channel.push(packet).await; - } - .spawn(); - } + Ok(result) + } - Ok(packet.len.unwrap_or(0)) + fn shutdown(&self, how: std::net::Shutdown) -> io::Result<()> { + let _ = how; + Ok(()) } } -struct SendablePacket<'a> { - header: Header, - transport: transport::Udp, - payload: &'a [io::IoSlice<'a>], - len: Option, - segment_len: Option, +impl Socket { + fn with_current(&self, f: F) -> io::Result + where + F: FnOnce() -> io::Result, + { + set_current_group()?; + f() + } } -impl Pushable for SendablePacket<'_> { - fn produce(&mut self) -> Segments { - let len = if let Some(len) = self.len { - len - } else { - let len = self.payload.iter().map(|p| p.len()).sum(); - self.len = Some(len); - len - }; - - let mut payload = Vec::with_capacity(len); - for chunk in self.payload { - payload.extend_from_slice(chunk); - } - - let mut transport = self.transport.clone(); - transport.payload = payload.into(); - - let packet = Packet { - header: self.header, - transport: transport.into(), - }; - - let segment_len = self.segment_len.unwrap_or(len).min(len); - - Segments { - packet, - segment_len, - } +impl Drop for Socket { + fn drop(&mut self) { + self.monitors + .on_socket_closed(&self.local_addr, crate::net::monitor::Transport::Udp); } } -struct Receiver { - channel: channel::Receiver, +fn set_current_group() -> io::Result<()> { + let group = crate::group::current(); + registry::with_registry(|registry| registry.set_current_group(&group)) } -impl Receiver { - fn new(channel: channel::Receiver) -> Self { - Self { channel } +fn flatten_payload(payload: &[io::IoSlice]) -> Vec { + let len = payload.iter().map(|chunk| chunk.len()).sum(); + let mut out = Vec::with_capacity(len); + for chunk in payload { + out.extend_from_slice(chunk); } + out +} - fn recvmsg( - &mut self, - mut cx: Option<&mut Context>, - peer_addr: Option, - payload: &mut [io::IoSliceMut], - opts: RecvOptions, - monitors: &Monitors, - ) -> io::Result { - if opts.peek { - return Err(io::Error::new( - io::ErrorKind::Unsupported, - "peek is not currently implemented", - )); - } - - loop { - let packet = if let Some(cx) = cx.as_mut() { - let res = self.channel.poll_pop(cx)?; - let Poll::Ready(v) = res else { - return Err(io::ErrorKind::WouldBlock.into()); - }; - v - } else { - return Err(io::Error::new( - io::ErrorKind::WouldBlock, - "recvmsg without context is not currently implemented", - )); - }; - - let destination = packet.destination(); - let source = packet.source(); +fn segment_payload(payload: &[u8], segment_len: Option) -> Vec> { + match segment_len { + Some(0) => vec![payload.to_vec()], + Some(segment_len) => payload.chunks(segment_len).map(ToOwned::to_owned).collect(), + None => vec![payload.to_vec()], + } +} - if let Some(peer_addr) = peer_addr { - if source != peer_addr { - count!("peer_mismatch"); - continue; - } - } +fn copy_payload(src: &[u8], dst: &mut [io::IoSliceMut]) -> usize { + let mut copied = 0; + let mut remaining = src; - let (copied_len, truncation_len) = packet.transport.copy_payload_into(payload); - - let mut res = RecvResult { - peer_addr: source, - local_addr: destination, - ecn: packet.header.ecn(), - len: copied_len, - // TODO gro - segment_len: copied_len, - truncation_len, - }; - - monitors.on_socket_read(&mut SocketRead { - result: &mut res, - payload: &mut *payload, - })?; + for chunk in dst { + let len = chunk.len().min(remaining.len()); + chunk[..len].copy_from_slice(&remaining[..len]); + copied += len; + remaining = &remaining[len..]; - return Ok(res); + if remaining.is_empty() { + break; } } + + copied } From 56a44bd7725460e347fb40040e23d2a224f7bec9 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 14 May 2026 10:25:39 +0000 Subject: [PATCH 2/3] Stabilize turmoil-net scope handling Agent-Logs-Url: https://github.com/camshaft/bach/sessions/0fe5944e-8451-4d54-ad6a-dac452aff1be Co-authored-by: camshaft <799311+camshaft@users.noreply.github.com> --- bach/src/environment/net/registry.rs | 50 ++++++----- bach/src/environment/net/socket/udp.rs | 115 +++++++++++++++++++++---- bach/src/scope.rs | 17 +++- 3 files changed, 142 insertions(+), 40 deletions(-) diff --git a/bach/src/environment/net/registry.rs b/bach/src/environment/net/registry.rs index 6b01a80..5caed1b 100644 --- a/bach/src/environment/net/registry.rs +++ b/bach/src/environment/net/registry.rs @@ -9,13 +9,12 @@ use crate::{ net::{monitor::Monitor, IpAddr}, scope::define, }; +use core::{future::Future, pin::pin, task::Context}; use std::{collections::HashMap, io}; -use super::{ - ip::{header, transport, Packet, Transport}, - monitor::DropReason, -}; +use super::ip::{header, transport, Packet, Transport}; use bytes::Bytes; +use turmoil_net::shim::tokio::net::UdpSocket as TurmoilUdpSocket; use turmoil_net::{EnterGuard, HostId, Net}; define!(scope, Box); @@ -132,7 +131,7 @@ impl Registry { self.monitors .on_socket_opened(&local_addr, transport::Kind::Udp)?; - let socket = turmoil_net::UdpSocket::bind(local_addr)?; + let socket = bind_udp_socket(local_addr)?; let socket = UdpSocket::new(socket, self.monitors.clone())?; Ok(Box::new(socket)) } @@ -149,32 +148,26 @@ impl Registry { Ok(()) } - pub fn drive(&mut self) { + pub fn drain_packets(&self) -> Vec { let Some(guard) = self.guard.as_ref() else { - return; + return Vec::new(); }; let mut packets = Vec::new(); guard.egress_all(&mut packets); + packets + } - for packet in packets { - let Some(monitor_packet) = monitor_packet(&packet) else { - guard.deliver(packet); - continue; - }; - - if self.monitors.on_packet_sent(&monitor_packet).is_drop() { - continue; - } - - if self.monitors.on_packet_received(&monitor_packet).is_drop() { - continue; - } - + pub fn deliver(&self, packet: turmoil_net::Packet) { + if let Some(guard) = self.guard.as_ref() { guard.deliver(packet); } } + pub fn monitors(&self) -> Monitors { + self.monitors.clone() + } + fn prepare(&mut self) -> io::Result<()> { if self.guard.is_some() { return Ok(()); @@ -225,7 +218,7 @@ impl Registry { } } -fn monitor_packet(packet: &turmoil_net::Packet) -> Option { +pub(crate) fn monitor_packet(packet: &turmoil_net::Packet) -> Option { let transport = match &packet.payload { turmoil_net::Transport::Udp(udp) => Transport::Udp(transport::Udp { source: udp.src_port, @@ -280,3 +273,16 @@ fn monitor_packet(packet: &turmoil_net::Packet) -> Option { packet.update_checksum(); Some(packet) } + +fn bind_udp_socket(local_addr: std::net::SocketAddr) -> io::Result { + let mut future = pin!(TurmoilUdpSocket::bind(local_addr)); + let waker = crate::task::waker::noop(); + let mut cx = Context::from_waker(&waker); + + match Future::poll(future.as_mut(), &mut cx) { + core::task::Poll::Ready(result) => result, + core::task::Poll::Pending => Err(io::Error::other( + "turmoil-net UDP bind unexpectedly returned pending", + )), + } +} diff --git a/bach/src/environment/net/socket/udp.rs b/bach/src/environment/net/socket/udp.rs index 9fec37e..50fe581 100644 --- a/bach/src/environment/net/socket/udp.rs +++ b/bach/src/environment/net/socket/udp.rs @@ -16,11 +16,13 @@ use core::{ }; use std::{ io, + panic::{self, AssertUnwindSafe}, sync::Mutex, }; +use turmoil_net::shim::tokio::net::UdpSocket as TurmoilUdpSocket; pub struct Socket { - inner: turmoil_net::UdpSocket, + inner: Option, local_addr: SocketAddr, peer_addr: Mutex>, monitors: Monitors, @@ -28,17 +30,15 @@ pub struct Socket { macro_rules! lock { ($lock:expr) => { - $lock - .lock() - .map_err(|e| io::Error::other(format!("{e}")))? + $lock.lock().map_err(|e| io::Error::other(format!("{e}")))? }; } impl Socket { - pub fn new(inner: turmoil_net::UdpSocket, monitors: Monitors) -> io::Result { + pub fn new(inner: TurmoilUdpSocket, monitors: Monitors) -> io::Result { let local_addr = inner.local_addr()?; Ok(Self { - inner, + inner: Some(inner), local_addr, peer_addr: Mutex::new(None), monitors, @@ -50,7 +50,7 @@ impl socket::Socket for Socket { fn poll_connect(&self, cx: &mut Context, peer_addr: SocketAddr) -> Poll> { set_current_group()?; - let mut future = pin!(self.inner.connect(peer_addr)); + let mut future = pin!(self.inner().connect(peer_addr)); match Future::poll(future.as_mut(), cx) { Poll::Ready(Ok(())) => { *lock!(self.peer_addr) = Some(peer_addr); @@ -65,7 +65,7 @@ impl socket::Socket for Socket { if let Some(peer_addr) = *lock!(self.peer_addr) { Ok(peer_addr) } else { - self.with_current(|| self.inner.peer_addr()) + self.with_current(|| self.inner().peer_addr()) } } @@ -134,28 +134,78 @@ impl socket::Socket for Socket { let payload = flatten_payload(payload); let segments = segment_payload(&payload, opts.segment_len); + let mut cx = cx; self.with_current(|| { let mut sent = 0; for segment in &segments { - let len = if let Some(cx) = cx { - let mut future = pin!(self.inner.send_to(segment, destination)); + let len = if let Some(cx) = cx.as_deref_mut() { + let mut future = pin!(self.inner().send_to(segment, destination)); match Future::poll(future.as_mut(), cx) { Poll::Ready(res) => res?, Poll::Pending => return Err(io::ErrorKind::WouldBlock.into()), } } else { - self.inner.try_send_to(segment, destination)? + self.inner().try_send_to(segment, destination)? }; sent += len; } - registry::with_registry(|registry| { - registry.drive(); - Ok(()) + let (monitors, packets) = registry::with_registry(|registry| { + Ok((registry.monitors(), registry.drain_packets())) })?; + let mut panic_payload = None; + + for packet in packets { + let Some(monitor_packet) = monitor_packet(&packet) else { + registry::with_registry(|registry| { + registry.deliver(packet); + Ok(()) + })?; + continue; + }; + + let sent = panic::catch_unwind(AssertUnwindSafe(|| { + monitors.on_packet_sent(&monitor_packet) + })); + let sent = match sent { + Ok(command) => command, + Err(payload) => { + panic_payload = Some(payload); + break; + } + }; + + if sent.is_drop() { + continue; + } + + let received = panic::catch_unwind(AssertUnwindSafe(|| { + monitors.on_packet_received(&monitor_packet) + })); + let received = match received { + Ok(command) => command, + Err(payload) => { + panic_payload = Some(payload); + break; + } + }; + + if received.is_drop() { + continue; + } + + registry::with_registry(|registry| { + registry.deliver(packet); + Ok(()) + })?; + } + + if let Some(payload) = panic_payload { + panic::resume_unwind(payload); + } Ok(sent) }) @@ -186,19 +236,19 @@ impl socket::Socket for Socket { } if opts.peek { - let mut future = pin!(self.inner.peek_from(&mut buffer)); + let mut future = pin!(self.inner().peek_from(&mut buffer)); match Future::poll(future.as_mut(), cx.expect("peek requires a task context")) { Poll::Ready(res) => res, Poll::Pending => Err(io::ErrorKind::WouldBlock.into()), } } else if let Some(cx) = cx { - let mut future = pin!(self.inner.recv_from(&mut buffer)); + let mut future = pin!(self.inner().recv_from(&mut buffer)); match Future::poll(future.as_mut(), cx) { Poll::Ready(res) => res, Poll::Pending => Err(io::ErrorKind::WouldBlock.into()), } } else { - self.inner.try_recv_from(&mut buffer) + self.inner().try_recv_from(&mut buffer) } })?; @@ -228,6 +278,10 @@ impl socket::Socket for Socket { } impl Socket { + fn inner(&self) -> &TurmoilUdpSocket { + self.inner.as_ref().expect("socket already dropped") + } + fn with_current(&self, f: F) -> io::Result where F: FnOnce() -> io::Result, @@ -241,6 +295,29 @@ impl Drop for Socket { fn drop(&mut self) { self.monitors .on_socket_closed(&self.local_addr, crate::net::monitor::Transport::Udp); + + let Some(inner) = self.inner.take() else { + return; + }; + + let registry_available = registry::scope::try_borrow_with(|scope| scope.is_some()); + if !registry_available { + std::mem::forget(inner); + return; + } + + let mut inner = Some(inner); + let result = panic::catch_unwind(AssertUnwindSafe(|| { + if set_current_group().is_ok() { + drop(inner.take()); + } + })); + + if let Some(inner) = inner.take() { + std::mem::forget(inner); + } + + let _ = result; } } @@ -283,3 +360,7 @@ fn copy_payload(src: &[u8], dst: &mut [io::IoSliceMut]) -> usize { copied } + +fn monitor_packet(packet: &turmoil_net::Packet) -> Option { + registry::monitor_packet(packet) +} diff --git a/bach/src/scope.rs b/bach/src/scope.rs index dd14292..b433048 100644 --- a/bach/src/scope.rs +++ b/bach/src/scope.rs @@ -20,9 +20,24 @@ macro_rules! define { #[allow(dead_code)] pub fn with R, R>(value: $ty, f: F) -> ($ty, R) { + struct Reset { + prev: Option<$ty>, + active: bool, + } + + impl Drop for Reset { + fn drop(&mut self) { + if self.active { + let _ = set(self.prev.take()); + } + } + } + let prev = set(Some(value)); + let mut reset = Reset { prev, active: true }; let res = f(); - let value = set(prev).unwrap(); + let value = set(reset.prev.take()).unwrap(); + reset.active = false; (value, res) } From a5305adee0972e68ba822e8d0bdbf75ec1161bcf Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 14 May 2026 10:30:12 +0000 Subject: [PATCH 3/3] Handle turmoil-net panic cleanup Agent-Logs-Url: https://github.com/camshaft/bach/sessions/0fe5944e-8451-4d54-ad6a-dac452aff1be Co-authored-by: camshaft <799311+camshaft@users.noreply.github.com> --- bach/src/environment/net/registry.rs | 2 ++ bach/src/environment/net/socket/udp.rs | 8 +++++++- 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/bach/src/environment/net/registry.rs b/bach/src/environment/net/registry.rs index 5caed1b..73af270 100644 --- a/bach/src/environment/net/registry.rs +++ b/bach/src/environment/net/registry.rs @@ -63,6 +63,8 @@ impl Registry { } pub fn set_pcap_dir>(&mut self, pcap: P) -> io::Result<()> { + // The first-pass turmoil-net backend does not emit PCAP files yet, but keep the + // builder method available so existing callers continue to compile. let _ = pcap.into(); Ok(()) } diff --git a/bach/src/environment/net/socket/udp.rs b/bach/src/environment/net/socket/udp.rs index 50fe581..386889e 100644 --- a/bach/src/environment/net/socket/udp.rs +++ b/bach/src/environment/net/socket/udp.rs @@ -117,6 +117,13 @@ impl socket::Socket for Socket { )); } + if opts.segment_len == Some(0) { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "segment_len must be greater than zero", + )); + } + let destination = if destination.ip().is_unspecified() || destination.port() == 0 { self.peer_addr()? } else { @@ -337,7 +344,6 @@ fn flatten_payload(payload: &[io::IoSlice]) -> Vec { fn segment_payload(payload: &[u8], segment_len: Option) -> Vec> { match segment_len { - Some(0) => vec![payload.to_vec()], Some(segment_len) => payload.chunks(segment_len).map(ToOwned::to_owned).collect(), None => vec![payload.to_vec()], }