This commit is contained in:
Philipp Krüger
2025-05-29 13:39:45 +02:00
parent a5a28dfc87
commit bc2358e616
7 changed files with 396 additions and 185 deletions
+1
View File
@@ -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 }
+2
View File
@@ -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)]
+63
View File
@@ -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<TestAddr> 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<EcnCodepoint>,
pub contents: Bytes,
pub segment_size: Option<usize>,
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(())
}
}
+229
View File
@@ -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<OwnedTransmit>,
receiver: mpsc::Receiver<OwnedTransmit>,
}
#[derive(Debug)]
pub struct VirtualSocketSender {
addr: SocketAddr,
sender: PollSender<OwnedTransmit>,
}
#[derive(Debug)]
pub struct Plug {
pub sender: mpsc::Sender<OwnedTransmit>,
pub receiver: mpsc::Receiver<OwnedTransmit>,
}
impl Plug {
pub fn testometer(
capacity: usize,
) -> (
Self,
mpsc::Sender<OwnedTransmit>,
mpsc::Receiver<OwnedTransmit>,
) {
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<SocketAddr>, 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<std::io::Result<()>> {
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<dyn crate::UdpSender>> {
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<std::io::Result<usize>> {
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<std::net::SocketAddr> {
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(())
}
}
+26
View File
@@ -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<SocketAddr, Plug>) -> Self {
Self {
task: AbortOnDropHandle::new(tokio::spawn(async move { Self::run(plugs).await })),
}
}
async fn run(mut plugs: BTreeMap<SocketAddr, Plug>) {
// let mut pending_transmits = Vec::<OwnedTransmit>::new();
// loop {
// for (_, plug) in plugs.iter_mut() {
// plug.receiver.poll_recv_many(cx, buffer, limit)
// }
// }
}
}
+70
View File
@@ -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(())
}
}
+5 -185
View File
@@ -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<VirtualNet>,
}
#[derive(Debug, Default)]
struct VirtualNet {
sockets: BTreeMap<SocketAddr, Mutex<VirtualSocketState>>,
}
#[derive(Debug, Default)]
struct VirtualSocketState {
datagrams: VecDeque<OwnedTransmit>,
wiretap: Option<UnboundedSender<OwnedTransmit>>,
wakers: Vec<Waker>,
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<VirtualNet>, addr: SocketAddr) -> Arc<Self> {
Arc::new(Self {
net: Arc::clone(net),
addr,
})
}
fn wiretap(&self) -> UnboundedReceiver<OwnedTransmit> {
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<EcnCodepoint>,
pub contents: Bytes,
pub segment_size: Option<usize>,
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<Self>) -> Pin<Box<dyn quinn::UdpPoller>> {
#[derive(Debug)]
struct UdpAlwaysReady;
impl quinn::UdpPoller for UdpAlwaysReady {
fn poll_writable(self: Pin<&mut Self>, _cx: &mut Context) -> Poll<std::io::Result<()>> {
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<std::io::Result<usize>> {
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<std::net::SocketAddr> {
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();