From c05f8b33ef0cd1f86a3745f993e391b6d04057e0 Mon Sep 17 00:00:00 2001 From: Luca Versari Date: Mon, 1 Dec 2025 02:22:57 +0100 Subject: [PATCH] Refactor net to be independent of the OS. --- pixie-uefi/Cargo.lock | 25 ++ pixie-uefi/Cargo.toml | 1 + pixie-uefi/src/flash.rs | 25 +- pixie-uefi/src/main.rs | 37 +- pixie-uefi/src/os/mod.rs | 112 +----- pixie-uefi/src/os/net.rs | 525 ----------------------------- pixie-uefi/src/os/net/interface.rs | 124 +++++++ pixie-uefi/src/os/net/mod.rs | 203 +++++++++++ pixie-uefi/src/os/net/speed.rs | 55 +++ pixie-uefi/src/os/net/tcp.rs | 185 ++++++++++ pixie-uefi/src/os/net/udp.rs | 102 ++++++ pixie-uefi/src/register.rs | 19 +- pixie-uefi/src/store.rs | 35 +- 13 files changed, 761 insertions(+), 687 deletions(-) delete mode 100644 pixie-uefi/src/os/net.rs create mode 100644 pixie-uefi/src/os/net/interface.rs create mode 100644 pixie-uefi/src/os/net/mod.rs create mode 100644 pixie-uefi/src/os/net/speed.rs create mode 100644 pixie-uefi/src/os/net/tcp.rs create mode 100644 pixie-uefi/src/os/net/udp.rs diff --git a/pixie-uefi/Cargo.lock b/pixie-uefi/Cargo.lock index 067ef4ed..15e849ad 100644 --- a/pixie-uefi/Cargo.lock +++ b/pixie-uefi/Cargo.lock @@ -286,6 +286,15 @@ dependencies = [ "stable_deref_trait", ] +[[package]] +name = "lock_api" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" +dependencies = [ + "scopeguard", +] + [[package]] name = "log" version = "0.4.28" @@ -371,6 +380,7 @@ dependencies = [ "pixie-shared", "postcard", "smoltcp", + "spin", "thingbuf", "uefi", ] @@ -456,6 +466,12 @@ dependencies = [ "winapi-util", ] +[[package]] +name = "scopeguard" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" + [[package]] name = "serde" version = "1.0.228" @@ -506,6 +522,15 @@ dependencies = [ "managed", ] +[[package]] +name = "spin" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d5fe4ccb98d9c292d56fec89a5e07da7fc4cf0dc11e156b41793132775d3e591" +dependencies = [ + "lock_api", +] + [[package]] name = "stable_deref_trait" version = "1.2.0" diff --git a/pixie-uefi/Cargo.toml b/pixie-uefi/Cargo.toml index 66cd29d3..61e72c32 100644 --- a/pixie-uefi/Cargo.toml +++ b/pixie-uefi/Cargo.toml @@ -24,6 +24,7 @@ managed = { version = "0.8.0", default-features = false, features = ["alloc"] } minicov = { version = "0.3.7", optional = true } postcard = { version = "1.1.3", default-features = false, features = ["alloc"] } smoltcp = { version = "0.12.0", default-features = false, features = ["alloc", "proto-ipv4", "medium-ethernet", "socket-udp", "socket-tcp", "socket-dhcpv4", "async", "socket-tcp-cubic"] } +spin = "0.10.0" thingbuf = { version = "0.1.6", default-features = false, features = ["alloc"] } uefi = { version = "0.36.1", features = ["alloc", "global_allocator", "panic_handler"] } diff --git a/pixie-uefi/src/flash.rs b/pixie-uefi/src/flash.rs index 9d9b6aba..a2de03ca 100644 --- a/pixie-uefi/src/flash.rs +++ b/pixie-uefi/src/flash.rs @@ -16,17 +16,18 @@ use uefi::proto::console::text::Color; use crate::os::boot_options::BootOptions; use crate::os::error::{Error, Result}; -use crate::os::{memory, TcpStream, UefiOS, PACKET_SIZE}; +use crate::os::net::{TcpStream, UdpSocket, ETH_PACKET_SIZE}; +use crate::os::{memory, UefiOS}; use crate::MIN_MEMORY; async fn fetch_image(stream: &TcpStream) -> Result { let req = TcpRequest::GetImage; let mut buf = postcard::to_allocvec(&req)?; - stream.send_u64_le(buf.len() as u64).await?; - stream.send(&buf).await?; - let len = stream.recv_u64_le().await?; + stream.write_u64_le(buf.len() as u64).await?; + stream.write_all(&buf).await?; + let len = stream.read_u64_le().await?; buf.resize(len as usize, 0); - stream.recv_exact(&mut buf).await?; + stream.read_exact(&mut buf).await?; Ok(postcard::from_bytes(&buf)?) } @@ -74,9 +75,9 @@ fn handle_packet( } pub async fn flash(os: UefiOS, server_addr: SocketAddrV4) -> Result<()> { - let stream = os.connect(server_addr).await?; + let stream = TcpStream::connect(server_addr).await?; let image = fetch_image(&stream).await?; - stream.close_send().await; + stream.shutdown().await; // TODO(virv): this could be better stream.force_close().await; @@ -160,8 +161,8 @@ pub async fn flash(os: UefiOS, server_addr: SocketAddrV4) -> Result<()> { info!("Disk scanned; {} chunks to fetch", stats.borrow().fetch); - let socket = os.udp_bind(Some(CHUNKS_PORT)).await?; - let mut buf = [0; PACKET_SIZE]; + let socket = UdpSocket::bind(Some(CHUNKS_PORT)).await?; + let mut buf = [0; ETH_PACKET_SIZE]; let mut received = BTreeMap::new(); @@ -177,7 +178,7 @@ pub async fn flash(os: UefiOS, server_addr: SocketAddrV4) -> Result<()> { BytesFmt(free_mem) ); while !chunks_info.is_empty() { - let recv = Box::pin(socket.recv(&mut buf)); + let recv = Box::pin(socket.recv_from(&mut buf)); let sleep = Box::pin(os.sleep_us(100_000)); match select(recv, sleep).await { Either::Left(((buf, _addr), _)) => { @@ -204,7 +205,7 @@ pub async fn flash(os: UefiOS, server_addr: SocketAddrV4) -> Result<()> { chunks_info.iter().take(40).map(|(hash, _)| *hash).collect(); stats.borrow_mut().requested += chunks.len(); let msg = postcard::to_allocvec(&UdpRequest::RequestChunks(chunks)).unwrap(); - socket.send(server_addr, &msg).await?; + socket.send_to(server_addr, &msg).await?; } } } @@ -221,7 +222,7 @@ pub async fn flash(os: UefiOS, server_addr: SocketAddrV4) -> Result<()> { let msg = UdpRequest::ActionProgress(stats.borrow().recv, stats.borrow().fetch); socket - .send(server_addr, &postcard::to_allocvec(&msg)?) + .send_to(server_addr, &postcard::to_allocvec(&msg)?) .await?; } Ok(()) diff --git a/pixie-uefi/src/main.rs b/pixie-uefi/src/main.rs index 051c3cf1..d8d878f7 100644 --- a/pixie-uefi/src/main.rs +++ b/pixie-uefi/src/main.rs @@ -14,7 +14,8 @@ use uefi::{entry, Status}; use crate::flash::flash; use crate::os::error::{Error, Result}; -use crate::os::{TcpStream, UefiOS, PACKET_SIZE}; +use crate::os::net::{TcpStream, UdpSocket, ETH_PACKET_SIZE}; +use crate::os::UefiOS; use crate::reboot_to_os::reboot_to_os; use crate::register::register; use crate::store::store; @@ -33,22 +34,22 @@ mod export_cov; const MIN_MEMORY: u64 = 32 << 20; async fn server_discover(os: UefiOS) -> Result { - let socket = os.udp_bind(None).await?; + let socket = UdpSocket::bind(None).await?; let task1 = async { let msg = postcard::to_allocvec(&UdpRequest::Discover).unwrap(); #[allow(unreachable_code)] Ok::<_, Error>(loop { socket - .send(SocketAddrV4::new(Ipv4Addr::BROADCAST, ACTION_PORT), &msg) + .send_to(SocketAddrV4::new(Ipv4Addr::BROADCAST, ACTION_PORT), &msg) .await?; os.sleep_us(1_000_000).await; }) }; let task2 = async { - let mut buf = [0; PACKET_SIZE]; - let (data, server) = socket.recv(&mut buf).await; + let mut buf = [0; ETH_PACKET_SIZE]; + let (data, server) = socket.recv_from(&mut buf).await; assert_eq!(data.len(), 0); Ok::<_, Error>(server) }; @@ -75,22 +76,22 @@ async fn shutdown(os: UefiOS) -> ! { async fn get_action(stream: &TcpStream) -> Result { let msg = postcard::to_allocvec(&TcpRequest::GetAction)?; - stream.send_u64_le(msg.len() as u64).await?; - stream.send(&msg).await?; + stream.write_u64_le(msg.len() as u64).await?; + stream.write_all(&msg).await?; - let len = stream.recv_u64_le().await? as usize; + let len = stream.read_u64_le().await? as usize; let mut buf = vec![0; len]; - stream.recv_exact(&mut buf).await?; + stream.read_exact(&mut buf).await?; let cmd = postcard::from_bytes(&buf)?; Ok(cmd) } async fn complete_action(stream: &TcpStream) -> Result<()> { let msg = postcard::to_allocvec(&TcpRequest::ActionComplete)?; - stream.send_u64_le(msg.len() as u64).await?; - stream.send(&msg).await?; + stream.write_u64_le(msg.len() as u64).await?; + stream.write_all(&msg).await?; - let len = stream.recv_u64_le().await?; + let len = stream.read_u64_le().await?; assert_eq!(len, 0); Ok(()) } @@ -101,10 +102,10 @@ async fn run(os: UefiOS) -> Result<()> { let mut last_was_wait = false; os.spawn("ping", async move { - let udp_socket = os.udp_bind(None).await.unwrap(); + let udp_socket = UdpSocket::bind(None).await.unwrap(); loop { udp_socket - .send(SocketAddrV4::new(*server.ip(), PING_PORT), b"pixie") + .send_to(SocketAddrV4::new(*server.ip(), PING_PORT), b"pixie") .await .unwrap(); os.sleep_us(10_000_000).await; @@ -119,9 +120,9 @@ async fn run(os: UefiOS) -> Result<()> { log::debug!("Sending request for command"); } - let tcp = os.connect(server).await?; + let tcp = TcpStream::connect(server).await?; let command = get_action(&tcp).await; - tcp.close_send().await; + tcp.shutdown().await; tcp.force_close().await; if let Err(e) = command { @@ -151,9 +152,9 @@ async fn run(os: UefiOS) -> Result<()> { Action::Flash => flash(os, server).await?, } - let tcp = os.connect(server).await?; + let tcp = TcpStream::connect(server).await?; complete_action(&tcp).await?; - tcp.close_send().await; + tcp.shutdown().await; tcp.force_close().await; if command == Action::Restart { diff --git a/pixie-uefi/src/os/mod.rs b/pixie-uefi/src/os/mod.rs index d9d918d1..8dd2cd18 100644 --- a/pixie-uefi/src/os/mod.rs +++ b/pixie-uefi/src/os/mod.rs @@ -3,11 +3,10 @@ use alloc::collections::VecDeque; use alloc::string::{String, ToString}; use alloc::sync::Arc; use alloc::vec::Vec; -use core::cell::{Ref, RefMut}; +use core::cell::RefMut; use core::ffi::c_void; use core::fmt::Write; use core::future::{poll_fn, Future}; -use core::net::SocketAddrV4; use core::ptr::NonNull; use core::task::Poll; @@ -15,16 +14,11 @@ use pixie_shared::util::BytesFmt; use uefi::boot::{EventType, ScopedProtocol, TimerTrigger, Tpl}; use uefi::proto::console::serial::Serial; use uefi::proto::console::text::{Color, Input, Key, Output}; -use uefi::proto::device_path::build::DevicePathBuilder; -use uefi::proto::device_path::text::{AllowShortcuts, DevicePathToText, DisplayOnly}; -use uefi::proto::device_path::DevicePath; -use uefi::proto::Protocol; -use uefi::{Event, Handle, Status}; +use uefi::{Event, Status}; use self::disk::Disk; use self::error::Result; use self::executor::{Executor, Task}; -use self::net::NetworkInterface; use self::sync::SyncRefCell; use self::timer::Timer; @@ -33,18 +27,15 @@ pub mod disk; pub mod error; mod executor; pub mod memory; -mod net; +pub mod net; mod sync; mod timer; -pub use net::{TcpStream, UdpHandle, PACKET_SIZE}; - struct UefiOSImpl { tasks: Vec>, input: ScopedProtocol, vga: ScopedProtocol, serial: Option>, - net: Option, messages: VecDeque<(f64, log::Level, String, String)>, ui_buf: Vec<(String, Color, Color)>, ui_pos: usize, @@ -163,7 +154,6 @@ impl UefiOS { input, vga, serial, - net: None, messages: VecDeque::new(), ui_buf: vec![], ui_pos: 0, @@ -175,8 +165,7 @@ impl UefiOS { log::set_logger(&UefiOS { cant_build: () }).unwrap(); log::set_max_level(log::LevelFilter::Trace); - let net = NetworkInterface::new(os); - os.borrow_mut().net = Some(net); + net::init(); os.spawn("init", async move { loop { @@ -202,37 +191,6 @@ impl UefiOS { } }); - os.spawn( - "[net_poll]", - poll_fn(move |cx| { - let mut os = os.borrow_mut(); - os.net.as_mut().unwrap().poll(); - // TODO(veluca): figure out whether we can suspend the task. - cx.waker().wake_by_ref(); - Poll::Pending - }), - ); - - os.spawn("[net_speed]", async move { - let mut prx = 0; - let mut ptx = 0; - let mut ptm = Timer::instant(); - loop { - { - let now = Timer::instant(); - let dt = (now - ptm).total_micros() as f64 / 1_000_000.0; - ptm = now; - - let mut net = os.net(); - net.vrx = ((net.rx - prx) as f64 / dt) as u64; - prx = net.rx; - net.vtx = ((net.tx - ptx) as f64 / dt) as u64; - ptx = net.tx; - } - os.sleep_us(1_000_000).await; - } - }); - os.spawn("[draw_ui]", async move { loop { os.draw_ui(); @@ -243,10 +201,6 @@ impl UefiOS { Executor::run() } - fn borrow(&self) -> Ref<'static, UefiOSImpl> { - Ref::map(OS.borrow(), |f| f.as_ref().unwrap()) - } - fn borrow_mut(&self) -> RefMut<'static, UefiOSImpl> { RefMut::map(OS.borrow_mut(), |f| f.as_mut().unwrap()) } @@ -255,21 +209,6 @@ impl UefiOS { RefMut::map(self.borrow_mut(), |f| &mut f.tasks) } - pub fn net(&self) -> RefMut<'static, NetworkInterface> { - RefMut::map(self.borrow_mut(), |f| f.net.as_mut().unwrap()) - } - - pub fn wait_for_ip(self) -> impl Future { - poll_fn(move |cx| { - if self.net().has_ip() { - Poll::Ready(()) - } else { - cx.waker().wake_by_ref(); - Poll::Pending - } - }) - } - /// Interrupt task execution. /// This is useful to yield the CPU to other tasks. pub fn schedule(&self) -> impl Future { @@ -308,44 +247,10 @@ impl UefiOS { uefi::boot::wait_for_event(&mut [e]).unwrap(); } - pub fn device_path_to_string(&self, device: &DevicePath) -> String { - let handle = uefi::boot::get_handle_for_protocol::().unwrap(); - let device_path_to_text = - uefi::boot::open_protocol_exclusive::(handle).unwrap(); - device_path_to_text - .convert_device_path_to_text(device, DisplayOnly(true), AllowShortcuts(true)) - .unwrap() - .to_string() - } - - /// Find the topmost device that implements this protocol. - fn handle_on_device(&self, device: &DevicePath) -> Option { - for i in 0..device.node_iter().count() { - let mut buf = vec![]; - let mut dev = DevicePathBuilder::with_vec(&mut buf); - for node in device.node_iter().take(i + 1) { - dev = dev.push(&node).unwrap(); - } - let mut dev = dev.finalize().unwrap(); - if let Ok(h) = uefi::boot::locate_device_path::

