--- a/crates/host/src/main.rs +++ b/crates/host/src/main.rs @@ -410,7 +410,7 @@ let (mut enc_r, mut enc_w) = tokio::io::split(reader.into_inner()); let sock = std::sync::Arc::new(sock); let s2 = sock.clone(); - let to_net = tokio::spawn(async move { + let to_net = async move { loop { let len = match enc_r.read_u16().await { Ok(l) => l as usize, @@ -424,8 +424,8 @@ return; } } - }); - let from_net = tokio::spawn(async move { + }; + let from_net = async move { let mut buf = vec![0u8; 65535]; loop { let n = match sock.recv(&mut buf).await { @@ -436,7 +436,7 @@ return; } } - }); + }; tokio::select! { _ = to_net => {} _ = from_net => {} @@ -510,3 +510,26 @@ }); } } + +#[cfg(test)] +mod udp_lifecycle_tests { + use super::*; + use tokio::io::AsyncReadExt; + + #[tokio::test] + async fn udp_peer_close_releases_both_halves_without_network_reply() { + let remote = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let (link, mut peer) = tokio::io::duplex(4096); + let args = Args::parse_from(["host", "--dev", "--allow-private", "--records-dir", "/tmp"]); + let handler = tokio::spawn(async move { handle_udp(&args, Box::new(link)).await }); + peer.write_all(format!("UDP {}\n", remote.local_addr().unwrap()).as_bytes()).await.unwrap(); + let mut ack = [0;3]; + peer.read_exact(&mut ack).await.unwrap(); + assert_eq!(&ack, b"OK\n"); + peer.shutdown().await.unwrap(); + handler.await.unwrap().unwrap(); + let mut rest = Vec::new(); + tokio::time::timeout(Duration::from_secs(1), peer.read_to_end(&mut rest)) + .await.expect("UDP writer survived a closed enclave association").unwrap(); + } +}