From bc2358e616eb2f3b504de4bd1bbb58452b01a168 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Philipp=20Kr=C3=BCger?= Date: Thu, 29 May 2025 13:39:45 +0200 Subject: [PATCH] WIP --- quinn/Cargo.toml | 1 + quinn/src/lib.rs | 2 + quinn/src/virtualnet.rs | 63 +++++++++ quinn/src/virtualnet/socket.rs | 229 +++++++++++++++++++++++++++++++++ quinn/src/virtualnet/switch.rs | 26 ++++ quinn/src/virtualnet/wire.rs | 70 ++++++++++ quinn/tests/web_test.rs | 190 +-------------------------- 7 files changed, 396 insertions(+), 185 deletions(-) create mode 100644 quinn/src/virtualnet.rs create mode 100644 quinn/src/virtualnet/socket.rs create mode 100644 quinn/src/virtualnet/switch.rs create mode 100644 quinn/src/virtualnet/wire.rs diff --git a/quinn/Cargo.toml b/quinn/Cargo.toml index 5cb3f7b52..9e92e2225 100644 --- a/quinn/Cargo.toml +++ b/quinn/Cargo.toml @@ -57,6 +57,7 @@ thiserror = { workspace = true } tracing = { workspace = true } tokio = { workspace = true } udp = { package = "iroh-quinn-udp", path = "../quinn-udp", version = "0.5", default-features = false, features = ["tracing"] } +tokio-util = { version = "0.7.15", features = ["rt"] } # Fix minimal dependencies for indirect deps async-global-executor = { workspace = true, optional = true } diff --git a/quinn/src/lib.rs b/quinn/src/lib.rs index b9d1c8918..f2cbf381e 100644 --- a/quinn/src/lib.rs +++ b/quinn/src/lib.rs @@ -61,6 +61,8 @@ mod runtime; mod send_stream; mod work_limiter; +pub mod virtualnet; + #[cfg(not(wasm_browser))] pub(crate) use std::time::{Duration, Instant}; #[cfg(wasm_browser)] diff --git a/quinn/src/virtualnet.rs b/quinn/src/virtualnet.rs new file mode 100644 index 000000000..e9b292307 --- /dev/null +++ b/quinn/src/virtualnet.rs @@ -0,0 +1,63 @@ +#![allow(missing_docs)] +use std::net::SocketAddr; + +use bytes::Bytes; +use udp::{EcnCodepoint, Transmit}; + +pub mod socket; +#[cfg(feature = "runtime-tokio")] +pub mod switch; +pub mod wire; + +pub struct TestAddr(pub u8); + +impl From for SocketAddr { + fn from(TestAddr(id): TestAddr) -> Self { + ([1, 1, 1, id], 42u16).into() + } +} + +#[derive(Debug, Clone)] +pub struct OwnedTransmit { + pub destination: SocketAddr, + pub ecn: Option, + pub contents: Bytes, + pub segment_size: Option, + pub src_ip: SocketAddr, +} + +impl OwnedTransmit { + fn new(src: SocketAddr, t: &udp::Transmit) -> Self { + Self { + destination: t.destination, + ecn: t.ecn.clone(), + contents: Bytes::copy_from_slice(t.contents), + segment_size: t.segment_size.clone(), + src_ip: SocketAddr::new(t.src_ip.unwrap_or(src.ip()), src.port()), + } + } + + fn as_quinn_transmit(&self) -> Transmit<'_> { + Transmit { + destination: self.destination, + ecn: self.ecn, + contents: self.contents.as_ref(), + segment_size: self.segment_size.clone(), + src_ip: Some(self.src_ip.ip()), + } + } + + fn receive_into( + &self, + buf: &mut std::io::IoSliceMut<'_>, + meta: &mut udp::RecvMeta, + ) -> std::io::Result<()> { + buf[..self.contents.len()].copy_from_slice(&self.contents); + meta.addr = self.src_ip; + meta.dst_ip = Some(self.destination.ip()); + meta.len = self.contents.len(); + meta.stride = self.contents.len(); + meta.ecn = self.ecn; + Ok(()) + } +} diff --git a/quinn/src/virtualnet/socket.rs b/quinn/src/virtualnet/socket.rs new file mode 100644 index 000000000..6e1eb758c --- /dev/null +++ b/quinn/src/virtualnet/socket.rs @@ -0,0 +1,229 @@ +use std::{ + net::SocketAddr, + pin::Pin, + task::{Context, Poll}, +}; + +use crate::AsyncUdpSocket; +use tokio::sync::mpsc; +use tokio_util::sync::PollSender; +use udp::Transmit; + +use super::OwnedTransmit; + +#[derive(Debug)] +pub struct VirtualSocket { + pub addr: SocketAddr, + sender: mpsc::Sender, + receiver: mpsc::Receiver, +} + +#[derive(Debug)] +pub struct VirtualSocketSender { + addr: SocketAddr, + sender: PollSender, +} + +#[derive(Debug)] +pub struct Plug { + pub sender: mpsc::Sender, + pub receiver: mpsc::Receiver, +} + +impl Plug { + pub fn testometer( + capacity: usize, + ) -> ( + Self, + mpsc::Sender, + mpsc::Receiver, + ) { + let (sender, test_receiver) = mpsc::channel(capacity); + let (test_sender, receiver) = mpsc::channel(capacity); + (Self { sender, receiver }, test_sender, test_receiver) + } +} + +impl VirtualSocket { + pub fn new(addr: impl Into, plug: Plug) -> Self { + Self { + addr: addr.into(), + sender: plug.sender, + receiver: plug.receiver, + } + } + + #[cfg(test)] + pub async fn receive_data(&mut self) -> std::io::Result<(SocketAddr, bytes::Bytes)> { + use std::io::IoSliceMut; + + let mut buf = [0u8; 1200]; + let mut bufs = [IoSliceMut::new(&mut buf)]; + let mut meta = [udp::RecvMeta::default()]; + + let num_datagrams = + std::future::poll_fn(|cx| self.poll_recv(cx, &mut bufs, &mut meta)).await?; + + debug_assert_eq!(num_datagrams, 1); // we don't support GSO/GRO(?) + + Ok(( + meta[0].addr, + bytes::Bytes::copy_from_slice(&buf[..meta[0].len]), + )) + } +} + +impl crate::UdpSender for VirtualSocketSender { + fn max_transmit_segments(&self) -> usize { + 1 + } + + fn poll_send( + mut self: Pin<&mut Self>, + transmit: &Transmit, + cx: &mut Context, + ) -> Poll> { + if let Err(_closed) = std::task::ready!(self.as_mut().sender.poll_reserve(cx)) { + return Poll::Ready(Err(std::io::Error::new( + std::io::ErrorKind::BrokenPipe, + "virtual socket closed when sending", + ))); + } + + self.try_send(transmit)?; + + Poll::Ready(Ok(())) + } + + fn try_send(mut self: Pin<&mut Self>, transmit: &Transmit) -> std::io::Result<()> { + let addr = self.addr; + if let Err(_closed) = self + .as_mut() + .sender + .send_item(OwnedTransmit::new(addr, transmit)) + { + return Err(std::io::Error::new( + std::io::ErrorKind::BrokenPipe, + "virtual socket closed when sending", + )); + } + + Ok(()) + } +} + +impl AsyncUdpSocket for VirtualSocket { + fn create_sender(&self) -> Pin> { + Box::pin(VirtualSocketSender { + addr: self.addr, + sender: PollSender::new(self.sender.clone()), + }) + } + + fn poll_recv( + &mut self, + cx: &mut Context, + bufs: &mut [std::io::IoSliceMut<'_>], + meta: &mut [udp::RecvMeta], + ) -> Poll> { + let mut transmits = Vec::new(); + let recvd = ready!(self.receiver.poll_recv_many(cx, &mut transmits, bufs.len())); + if recvd == 0 { + return Poll::Ready(Err(std::io::Error::new( + std::io::ErrorKind::BrokenPipe, + "virtual socket closed when receiving", + ))); + } + + for ((t, buf), meta) in transmits + .into_iter() + .zip(bufs.iter_mut()) + .zip(meta.iter_mut()) + { + if buf.len() >= t.contents.len() { + tracing::debug!( + "Received from {:?} to {:?}: {} bytes", + t.src_ip, + t.destination, + t.contents.len() + ); + t.receive_into(buf, meta)?; + } + } + + Poll::Ready(Ok(recvd)) + } + + fn local_addr(&self) -> std::io::Result { + Ok(self.addr) + } +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use crate::{ + virtualnet::{OwnedTransmit, TestAddr}, + AsyncUdpSocket, + }; + + use super::{Plug, VirtualSocket}; + + #[tokio::test] + async fn test_recv() -> anyhow::Result<()> { + let (plug, sender, _receiver) = Plug::testometer(64); + let mut socket = VirtualSocket::new(TestAddr(11), plug); + + let other_addr = TestAddr(99).into(); + let contents = Bytes::copy_from_slice(b"Hello, world!"); + let transmit = OwnedTransmit { + contents: contents.clone(), + destination: socket.addr, + ecn: None, + segment_size: None, + src_ip: other_addr, + }; + + sender.send(transmit).await?; + + let (source_addr, received) = socket.receive_data().await?; + + assert_eq!(received, contents); + assert_eq!(source_addr, other_addr); + + Ok(()) + } + + #[tokio::test] + async fn test_send() -> anyhow::Result<()> { + let (plug, _sender, mut receiver) = Plug::testometer(64); + let socket = VirtualSocket::new(TestAddr(11), plug); + let mut socket_sender = socket.create_sender(); + + let other_addr = TestAddr(99).into(); + let contents = Bytes::copy_from_slice(b"Hello, world!"); + let transmit = OwnedTransmit { + contents: contents.clone(), + destination: other_addr, + ecn: None, + segment_size: None, + src_ip: socket.addr, + }; + + std::future::poll_fn(|cx| { + socket_sender + .as_mut() + .poll_send(&transmit.as_quinn_transmit(), cx) + }) + .await?; + + assert!(!receiver.is_empty()); + let received = receiver.recv().await.unwrap(); + assert_eq!(received.src_ip, socket.addr); + assert_eq!(received.destination, other_addr); + assert_eq!(received.contents, contents); + + Ok(()) + } +} diff --git a/quinn/src/virtualnet/switch.rs b/quinn/src/virtualnet/switch.rs new file mode 100644 index 000000000..22b8a614d --- /dev/null +++ b/quinn/src/virtualnet/switch.rs @@ -0,0 +1,26 @@ +use std::{collections::BTreeMap, net::SocketAddr}; + +use tokio_util::task::AbortOnDropHandle; + +use super::socket::Plug; + +pub struct Switch { + task: AbortOnDropHandle<()>, +} + +impl Switch { + pub fn new(plugs: BTreeMap) -> Self { + Self { + task: AbortOnDropHandle::new(tokio::spawn(async move { Self::run(plugs).await })), + } + } + + async fn run(mut plugs: BTreeMap) { + // let mut pending_transmits = Vec::::new(); + // loop { + // for (_, plug) in plugs.iter_mut() { + // plug.receiver.poll_recv_many(cx, buffer, limit) + // } + // } + } +} diff --git a/quinn/src/virtualnet/wire.rs b/quinn/src/virtualnet/wire.rs new file mode 100644 index 000000000..1f65780ad --- /dev/null +++ b/quinn/src/virtualnet/wire.rs @@ -0,0 +1,70 @@ +use tokio::sync::mpsc; + +use super::socket::Plug; + +pub struct Wire { + start: Plug, + end: Plug, +} + +impl Wire { + pub fn new(capacity: usize) -> Self { + // start -> end transmission + let (start_sender, end_receiver) = mpsc::channel(capacity); + // end -> start transmission + let (end_sender, start_receiver) = mpsc::channel(capacity); + let start = Plug { + sender: start_sender, + receiver: start_receiver, + }; + let end = Plug { + sender: end_sender, + receiver: end_receiver, + }; + + Self { start, end } + } +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use crate::{ + virtualnet::{socket::VirtualSocket, OwnedTransmit, TestAddr}, + AsyncUdpSocket, + }; + + use super::Wire; + + #[tokio::test] + async fn test_wire_plugging() -> std::io::Result<()> { + let wire = Wire::new(64); + let socket0 = VirtualSocket::new(TestAddr(11), wire.start); + let mut socket1 = VirtualSocket::new(TestAddr(99), wire.end); + + let contents = Bytes::copy_from_slice(b"Hello, world!"); + let transmit = OwnedTransmit { + contents: contents.clone(), + destination: socket1.addr, + ecn: None, + segment_size: None, + src_ip: socket0.addr, + }; + + let mut socket_sender = socket0.create_sender(); + std::future::poll_fn(|cx| { + socket_sender + .as_mut() + .poll_send(&transmit.as_quinn_transmit(), cx) + }) + .await?; + + let (source_addr, received) = socket1.receive_data().await?; + + assert_eq!(source_addr, socket0.addr); + assert_eq!(received, contents); + + Ok(()) + } +} diff --git a/quinn/tests/web_test.rs b/quinn/tests/web_test.rs index ab462dde2..cacc590b2 100644 --- a/quinn/tests/web_test.rs +++ b/quinn/tests/web_test.rs @@ -3,200 +3,20 @@ use std::{ net::SocketAddr, pin::Pin, sync::{Arc, Mutex}, - task::{Context, Poll, Waker}, + task::{ready, Context, Poll, Waker}, }; use bytes::Bytes; +use iroh_quinn::{AsyncUdpSocket, Endpoint, Runtime as _}; use proto::{ClientConfig, EndpointConfig, ServerConfig}; -use quinn::{AsyncUdpSocket, Endpoint, Runtime as _}; -use tokio::sync::mpsc::{UnboundedReceiver, UnboundedSender}; +use tokio::sync::mpsc; +use tokio_util::{sync::PollSender, task::AbortOnDropHandle}; use udp::{EcnCodepoint, Transmit}; -#[derive(Debug)] -struct VirtualSocket { - addr: SocketAddr, - net: Arc, -} - -#[derive(Debug, Default)] -struct VirtualNet { - sockets: BTreeMap>, -} - -#[derive(Debug, Default)] -struct VirtualSocketState { - datagrams: VecDeque, - wiretap: Option>, - wakers: Vec, - paused: bool, -} - -impl VirtualNet { - fn add_socket(&mut self, id: u8) -> SocketAddr { - let addr = SocketAddr::from(([192, 168, 0, id], 1234u16)); - self.sockets.insert(addr, Default::default()); - addr - } -} - -impl VirtualSocket { - fn new(net: &Arc, addr: SocketAddr) -> Arc { - Arc::new(Self { - net: Arc::clone(net), - addr, - }) - } - - fn wiretap(&self) -> UnboundedReceiver { - let (sender, receiver) = tokio::sync::mpsc::unbounded_channel(); - self.socket_state().wiretap = Some(sender); - receiver - } - - fn set_paused(&self, paused: bool) { - let mut socket_state = self.socket_state(); - - socket_state.paused = paused; - if !paused { - while let Some(waker) = socket_state.wakers.pop() { - waker.wake(); - } - } - } - - fn socket_state(&self) -> std::sync::MutexGuard<'_, VirtualSocketState> { - self.net - .sockets - .get(&self.addr) - .expect("socket missing") - .lock() - .expect("poisoned") - } -} - -#[derive(Debug, Clone)] -pub struct OwnedTransmit { - pub destination: SocketAddr, - pub ecn: Option, - pub contents: Bytes, - pub segment_size: Option, - pub src_ip: SocketAddr, -} - -impl OwnedTransmit { - fn new(src: SocketAddr, t: &udp::Transmit) -> Self { - Self { - destination: t.destination, - ecn: t.ecn.clone(), - contents: Bytes::copy_from_slice(t.contents), - segment_size: t.segment_size.clone(), - src_ip: SocketAddr::new(t.src_ip.unwrap_or(src.ip()), src.port()), - } - } - - fn as_quinn_transmit(&self) -> Transmit<'_> { - Transmit { - destination: self.destination, - ecn: self.ecn, - contents: self.contents.as_ref(), - segment_size: self.segment_size.clone(), - src_ip: Some(self.src_ip.ip()), - } - } -} - -impl AsyncUdpSocket for VirtualSocket { - fn create_io_poller(self: Arc) -> Pin> { - #[derive(Debug)] - struct UdpAlwaysReady; - - impl quinn::UdpPoller for UdpAlwaysReady { - fn poll_writable(self: Pin<&mut Self>, _cx: &mut Context) -> Poll> { - Poll::Ready(Ok(())) - } - } - - Box::pin(UdpAlwaysReady) - } - - fn try_send(&self, transmit: &udp::Transmit) -> std::io::Result<()> { - // If there's no socket, then there's no point to ever send, since nobody could ever receive - let Some(send) = self.net.sockets.get(&transmit.destination) else { - return Ok(()); - }; - - let mut socket_state = send.lock().expect("poisoned"); - let transmit = OwnedTransmit::new(self.addr, transmit); - socket_state.datagrams.push_back(transmit.clone()); - if let Some(sender) = &socket_state.wiretap { - if sender.send(transmit).is_err() { - socket_state.wiretap = None; - } - } - if !socket_state.paused { - while let Some(waker) = socket_state.wakers.pop() { - waker.wake(); - } - } - Ok(()) - } - - fn poll_recv( - &self, - cx: &mut Context, - bufs: &mut [std::io::IoSliceMut<'_>], - meta: &mut [udp::RecvMeta], - ) -> Poll> { - let mut socket_state = self.socket_state(); - - if socket_state.paused { - socket_state.wakers.push(cx.waker().clone()); - return Poll::Pending; - } - - let mut num_msgs = 0; - while let Some(t) = socket_state.datagrams.pop_front() { - if bufs.len() <= num_msgs || meta.len() <= num_msgs { - break; - } - - let buf = &mut bufs[num_msgs]; - let meta = &mut meta[num_msgs]; - - if buf.len() >= t.contents.len() { - tracing::debug!( - "Received from {:?} to {:?}: {} bytes", - t.src_ip, - t.destination, - t.contents.len() - ); - buf[..t.contents.len()].copy_from_slice(&t.contents); - meta.addr = t.src_ip; - meta.dst_ip = Some(t.destination.ip()); - meta.len = t.contents.len(); - meta.stride = t.contents.len(); - meta.ecn = t.ecn; - num_msgs += 1; - } - } - - if num_msgs > 0 { - Poll::Ready(Ok(num_msgs)) - } else { - socket_state.wakers.push(cx.waker().clone()); - Poll::Pending - } - } - - fn local_addr(&self) -> std::io::Result { - Ok(self.addr) - } -} - #[tokio::test] async fn test_connect_with_virtual_socket() { tracing_subscriber::fmt::init(); - let runtime = Arc::new(quinn::TokioRuntime); + let runtime = Arc::new(iroh_quinn::TokioRuntime); // Virtual sockets setup let mut net = VirtualNet::default();