(&mut dev) { - return Some(h); - } - } - None - } - pub fn open_first_disk(&self) -> Disk { Disk::new(*self) } - pub async fn connect(&self, addr: SocketAddrV4) -> Result { - TcpStream::new(*self, addr).await - } - - pub async fn udp_bind(&self, port: Option) -> Result { - UdpHandle::new(*self, port).await - } - pub fn read_key(&self) -> impl Future> + '_ { poll_fn(move |cx| { let key = self.borrow_mut().input.read_key(); @@ -369,7 +274,7 @@ impl UefiOS { // Write the header. { let time = Timer::micros() as f32 * 0.000_001; - let ip = self.net().ip(); + let ip = net::ip(); let mut os = self.borrow_mut(); let mode = os.vga.current_mode().unwrap().unwrap(); @@ -386,8 +291,7 @@ impl UefiOS { os.maybe_advance_to_col(3 * cols / 5); - let vrx = os.net.as_ref().unwrap().vrx; - let vtx = os.net.as_ref().unwrap().vtx; + let (vtx, vrx) = net::speed(); os.write_with_color( &format!("rx: {}/s tx: {}/s\n\n", BytesFmt(vrx), BytesFmt(vtx)), Color::White, @@ -437,10 +341,6 @@ impl UefiOS { } pub fn force_ui_redraw(&self) { - // TODO(virv): during network initialization we already start logging - if self.borrow().net.is_none() { - return; - } self.draw_ui() } diff --git a/pixie-uefi/src/os/net.rs b/pixie-uefi/src/os/net.rs deleted file mode 100644 index 28a108a7..00000000 --- a/pixie-uefi/src/os/net.rs +++ /dev/null @@ -1,525 +0,0 @@ -use alloc::boxed::Box; -use core::future::poll_fn; -use core::net::{IpAddr, Ipv4Addr, SocketAddrV4}; -use core::task::Poll; - -use futures::future::select; -use smoltcp::iface::{Config, Interface, PollResult, SocketHandle, SocketSet}; -use smoltcp::phy::{Device, DeviceCapabilities, Medium, RxToken, TxToken}; -use smoltcp::socket::dhcpv4::{Event, Socket as Dhcpv4Socket}; -use smoltcp::socket::tcp::{Socket as TcpSocket, State}; -use smoltcp::socket::udp::{self, Socket as UdpSocket}; -use smoltcp::storage::RingBuffer; -use smoltcp::time::{Duration, Instant}; -use smoltcp::wire::{DhcpOption, HardwareAddress, IpCidr, IpEndpoint}; -use uefi::boot::ScopedProtocol; -use uefi::proto::network::snp::{ReceiveFlags, SimpleNetwork}; -use uefi::Status; - -use super::error::{Error, Result}; -use super::timer::Timer; -use super::UefiOS; -use crate::os::boot_options::BootOptions; -use crate::os::timer::rdtsc; - -pub const PACKET_SIZE: usize = 1514; - -type Snp = &'static ScopedProtocol; - -struct SnpDevice { - snp: Snp, - tx_buf: [u8; PACKET_SIZE], - // Received packets might contain Ethernet-related padding (up to 4 bytes). - rx_buf: [u8; PACKET_SIZE + 4], -} - -impl SnpDevice { - fn new(snp: Snp) -> SnpDevice { - // Shut down the SNP protocol if needed. - let _ = snp.shutdown(); - let _ = snp.stop(); - // Initialize. - snp.start().unwrap(); - snp.initialize(0, 0).unwrap(); - // Enable packet reception. - snp.receive_filters( - ReceiveFlags::UNICAST | ReceiveFlags::BROADCAST, - ReceiveFlags::empty(), - true, - None, - ) - .unwrap(); - - SnpDevice { - snp, - tx_buf: [0; PACKET_SIZE], - rx_buf: [0; PACKET_SIZE + 4], - } - } -} - -impl Drop for SnpDevice { - fn drop(&mut self) { - self.snp.stop().unwrap() - } -} - -struct SnpRxToken<'a> { - packet: &'a mut [u8], -} - -struct SnpTxToken<'a> { - snp: Snp, - buf: &'a mut [u8], -} - -impl TxToken for SnpTxToken<'_> { - fn consume(self, len: usize, f: F) -> R - where - F: FnOnce(&mut [u8]) -> R, - { - assert!(len <= self.buf.len()); - let payload = &mut self.buf[..len]; - let ret = f(payload); - let snp = self.snp; - snp.transmit(0, payload, None, None, None) - .expect("Failed to transmit frame"); - // Wait until sending is complete. - while snp.get_recycled_transmit_buffer_status().unwrap().is_none() {} - ret - } -} - -impl RxToken for SnpRxToken<'_> { - fn consume(self, f: F) -> R - where - F: FnOnce(&[u8]) -> R, - { - f(self.packet) - } -} - -impl Device for SnpDevice { - type TxToken<'d> = SnpTxToken<'d>; - type RxToken<'d> = SnpRxToken<'d>; - - fn receive(&mut self, _: Instant) -> Option<(SnpRxToken<'_>, SnpTxToken<'_>)> { - let rec = self.snp.receive(&mut self.rx_buf, None, None, None, None); - if rec == Err(Status::NOT_READY.into()) { - return None; - } - Some(( - SnpRxToken { - packet: &mut self.rx_buf[..rec.unwrap()], - }, - SnpTxToken { - snp: self.snp, - buf: &mut self.tx_buf, - }, - )) - } - - fn transmit(&mut self, _: Instant) -> Option> { - Some(SnpTxToken { - snp: self.snp, - buf: &mut self.tx_buf, - }) - } - - fn capabilities(&self) -> DeviceCapabilities { - let mut caps = DeviceCapabilities::default(); - caps.medium = Medium::Ethernet; - let mode = self.snp.mode(); - assert!(mode.media_header_size == 14); - caps.max_transmission_unit = - PACKET_SIZE.min((mode.max_packet_size + mode.media_header_size) as usize); - caps.max_burst_size = Some(1); - caps - } -} - -pub struct NetworkInterface { - interface: Interface, - device: SnpDevice, - socket_set: SocketSet<'static>, - dhcp_socket_handle: SocketHandle, - ephemeral_port_counter: u64, - pub(super) rx: u64, - pub(super) tx: u64, - pub(super) vrx: u64, - pub(super) vtx: u64, -} - -impl NetworkInterface { - pub fn new(os: UefiOS) -> NetworkInterface { - let curopt = BootOptions::get(BootOptions::current()); - let (descr, device) = BootOptions::boot_entry_info(&curopt[..]); - log::info!( - "Configuring network on interface used for booting ({} -- {})", - descr, - os.device_path_to_string(device) - ); - - let snp_handle = if let Some(handle) = os.handle_on_device::(device) { - handle - } else { - log::info!("SNP handle not found on device, falling back to first SNP handle"); - uefi::boot::find_handles::().unwrap()[0] - }; - - let snp = uefi::boot::open_protocol_exclusive::(snp_handle).unwrap(); - let mut device = SnpDevice::new(Box::leak(Box::new(snp))); - - let hw_addr = HardwareAddress::Ethernet(smoltcp::wire::EthernetAddress::from_bytes( - &device.snp.mode().current_address.0[..6], - )); - - let mut interface_config = Config::new(hw_addr); - interface_config.random_seed = rdtsc() as u64; - let now = Timer::instant(); - let interface = Interface::new(interface_config, &mut device, now); - let mut dhcp_socket = Dhcpv4Socket::new(); - dhcp_socket.set_outgoing_options(&[DhcpOption { - kind: 60, - data: b"pixie", - }]); - let mut socket_set = SocketSet::new(vec![]); - let dhcp_socket_handle = socket_set.add(dhcp_socket); - - NetworkInterface { - interface, - dhcp_socket_handle, - device, - socket_set, - ephemeral_port_counter: 0, - rx: 0, - tx: 0, - vrx: 0, - vtx: 0, - } - } - - fn get_ephemeral_port(&mut self) -> u16 { - let ans = self.ephemeral_port_counter; - self.ephemeral_port_counter += 1; - ((ans % (60999 - 49152)) + 49152) as u16 - } - - pub fn has_ip(&self) -> bool { - self.ip().is_some() - } - - pub fn ip(&self) -> Option { - self.interface.ipv4_addr() - } - - pub(super) fn poll(&mut self) -> bool { - let now = Timer::instant(); - let status = self - .interface - .poll(now, &mut self.device, &mut self.socket_set); - if status == PollResult::None { - return false; - } - - let dhcp_status = self - .socket_set - .get_mut::(self.dhcp_socket_handle) - .poll(); - - if let Some(dhcp_status) = dhcp_status { - if let Event::Configured(config) = dhcp_status { - self.interface.update_ip_addrs(|a| { - a.push(IpCidr::Ipv4(config.address)).unwrap(); - }); - if let Some(router) = config.router { - self.interface - .routes_mut() - .add_default_ipv4_route(router) - .unwrap(); - } - } else { - self.interface.update_ip_addrs(|a| { - a.clear(); - }); - self.interface.routes_mut().remove_default_ipv4_route(); - } - } - - true - } -} - -pub struct TcpStream { - handle: SocketHandle, - os: UefiOS, -} - -// TODO(veluca): we may leak a fair bit of sockets here. It doesn't really matter, as we won't -// create that many, but still it would be nice to fix eventually. -// Also, trying to use a closed connection may result in panics. -impl TcpStream { - pub async fn new(os: UefiOS, addr: SocketAddrV4) -> Result { - os.wait_for_ip().await; - const TCP_BUF_SIZE: usize = 1 << 22; - let mut tcp_socket = TcpSocket::new( - RingBuffer::new(vec![0; TCP_BUF_SIZE]), - RingBuffer::new(vec![0; TCP_BUF_SIZE]), - ); - tcp_socket.set_congestion_control(smoltcp::socket::tcp::CongestionControl::Cubic); - tcp_socket.set_timeout(Some(Duration::from_secs(5))); - tcp_socket.set_keep_alive(Some(Duration::from_secs(1))); - let sport = os.net().get_ephemeral_port(); - tcp_socket.connect( - os.net().interface.context(), - IpEndpoint { - addr: (*addr.ip()).into(), - port: addr.port(), - }, - sport, - )?; - - let handle = os.net().socket_set.add(tcp_socket); - - let ret = TcpStream { handle, os }; - - ret.wait_for_state(|state| match state { - State::Established => Some(Ok(())), - State::Closed => Some(Err(Error::msg("connection refused"))), - _ => None, - }) - .await?; - - Ok(ret) - } - - pub async fn wait_for_state(&self, f: impl Fn(State) -> Option) -> T { - poll_fn(move |cx| { - let state = self - .os - .net() - .socket_set - .get_mut::(self.handle) - .state(); - let res = f(state); - if let Some(res) = res { - Poll::Ready(res) - } else { - cx.waker().wake_by_ref(); - Poll::Pending - } - }) - .await - } - - pub async fn wait_until_closed(&self) { - self.wait_for_state(|s| if s == State::Closed { Some(()) } else { None }) - .await; - self.os.net().socket_set.remove(self.handle); - } - - async fn fail_if_closed(&self) -> Result<()> { - self.wait_until_closed().await; - Err(Error::msg("connection closed")) - } - - pub async fn send(&self, data: &[u8]) -> Result<()> { - if data.is_empty() { - return Ok(()); - } - - let mut pos = 0; - let send = poll_fn(move |cx| { - let mut net = self.os.net(); - let socket = net.socket_set.get_mut::(self.handle); - let sent = socket.send_slice(&data[pos..]); - if let Err(err) = sent { - return Poll::Ready(Err(err.into())); - } - pos += sent.unwrap(); - if pos < data.len() { - socket.register_send_waker(cx.waker()); - net.tx += sent.unwrap() as u64; - Poll::Pending - } else { - net.tx += sent.unwrap() as u64; - Poll::Ready(Ok(())) - } - }); - - select(send, Box::pin(self.fail_if_closed())) - .await - .factor_first() - .0 - } - - /// Returns the number of bytes received (0 if connection is closed on the other end without - /// receiving any data. - pub async fn recv(&self, data: &mut [u8]) -> Result { - poll_fn(move |cx| { - let mut net = self.os.net(); - let socket = net.socket_set.get_mut::(self.handle); - if !socket.may_recv() { - return Poll::Ready(Ok(0)); - } - let recvd = socket.recv_slice(data); - if recvd == Err(smoltcp::socket::tcp::RecvError::Finished) { - return Poll::Ready(Ok(0)); - } - if let Err(err) = recvd { - return Poll::Ready(Err(err.into())); - } - if recvd.unwrap() == 0 { - socket.register_recv_waker(cx.waker()); - Poll::Pending - } else { - net.rx += recvd.unwrap() as u64; - Poll::Ready(Ok(recvd.unwrap())) - } - }) - .await - } - - pub async fn recv_exact(&self, data: &mut [u8]) -> Result<()> { - let mut pos = 0; - while pos < data.len() { - let recvd = self.recv(&mut data[pos..]).await?; - if recvd == 0 { - return Err(Error::msg("connection closed")); - } - pos += recvd; - } - Ok(()) - } - - pub async fn send_u64_le(&self, data: u64) -> Result<()> { - self.send(&data.to_le_bytes()).await - } - - pub async fn recv_u64_le(&self) -> Result { - let mut buf = [0; 8]; - self.recv_exact(&mut buf).await?; - Ok(u64::from_le_bytes(buf)) - } - - pub async fn close_send(&self) { - { - self.os - .net() - .socket_set - .get_mut::(self.handle) - .close(); - } - self.wait_for_state(|state| match state { - State::Closed | State::Closing | State::FinWait1 | State::FinWait2 => Some(()), - _ => None, - }) - .await - } - - pub async fn force_close(&self) { - { - self.os - .net() - .socket_set - .get_mut::(self.handle) - .abort(); - } - self.wait_until_closed().await; - } -} - -pub struct UdpHandle { - handle: SocketHandle, - os: UefiOS, -} - -impl UdpHandle { - pub async fn new(os: UefiOS, listen_port: Option) -> Result { - os.wait_for_ip().await; - const UDP_BUF_SIZE: usize = 1 << 22; - const UDP_PACKET_BUF_SIZE: usize = 1 << 10; - let rx_buffer = udp::PacketBuffer::new( - vec![udp::PacketMetadata::EMPTY; UDP_PACKET_BUF_SIZE], - vec![0; UDP_BUF_SIZE], - ); - let tx_buffer = udp::PacketBuffer::new( - vec![udp::PacketMetadata::EMPTY; UDP_PACKET_BUF_SIZE], - vec![0; UDP_BUF_SIZE], - ); - - let mut udp_socket = UdpSocket::new(rx_buffer, tx_buffer); - let sport = if let Some(p) = listen_port { - p - } else { - os.net().get_ephemeral_port() - }; - udp_socket.bind(sport)?; - - let handle = os.net().socket_set.add(udp_socket); - - let ret = UdpHandle { handle, os }; - Ok(ret) - } - - pub async fn send(&self, addr: SocketAddrV4, data: &[u8]) -> Result<()> { - let endpoint = IpEndpoint { - addr: (*addr.ip()).into(), - port: addr.port(), - }; - - Ok(poll_fn(move |cx| { - let mut net = self.os.net(); - let socket = net.socket_set.get_mut::(self.handle); - if !socket.can_send() { - socket.register_send_waker(cx.waker()); - Poll::Pending - } else { - let status = socket.send_slice(data, endpoint); - net.tx = net.tx.wrapping_add(data.len() as u64); - Poll::Ready(status) - } - }) - .await?) - } - - pub async fn recv<'a>(&self, buf: &'a mut [u8; PACKET_SIZE]) -> (&'a mut [u8], SocketAddrV4) { - let buf2 = &mut *buf; - let (len, addr) = poll_fn(move |cx| { - let mut net = self.os.net(); - let socket = net.socket_set.get_mut::(self.handle); - if !socket.can_recv() { - socket.register_recv_waker(cx.waker()); - Poll::Pending - } else { - // Cannot fail if can_recv() returned true. - let recvd = socket.recv_slice(buf2).unwrap(); - let IpAddr::V4(ip) = (recvd.1).endpoint.addr.into() else { - unreachable!(); - }; - let port = (recvd.1).endpoint.port; - Poll::Ready((recvd.0, SocketAddrV4::new(ip, port))) - } - }) - .await; - - self.os.net().rx += len as u64; - - (&mut buf[..len], addr) - } - - pub fn close(&mut self) { - self.os - .net() - .socket_set - .get_mut::(self.handle) - .close(); - self.os.net().socket_set.remove(self.handle); - } -} - -impl Drop for UdpHandle { - fn drop(&mut self) { - self.close() - } -} diff --git a/pixie-uefi/src/os/net/interface.rs b/pixie-uefi/src/os/net/interface.rs new file mode 100644 index 00000000..9afe7868 --- /dev/null +++ b/pixie-uefi/src/os/net/interface.rs @@ -0,0 +1,124 @@ +use smoltcp::phy::{Device, DeviceCapabilities, Medium, RxToken, TxToken}; +use smoltcp::time::Instant; +use uefi::boot::ScopedProtocol; +use uefi::proto::network::snp::{ReceiveFlags, SimpleNetwork}; +use uefi::Status; + +use super::ETH_PACKET_SIZE; + +type Snp = ScopedProtocol; + +pub struct SnpDevice { + snp: Snp, + tx_buf: [u8; ETH_PACKET_SIZE], + // Received packets might contain Ethernet-related padding (up to 4 bytes). + rx_buf: [u8; ETH_PACKET_SIZE + 4], +} + +// SAFETY: we never create threads anyway. +unsafe impl Send for SnpDevice {} + +impl SnpDevice { + pub fn new(snp: Snp) -> SnpDevice { + // Shut down the SNP protocol if needed. + let _ = snp.shutdown(); + let _ = snp.stop(); + // Initialize. + snp.start().unwrap(); + snp.initialize(0, 0).unwrap(); + // Enable packet reception. + snp.receive_filters( + ReceiveFlags::UNICAST | ReceiveFlags::BROADCAST, + ReceiveFlags::empty(), + true, + None, + ) + .unwrap(); + + SnpDevice { + snp, + tx_buf: [0; ETH_PACKET_SIZE], + rx_buf: [0; ETH_PACKET_SIZE + 4], + } + } +} + +impl Drop for SnpDevice { + fn drop(&mut self) { + self.snp.stop().unwrap() + } +} + +pub struct SnpRxToken<'a> { + packet: &'a mut [u8], +} + +pub struct SnpTxToken<'a> { + snp: &'a Snp, + buf: &'a mut [u8], +} + +impl TxToken for SnpTxToken<'_> { + fn consume(self, len: usize, f: F) -> R + where + F: FnOnce(&mut [u8]) -> R, + { + assert!(len <= self.buf.len()); + let payload = &mut self.buf[..len]; + let ret = f(payload); + let snp = self.snp; + snp.transmit(0, payload, None, None, None) + .expect("Failed to transmit frame"); + // Wait until sending is complete. + while snp.get_recycled_transmit_buffer_status().unwrap().is_none() {} + ret + } +} + +impl RxToken for SnpRxToken<'_> { + fn consume(self, f: F) -> R + where + F: FnOnce(&[u8]) -> R, + { + f(self.packet) + } +} + +impl Device for SnpDevice { + type TxToken<'d> = SnpTxToken<'d>; + type RxToken<'d> = SnpRxToken<'d>; + + fn receive(&mut self, _: Instant) -> Option<(SnpRxToken<'_>, SnpTxToken<'_>)> { + let rec = self.snp.receive(&mut self.rx_buf, None, None, None, None); + if rec == Err(Status::NOT_READY.into()) { + return None; + } + Some(( + SnpRxToken { + packet: &mut self.rx_buf[..rec.unwrap()], + }, + SnpTxToken { + snp: &self.snp, + buf: &mut self.tx_buf, + }, + )) + } + + fn transmit(&mut self, _: Instant) -> Option> { + Some(SnpTxToken { + snp: &self.snp, + buf: &mut self.tx_buf, + }) + } + + fn capabilities(&self) -> DeviceCapabilities { + let mut caps = DeviceCapabilities::default(); + caps.medium = Medium::Ethernet; + let mode = self.snp.mode(); + assert!(mode.media_header_size == 14); + caps.max_transmission_unit = + ETH_PACKET_SIZE.min((mode.max_packet_size + mode.media_header_size) as usize); + caps.max_burst_size = Some(1); + caps + } +} diff --git a/pixie-uefi/src/os/net/mod.rs b/pixie-uefi/src/os/net/mod.rs new file mode 100644 index 00000000..9f3035dd --- /dev/null +++ b/pixie-uefi/src/os/net/mod.rs @@ -0,0 +1,203 @@ +use alloc::string::{String, ToString}; +use core::future::{poll_fn, Future}; +use core::net::Ipv4Addr; +use core::sync::atomic::{AtomicU64, Ordering}; +use core::task::Poll; + +use smoltcp::iface::{Config, Interface, PollResult, SocketHandle, SocketSet}; +use smoltcp::socket::dhcpv4::{Event, Socket as Dhcpv4Socket}; +use smoltcp::wire::{DhcpOption, HardwareAddress, IpCidr}; +use spin::Mutex; +use uefi::proto::device_path::build::DevicePathBuilder; +use uefi::proto::device_path::text::{AllowShortcuts, DevicePathToText, DisplayOnly}; +use uefi::proto::device_path::DevicePath; +use uefi::proto::network::snp::SimpleNetwork; +use uefi::proto::Protocol; +use uefi::Handle; + +use super::timer::Timer; +use crate::os::boot_options::BootOptions; +use crate::os::net::interface::SnpDevice; +use crate::os::net::speed::{RX_SPEED, TX_SPEED}; +pub use crate::os::net::tcp::TcpStream; +pub use crate::os::net::udp::UdpSocket; +use crate::os::timer::rdtsc; +use crate::os::UefiOS; + +mod interface; +mod speed; +mod tcp; +mod udp; + +pub const ETH_PACKET_SIZE: usize = 1514; + +static EPHEMERAL_PORT_COUNTER: AtomicU64 = AtomicU64::new(0); + +struct NetworkData { + interface: Interface, + device: SnpDevice, + socket_set: SocketSet<'static>, + dhcp_socket_handle: SocketHandle, +} + +static NETWORK_DATA: Mutex> = Mutex::new(None); + +pub fn speed() -> (u64, u64) { + (TX_SPEED.bytes_per_second(), RX_SPEED.bytes_per_second()) +} + +fn with_net T>(f: F) -> T { + let mut mg = NETWORK_DATA.try_lock().expect("Network is locked"); + f(mg.as_mut().expect("Network is not initialized")) +} + +fn device_path_to_string(device: &DevicePath) -> String { + let handle = uefi::boot::get_handle_for_protocol::().unwrap(); + let device_path_to_text = + uefi::boot::open_protocol_exclusive::(handle).unwrap(); + device_path_to_text + .convert_device_path_to_text(device, DisplayOnly(true), AllowShortcuts(true)) + .unwrap() + .to_string() +} + +/// Find the topmost device that implements this protocol. +fn handle_on_device(device: &DevicePath) -> Option { + for i in 0..device.node_iter().count() { + let mut buf = vec![]; + let mut dev = DevicePathBuilder::with_vec(&mut buf); + for node in device.node_iter().take(i + 1) { + dev = dev.push(&node).unwrap(); + } + let mut dev = dev.finalize().unwrap(); + if let Ok(h) = uefi::boot::locate_device_path::

(&mut dev) { + return Some(h); + } + } + None +} + +pub(super) fn init() { + // TODO(veluca): remove the use of `os` once we move spawning to the scheduler. + let os = UefiOS { cant_build: () }; + + let curopt = BootOptions::get(BootOptions::current()); + let (descr, device) = BootOptions::boot_entry_info(&curopt[..]); + log::info!( + "Configuring network on interface used for booting ({} -- {})", + descr, + device_path_to_string(device) + ); + + let snp_handle = if let Some(handle) = handle_on_device::(device) { + handle + } else { + log::info!("SNP handle not found on device, falling back to first SNP handle"); + uefi::boot::find_handles::().unwrap()[0] + }; + + let snp = uefi::boot::open_protocol_exclusive::(snp_handle).unwrap(); + + let hw_addr = HardwareAddress::Ethernet(smoltcp::wire::EthernetAddress::from_bytes( + &snp.mode().current_address.0[..6], + )); + + let mut device = SnpDevice::new(snp); + + let mut interface_config = Config::new(hw_addr); + interface_config.random_seed = rdtsc() as u64; + let now = Timer::instant(); + let interface = Interface::new(interface_config, &mut device, now); + let mut dhcp_socket = Dhcpv4Socket::new(); + dhcp_socket.set_outgoing_options(&[DhcpOption { + kind: 60, + data: b"pixie", + }]); + let mut socket_set = SocketSet::new(vec![]); + let dhcp_socket_handle = socket_set.add(dhcp_socket); + + *NETWORK_DATA.lock() = Some(NetworkData { + interface, + device, + socket_set, + dhcp_socket_handle, + }); + + os.spawn( + "[net_poll]", + poll_fn(move |cx| { + poll(); + // TODO(veluca): figure out whether we can suspend the task. + cx.waker().wake_by_ref(); + Poll::Pending + }), + ); + + speed::spawn_update_network_speed_task(os); +} + +pub fn wait_for_ip() -> impl Future { + poll_fn(move |cx| { + if ip().is_some() { + Poll::Ready(()) + } else { + cx.waker().wake_by_ref(); + Poll::Pending + } + }) +} + +pub fn ip() -> Option { + let mg = NETWORK_DATA.try_lock().unwrap(); + mg.as_ref().and_then(|n| n.interface.ipv4_addr()) +} + +fn get_ephemeral_port() -> u16 { + let ans = EPHEMERAL_PORT_COUNTER.fetch_add(1, Ordering::Relaxed); + ((ans % (60999 - 49152)) + 49152) as u16 +} + +fn poll() { + let now = Timer::instant(); + + let mut data = NETWORK_DATA.lock(); + + let Some(NetworkData { + interface, + device, + socket_set, + dhcp_socket_handle, + }) = data.as_mut() + else { + return; + }; + + let status = interface.poll(now, device, socket_set); + + if status == PollResult::None { + return; + } + + let dhcp_status = socket_set + .get_mut::(*dhcp_socket_handle) + .poll(); + + if let Some(dhcp_status) = dhcp_status { + if let Event::Configured(config) = dhcp_status { + interface.update_ip_addrs(|a| { + a.push(IpCidr::Ipv4(config.address)).unwrap(); + }); + if let Some(router) = config.router { + interface + .routes_mut() + .add_default_ipv4_route(router) + .unwrap(); + } + } else { + interface.update_ip_addrs(|a| { + a.clear(); + }); + interface.routes_mut().remove_default_ipv4_route(); + } + } +} diff --git a/pixie-uefi/src/os/net/speed.rs b/pixie-uefi/src/os/net/speed.rs new file mode 100644 index 00000000..79555288 --- /dev/null +++ b/pixie-uefi/src/os/net/speed.rs @@ -0,0 +1,55 @@ +use core::sync::atomic::{AtomicU64, Ordering}; + +use crate::os::timer::Timer; +use crate::os::UefiOS; + +pub struct NetSpeed { + total: AtomicU64, + bytes_per_second: AtomicU64, + last: AtomicU64, + last_update_micros: AtomicU64, +} + +impl NetSpeed { + const fn new() -> Self { + Self { + total: AtomicU64::new(0), + last: AtomicU64::new(0), + bytes_per_second: AtomicU64::new(0), + last_update_micros: AtomicU64::new(0), + } + } + + pub fn add_bytes(&self, count: usize) { + self.total.fetch_add(count as u64, Ordering::Relaxed); + } + + pub fn bytes_per_second(&self) -> u64 { + self.bytes_per_second.load(Ordering::Relaxed) + } + + pub fn update_speed(&self) { + let micros = Timer::micros() as u64; + let total = self.total.load(Ordering::Relaxed); + let last_micros = self.last_update_micros.swap(micros, Ordering::Relaxed); + let last = self.last.swap(total, Ordering::Relaxed); + let elapsed = micros.saturating_sub(last_micros).max(1); + let bytes = total.saturating_sub(last); + let bytes_per_second = bytes * 1_000_000 / elapsed; + self.bytes_per_second + .store(bytes_per_second, Ordering::Relaxed); + } +} + +pub(super) static TX_SPEED: NetSpeed = NetSpeed::new(); +pub(super) static RX_SPEED: NetSpeed = NetSpeed::new(); + +pub(super) fn spawn_update_network_speed_task(os: UefiOS) { + os.spawn("[net_speed]", async move { + loop { + TX_SPEED.update_speed(); + RX_SPEED.update_speed(); + os.sleep_us(1_000_000).await; + } + }); +} diff --git a/pixie-uefi/src/os/net/tcp.rs b/pixie-uefi/src/os/net/tcp.rs new file mode 100644 index 00000000..1afc9b9b --- /dev/null +++ b/pixie-uefi/src/os/net/tcp.rs @@ -0,0 +1,185 @@ +use alloc::boxed::Box; +use core::future::{poll_fn, Future}; +use core::net::SocketAddrV4; +use core::task::Poll; + +use futures::future::select; +use smoltcp::iface::SocketHandle; +use smoltcp::socket::tcp::{Socket as TcpSocket, State}; +use smoltcp::storage::RingBuffer; +use smoltcp::time::Duration; +use smoltcp::wire::IpEndpoint; + +use crate::os::error::{Error, Result}; +use crate::os::net::speed::{RX_SPEED, TX_SPEED}; +use crate::os::net::with_net; + +pub struct TcpStream { + handle: SocketHandle, +} + +// TODO(veluca): we may leak a fair bit of sockets here. It doesn't really matter, as we won't +// create that many, but still it would be nice to fix eventually. +// Also, trying to use a closed connection may result in panics. +impl TcpStream { + pub async fn connect(addr: SocketAddrV4) -> Result { + super::wait_for_ip().await; + const TCP_BUF_SIZE: usize = 1 << 22; + let mut tcp_socket = TcpSocket::new( + RingBuffer::new(vec![0; TCP_BUF_SIZE]), + RingBuffer::new(vec![0; TCP_BUF_SIZE]), + ); + tcp_socket.set_congestion_control(smoltcp::socket::tcp::CongestionControl::Cubic); + tcp_socket.set_timeout(Some(Duration::from_secs(5))); + tcp_socket.set_keep_alive(Some(Duration::from_secs(1))); + let sport = super::get_ephemeral_port(); + let handle = with_net(|net| { + tcp_socket.connect( + net.interface.context(), + IpEndpoint { + addr: (*addr.ip()).into(), + port: addr.port(), + }, + sport, + )?; + + Ok::<_, Error>(net.socket_set.add(tcp_socket)) + })?; + + let ret = TcpStream { handle }; + + ret.wait_for_state(|state| match state { + State::Established => Poll::Ready(Ok(())), + State::Closed => Poll::Ready(Err(Error::msg("connection refused"))), + _ => Poll::Pending, + }) + .await?; + Ok(ret) + } + + fn wait_for_state<'a, T>( + &'a self, + f: impl Fn(State) -> Poll + 'a, + ) -> impl Future + 'a { + poll_fn(move |cx| { + let state = with_net(|n| n.socket_set.get_mut::(self.handle).state()); + let res = f(state); + if matches!(res, Poll::Pending) { + cx.waker().wake_by_ref(); + } + res + }) + } + + async fn wait_until_closed(&self) { + self.wait_for_state(|s| { + if s == State::Closed { + Poll::Ready(()) + } else { + Poll::Pending + } + }) + .await; + with_net(|n| n.socket_set.remove(self.handle)); + } + + async fn fail_if_closed(&self) -> Result<()> { + self.wait_until_closed().await; + Err(Error::msg("connection closed")) + } + + pub async fn write_all(&self, data: &[u8]) -> Result<()> { + if data.is_empty() { + return Ok(()); + } + + let mut pos = 0; + let send = poll_fn(move |cx| { + with_net(|net| { + let socket = net.socket_set.get_mut::(self.handle); + let sent = socket.send_slice(&data[pos..]); + if let Err(err) = sent { + return Poll::Ready(Err(err.into())); + } + pos += sent.unwrap(); + TX_SPEED.add_bytes(sent.unwrap()); + if pos < data.len() { + socket.register_send_waker(cx.waker()); + Poll::Pending + } else { + Poll::Ready(Ok(())) + } + }) + }); + + select(send, Box::pin(self.fail_if_closed())) + .await + .factor_first() + .0 + } + + /// Returns the number of bytes received (0 if connection is closed on the other end without + /// receiving any data. + pub fn read<'a>(&'a self, data: &'a mut [u8]) -> impl Future> + 'a { + poll_fn(move |cx| { + with_net(|net| { + let socket = net.socket_set.get_mut::(self.handle); + if !socket.may_recv() { + return Poll::Ready(Ok(0)); + } + let recvd = socket.recv_slice(data); + if recvd == Err(smoltcp::socket::tcp::RecvError::Finished) { + return Poll::Ready(Ok(0)); + } + if let Err(err) = recvd { + return Poll::Ready(Err(err.into())); + } + if recvd.unwrap() == 0 { + socket.register_recv_waker(cx.waker()); + Poll::Pending + } else { + RX_SPEED.add_bytes(recvd.unwrap()); + Poll::Ready(Ok(recvd.unwrap())) + } + }) + }) + } + + pub async fn read_exact(&self, data: &mut [u8]) -> Result<()> { + let mut pos = 0; + while pos < data.len() { + let recvd = self.read(&mut data[pos..]).await?; + if recvd == 0 { + return Err(Error::msg("connection closed")); + } + pos += recvd; + } + Ok(()) + } + + pub async fn write_u64_le(&self, data: u64) -> Result<()> { + self.write_all(&data.to_le_bytes()).await + } + + pub async fn read_u64_le(&self) -> Result { + let mut buf = [0; 8]; + self.read_exact(&mut buf).await?; + Ok(u64::from_le_bytes(buf)) + } + + pub async fn shutdown(&self) { + with_net(|n| { + n.socket_set.get_mut::(self.handle).close(); + }); + self.wait_for_state(|state| match state { + State::Closed | State::Closing | State::FinWait1 | State::FinWait2 => Poll::Ready(()), + _ => Poll::Pending, + }) + .await + } + + pub async fn force_close(self) { + with_net(|n| n.socket_set.get_mut::(self.handle).abort()); + self.wait_until_closed().await; + } +} diff --git a/pixie-uefi/src/os/net/udp.rs b/pixie-uefi/src/os/net/udp.rs new file mode 100644 index 00000000..2cfdf62e --- /dev/null +++ b/pixie-uefi/src/os/net/udp.rs @@ -0,0 +1,102 @@ +use core::future::{poll_fn, Future}; +use core::net::{IpAddr, SocketAddrV4}; +use core::task::Poll; + +use smoltcp::iface::SocketHandle; +use smoltcp::socket::udp::{self, Socket}; +use smoltcp::wire::IpEndpoint; + +use crate::os::error::Result; +use crate::os::net::speed::{RX_SPEED, TX_SPEED}; +use crate::os::net::{with_net, ETH_PACKET_SIZE}; + +pub struct UdpSocket { + handle: SocketHandle, +} + +impl UdpSocket { + pub async fn bind(listen_port: Option) -> Result { + super::wait_for_ip().await; + const UDP_BUF_SIZE: usize = 1 << 22; + const UDP_PACKET_BUF_SIZE: usize = 1 << 10; + let rx_buffer = udp::PacketBuffer::new( + vec![udp::PacketMetadata::EMPTY; UDP_PACKET_BUF_SIZE], + vec![0; UDP_BUF_SIZE], + ); + let tx_buffer = udp::PacketBuffer::new( + vec![udp::PacketMetadata::EMPTY; UDP_PACKET_BUF_SIZE], + vec![0; UDP_BUF_SIZE], + ); + + let mut udp_socket = Socket::new(rx_buffer, tx_buffer); + let sport = listen_port.unwrap_or_else(super::get_ephemeral_port); + udp_socket.bind(sport)?; + + let handle = with_net(|n| n.socket_set.add(udp_socket)); + + Ok(UdpSocket { handle }) + } + + pub fn send_to<'a>( + &'a self, + addr: SocketAddrV4, + data: &'a [u8], + ) -> impl Future> + 'a { + let endpoint = IpEndpoint { + addr: (*addr.ip()).into(), + port: addr.port(), + }; + + poll_fn(move |cx| { + with_net(|net| { + let socket = net.socket_set.get_mut::(self.handle); + if !socket.can_send() { + socket.register_send_waker(cx.waker()); + Poll::Pending + } else { + let status = socket.send_slice(data, endpoint); + TX_SPEED.add_bytes(data.len()); + Poll::Ready(status.map_err(|e| e.into())) + } + }) + }) + } + + pub async fn recv_from<'a>( + &self, + buf: &'a mut [u8; ETH_PACKET_SIZE], + ) -> (&'a mut [u8], SocketAddrV4) { + let buf2 = &mut *buf; + let (len, addr) = poll_fn(move |cx| { + with_net(|net| { + let socket = net.socket_set.get_mut::(self.handle); + if !socket.can_recv() { + socket.register_recv_waker(cx.waker()); + Poll::Pending + } else { + // Cannot fail if can_recv() returned true. + let recvd = socket.recv_slice(buf2).unwrap(); + let IpAddr::V4(ip) = (recvd.1).endpoint.addr.into() else { + unreachable!(); + }; + let port = (recvd.1).endpoint.port; + Poll::Ready((recvd.0, SocketAddrV4::new(ip, port))) + } + }) + }) + .await; + + RX_SPEED.add_bytes(len); + + (&mut buf[..len], addr) + } +} + +impl Drop for UdpSocket { + fn drop(&mut self) { + with_net(|net| { + net.socket_set.get_mut::(self.handle).close(); + net.socket_set.remove(self.handle); + }) + } +} diff --git a/pixie-uefi/src/register.rs b/pixie-uefi/src/register.rs index 1fcb53e9..5b0dee76 100644 --- a/pixie-uefi/src/register.rs +++ b/pixie-uefi/src/register.rs @@ -9,7 +9,8 @@ use pixie_shared::{HintPacket, RegistrationInfo, TcpRequest, HINT_PORT}; use uefi::proto::console::text::{Color, Key, ScanCode}; use crate::os::error::{Error, Result}; -use crate::os::{UefiOS, PACKET_SIZE}; +use crate::os::net::{TcpStream, UdpSocket, ETH_PACKET_SIZE}; +use crate::os::UefiOS; #[derive(Debug, Default)] struct Data { @@ -61,8 +62,8 @@ pub async fn register(os: UefiOS, server_addr: SocketAddrV4) -> Result<()> { ); }); - let udp = os.udp_bind(Some(HINT_PORT)).await?; - let mut buf = [0; PACKET_SIZE]; + let udp = UdpSocket::bind(Some(HINT_PORT)).await?; + let mut buf = [0; ETH_PACKET_SIZE]; let mut hint = true; let mut images = Vec::new(); @@ -71,7 +72,7 @@ pub async fn register(os: UefiOS, server_addr: SocketAddrV4) -> Result<()> { loop { let key = if hint { loop { - let recv = Box::pin(udp.recv(&mut buf)); + let recv = Box::pin(udp.recv_from(&mut buf)); let key = Box::pin(os.read_key()); match select(recv, key).await { Either::Left(((buf, _), _)) => { @@ -161,12 +162,12 @@ pub async fn register(os: UefiOS, server_addr: SocketAddrV4) -> Result<()> { let msg = TcpRequest::Register(data.borrow().station.clone()); let buf = postcard::to_allocvec(&msg)?; - let stream = os.connect(server_addr).await?; - stream.send_u64_le(buf.len() as u64).await?; - stream.send(&buf).await?; - let len = stream.recv_u64_le().await?; + let stream = TcpStream::connect(server_addr).await?; + stream.write_u64_le(buf.len() as u64).await?; + stream.write_all(&buf).await?; + let len = stream.read_u64_le().await?; assert_eq!(len, 0); - stream.close_send().await; + stream.shutdown().await; // TODO(virv): this could be better stream.force_close().await; diff --git a/pixie-uefi/src/store.rs b/pixie-uefi/src/store.rs index 61a8e0b5..432207e1 100644 --- a/pixie-uefi/src/store.rs +++ b/pixie-uefi/src/store.rs @@ -11,7 +11,8 @@ use uefi::proto::console::text::Color; use crate::os::boot_options::BootOptions; use crate::os::error::{Error, Result}; -use crate::os::{memory, TcpStream, UefiOS}; +use crate::os::net::{TcpStream, UdpSocket}; +use crate::os::{memory, UefiOS}; use crate::{parse_disk, MIN_MEMORY}; #[derive(Debug)] @@ -23,9 +24,9 @@ pub struct ChunkInfo { async fn save_image(stream: &TcpStream, image: Image) -> Result<()> { let req = TcpRequest::UploadImage(image); let buf = postcard::to_allocvec(&req)?; - stream.send_u64_le(buf.len() as u64).await?; - stream.send(&buf).await?; - let len = stream.recv_u64_le().await?; + stream.write_u64_le(buf.len() as u64).await?; + stream.write_all(&buf).await?; + let len = stream.read_u64_le().await?; assert_eq!(len, 0); Ok(()) } @@ -76,9 +77,9 @@ pub async fn store(os: UefiOS, server_address: SocketAddrV4) -> Result<()> { BytesFmt(chunks.iter().map(|x| x.size as u64).sum::()) ); - let udp = os.udp_bind(None).await?; - let stream_get_csize = os.connect(server_address).await?; - let stream_upload_chunk = os.connect(server_address).await?; + let udp = UdpSocket::bind(None).await?; + let stream_get_csize = TcpStream::connect(server_address).await?; + let stream_upload_chunk = TcpStream::connect(server_address).await?; let total = chunks.len(); @@ -121,8 +122,8 @@ pub async fn store(os: UefiOS, server_address: SocketAddrV4) -> Result<()> { while let Some((chunk, cdata)) = rx1.recv().await { let req = TcpRequest::HasChunk(chunk.hash); let buf = postcard::to_allocvec(&req)?; - stream_get_csize.send_u64_le(buf.len() as u64).await?; - stream_get_csize.send(&buf).await?; + stream_get_csize.write_u64_le(buf.len() as u64).await?; + stream_get_csize.write_all(&buf).await?; tx2.send((chunk, cdata)).await.expect("receiver dropped"); } Ok(()) @@ -131,9 +132,9 @@ pub async fn store(os: UefiOS, server_address: SocketAddrV4) -> Result<()> { let task3 = async { let tx3 = tx3; while let Some((chunk, cdata)) = rx2.recv().await { - let len = stream_get_csize.recv_u64_le().await?; + let len = stream_get_csize.read_u64_le().await?; let mut buf = vec![0; len as usize]; - stream_get_csize.recv(&mut buf).await?; + stream_get_csize.read_exact(&mut buf).await?; let has_chunk: bool = postcard::from_bytes(&buf)?; tx3.send((chunk, cdata, has_chunk)) .await @@ -148,8 +149,8 @@ pub async fn store(os: UefiOS, server_address: SocketAddrV4) -> Result<()> { if !has_chunk { let req = TcpRequest::UploadChunk(cdata); let buf = postcard::to_allocvec(&req)?; - stream_upload_chunk.send_u64_le(buf.len() as u64).await?; - stream_upload_chunk.send(&buf).await?; + stream_upload_chunk.write_u64_le(buf.len() as u64).await?; + stream_upload_chunk.write_all(&buf).await?; } tx4.send((chunk, has_chunk)) .await @@ -169,7 +170,7 @@ pub async fn store(os: UefiOS, server_address: SocketAddrV4) -> Result<()> { let mut chunks = Vec::new(); while let Some((chunk, has_chunk)) = rx4.recv().await { if !has_chunk { - let len = stream_upload_chunk.recv_u64_le().await?; + let len = stream_upload_chunk.read_u64_le().await?; assert_eq!(len, 0); } chunks.push(chunk); @@ -183,7 +184,7 @@ pub async fn store(os: UefiOS, server_address: SocketAddrV4) -> Result<()> { tsize: total_size, tcsize: total_csize, }); - udp.send( + udp.send_to( server_address, &postcard::to_allocvec(&UdpRequest::ActionProgress(chunks.len(), total))?, ) @@ -204,10 +205,10 @@ pub async fn store(os: UefiOS, server_address: SocketAddrV4) -> Result<()> { ) .await?; - stream_get_csize.close_send().await; + stream_get_csize.shutdown().await; stream_get_csize.force_close().await; - stream_upload_chunk.close_send().await; + stream_upload_chunk.shutdown().await; // TODO(virv): this could be better stream_upload_chunk.force_close().await;