diff --git a/Cargo.lock b/Cargo.lock index 5c63565..6250de0 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -546,6 +546,7 @@ dependencies = [ "serde_yaml", "sha2", "simple-dns", + "socket2 0.5.10", "tempfile", "thiserror 2.0.18", "tokio", @@ -1402,6 +1403,16 @@ version = "1.15.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" +[[package]] +name = "socket2" +version = "0.5.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e22376abed350d73dd1cd119b57ffccad95b4e585a7cda43e286245ce23c0678" +dependencies = [ + "libc", + "windows-sys 0.52.0", +] + [[package]] name = "socket2" version = "0.6.2" @@ -1518,7 +1529,7 @@ dependencies = [ "mio", "pin-project-lite", "signal-hook-registry", - "socket2", + "socket2 0.6.2", "tokio-macros", "windows-sys 0.61.2", ] @@ -1770,6 +1781,15 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets 0.52.6", +] + [[package]] name = "windows-sys" version = "0.59.0" diff --git a/Cargo.toml b/Cargo.toml index b136847..18fcce1 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -24,6 +24,7 @@ rtnetlink = "0.14" tokio = { version = "1", features = ["rt", "macros", "signal", "sync", "net", "time"] } futures = "0.3" simple-dns = "0.9" +socket2 = { version = "0.5", features = ["all"] } [dev-dependencies] tempfile = "3.15" diff --git a/src/config/mod.rs b/src/config/mod.rs index bc2fd0d..5e51955 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -517,6 +517,7 @@ transports: {} let single = TransportInstances::Single(UdpConfig { bind_addr: Some("0.0.0.0:4000".to_string()), mtu: None, + ..Default::default() }); let items: Vec<_> = single.iter().collect(); assert_eq!(items.len(), 1); diff --git a/src/config/transport.rs b/src/config/transport.rs index 32a4b78..fbca658 100644 --- a/src/config/transport.rs +++ b/src/config/transport.rs @@ -13,6 +13,12 @@ const DEFAULT_UDP_BIND_ADDR: &str = "0.0.0.0:4000"; /// Default UDP MTU (IPv6 minimum). const DEFAULT_UDP_MTU: u16 = 1280; +/// Default UDP receive buffer size (2 MB). +const DEFAULT_UDP_RECV_BUF: usize = 2 * 1024 * 1024; + +/// Default UDP send buffer size (2 MB). +const DEFAULT_UDP_SEND_BUF: usize = 2 * 1024 * 1024; + /// UDP transport instance configuration. #[derive(Debug, Clone, Default, Serialize, Deserialize)] #[serde(deny_unknown_fields)] @@ -24,6 +30,14 @@ pub struct UdpConfig { /// UDP MTU (`mtu`). Defaults to 1280 (IPv6 minimum). #[serde(default, skip_serializing_if = "Option::is_none")] pub mtu: Option, + + /// UDP receive buffer size in bytes (`recv_buf_size`). Defaults to 2 MB. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub recv_buf_size: Option, + + /// UDP send buffer size in bytes (`send_buf_size`). Defaults to 2 MB. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub send_buf_size: Option, } impl UdpConfig { @@ -36,6 +50,16 @@ impl UdpConfig { pub fn mtu(&self) -> u16 { self.mtu.unwrap_or(DEFAULT_UDP_MTU) } + + /// Get the receive buffer size, using default if not configured. + pub fn recv_buf_size(&self) -> usize { + self.recv_buf_size.unwrap_or(DEFAULT_UDP_RECV_BUF) + } + + /// Get the send buffer size, using default if not configured. + pub fn send_buf_size(&self) -> usize { + self.send_buf_size.unwrap_or(DEFAULT_UDP_SEND_BUF) + } } /// Transport instances - either a single config or named instances. diff --git a/src/node/tests/handshake.rs b/src/node/tests/handshake.rs index 6fadd24..c50573a 100644 --- a/src/node/tests/handshake.rs +++ b/src/node/tests/handshake.rs @@ -20,6 +20,7 @@ async fn test_two_node_handshake_udp() { let udp_config = UdpConfig { bind_addr: Some("127.0.0.1:0".to_string()), mtu: Some(1280), + ..Default::default() }; let (packet_tx_a, mut packet_rx_a) = packet_channel(64); @@ -255,6 +256,7 @@ async fn test_run_rx_loop_handshake() { let udp_config = UdpConfig { bind_addr: Some("127.0.0.1:0".to_string()), mtu: Some(1280), + ..Default::default() }; let (packet_tx_a, packet_rx_a) = packet_channel(64); @@ -445,6 +447,7 @@ async fn test_cross_connection_both_initiate() { let udp_config = UdpConfig { bind_addr: Some("127.0.0.1:0".to_string()), mtu: Some(1280), + ..Default::default() }; let (packet_tx_a, mut packet_rx_a) = packet_channel(64); diff --git a/src/node/tests/spanning_tree.rs b/src/node/tests/spanning_tree.rs index 1a8936f..eb5d585 100644 --- a/src/node/tests/spanning_tree.rs +++ b/src/node/tests/spanning_tree.rs @@ -25,6 +25,7 @@ pub(super) async fn make_test_node() -> TestNode { let udp_config = UdpConfig { bind_addr: Some("127.0.0.1:0".to_string()), mtu: Some(1280), + ..Default::default() }; let (packet_tx, packet_rx) = packet_channel(256); diff --git a/src/transport/udp.rs b/src/transport/udp.rs index 2668122..eda0f51 100644 --- a/src/transport/udp.rs +++ b/src/transport/udp.rs @@ -7,6 +7,7 @@ use super::{ TransportId, TransportState, TransportType, }; use crate::config::UdpConfig; +use socket2::{Domain, Protocol, Socket, Type}; use std::net::SocketAddr; use std::sync::Arc; use tokio::net::UdpSocket; @@ -89,11 +90,33 @@ impl UdpTransport { .parse() .map_err(|e| TransportError::StartFailed(format!("invalid bind address: {}", e)))?; - // Bind socket - let socket = UdpSocket::bind(bind_addr) - .await + // Create socket via socket2 for buffer size control + let domain = if bind_addr.is_ipv4() { Domain::IPV4 } else { Domain::IPV6 }; + let sock2 = Socket::new(domain, Type::DGRAM, Some(Protocol::UDP)) + .map_err(|e| TransportError::StartFailed(format!("socket create failed: {}", e)))?; + sock2.set_nonblocking(true) + .map_err(|e| TransportError::StartFailed(format!("set nonblocking failed: {}", e)))?; + sock2.bind(&bind_addr.into()) .map_err(|e| TransportError::StartFailed(format!("bind failed: {}", e)))?; + // Set socket buffer sizes + let recv_buf = self.config.recv_buf_size(); + let send_buf = self.config.send_buf_size(); + sock2.set_recv_buffer_size(recv_buf) + .map_err(|e| TransportError::StartFailed(format!("set recv buffer: {}", e)))?; + sock2.set_send_buffer_size(send_buf) + .map_err(|e| TransportError::StartFailed(format!("set send buffer: {}", e)))?; + + let actual_recv = sock2.recv_buffer_size() + .map_err(|e| TransportError::StartFailed(format!("get recv buffer: {}", e)))?; + let actual_send = sock2.send_buffer_size() + .map_err(|e| TransportError::StartFailed(format!("get send buffer: {}", e)))?; + + // Convert to tokio UdpSocket + let std_socket: std::net::UdpSocket = sock2.into(); + let socket = UdpSocket::from_std(std_socket) + .map_err(|e| TransportError::StartFailed(format!("tokio socket failed: {}", e)))?; + self.local_addr = Some( socket .local_addr() @@ -119,11 +142,15 @@ impl UdpTransport { info!( name = %name, local_addr = %self.local_addr.unwrap(), + recv_buf = actual_recv, + send_buf = actual_send, "UDP transport started" ); } else { info!( local_addr = %self.local_addr.unwrap(), + recv_buf = actual_recv, + send_buf = actual_send, "UDP transport started" ); } @@ -309,6 +336,8 @@ mod tests { UdpConfig { bind_addr: Some(format!("127.0.0.1:{}", port)), mtu: Some(1280), + recv_buf_size: None, + send_buf_size: None, } }