diff --git a/Cargo.toml b/Cargo.toml index a25b225..5a26ed4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -18,7 +18,7 @@ tracing-subscriber = { version = "0.3", features = ["env-filter"] } tun = { version = "0.7", features = ["async"] } libc = "0.2" rtnetlink = "0.14" -tokio = { version = "1", features = ["rt", "macros", "signal", "sync"] } +tokio = { version = "1", features = ["rt", "macros", "signal", "sync", "net", "time"] } futures = "0.3" [dev-dependencies] diff --git a/src/config.rs b/src/config.rs index 71340a9..03e77e4 100644 --- a/src/config.rs +++ b/src/config.rs @@ -72,6 +72,12 @@ const DEFAULT_TUN_NAME: &str = "fips0"; /// Default TUN MTU (IPv6 minimum). const DEFAULT_TUN_MTU: u16 = 1280; +/// Default UDP bind address. +const DEFAULT_UDP_BIND_ADDR: &str = "0.0.0.0:4000"; + +/// Default UDP MTU (IPv6 minimum). +const DEFAULT_UDP_MTU: u16 = 1280; + /// TUN interface configuration (`tun.*`). #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct TunConfig { @@ -101,6 +107,34 @@ impl TunConfig { } } +/// UDP transport configuration (`udp.*`). +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct UdpConfig { + /// Enable UDP transport (`udp.enabled`). + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + pub enabled: bool, + + /// Bind address (`udp.bind_addr`). Defaults to "0.0.0.0:4000". + #[serde(default, skip_serializing_if = "Option::is_none")] + pub bind_addr: Option, + + /// UDP MTU (`udp.mtu`). Defaults to 1280 (IPv6 minimum). + #[serde(default, skip_serializing_if = "Option::is_none")] + pub mtu: Option, +} + +impl UdpConfig { + /// Get the bind address, using default if not configured. + pub fn bind_addr(&self) -> &str { + self.bind_addr.as_deref().unwrap_or(DEFAULT_UDP_BIND_ADDR) + } + + /// Get the UDP MTU, using default if not configured. + pub fn mtu(&self) -> u16 { + self.mtu.unwrap_or(DEFAULT_UDP_MTU) + } +} + /// Root configuration structure. #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct Config { @@ -111,6 +145,10 @@ pub struct Config { /// TUN interface configuration (`tun.*`). #[serde(default)] pub tun: TunConfig, + + /// UDP transport configuration (`udp.*`). + #[serde(default)] + pub udp: UdpConfig, } impl Config { @@ -209,6 +247,16 @@ impl Config { if other.tun.mtu.is_some() { self.tun.mtu = other.tun.mtu; } + // Merge udp section + if other.udp.enabled { + self.udp.enabled = true; + } + if other.udp.bind_addr.is_some() { + self.udp.bind_addr = other.udp.bind_addr; + } + if other.udp.mtu.is_some() { + self.udp.mtu = other.udp.mtu; + } } /// Create an Identity from this configuration. diff --git a/src/lib.rs b/src/lib.rs index fb7739b..1cdcf0d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -22,7 +22,7 @@ pub use identity::{ }; // Re-export config types -pub use config::{Config, ConfigError, IdentityConfig, TunConfig}; +pub use config::{Config, ConfigError, IdentityConfig, TunConfig, UdpConfig}; // Re-export tree types pub use tree::{ParentDeclaration, TreeCoordinate, TreeError, TreeState}; @@ -32,9 +32,11 @@ pub use bloom::{BloomError, BloomFilter, BloomState}; // Re-export transport types pub use transport::{ - DiscoveredPeer, Link, LinkDirection, LinkId, LinkState, LinkStats, Transport, TransportAddr, - TransportError, TransportId, TransportState, TransportType, + packet_channel, DiscoveredPeer, Link, LinkDirection, LinkId, LinkState, LinkStats, PacketRx, + PacketTx, ReceivedPacket, Transport, TransportAddr, TransportError, TransportId, + TransportState, TransportType, }; +pub use transport::udp::UdpTransport; // Re-export protocol types pub use protocol::{ diff --git a/src/transport.rs b/src/transport/mod.rs similarity index 83% rename from src/transport.rs rename to src/transport/mod.rs index ea5c31d..ee533c0 100644 --- a/src/transport.rs +++ b/src/transport/mod.rs @@ -4,11 +4,76 @@ //! underlying communication mechanisms (UDP, Ethernet, Tor, etc.) over //! which FIPS links are established. +pub mod udp; + use secp256k1::XOnlyPublicKey; use std::fmt; -use std::time::Duration; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; use thiserror::Error; +// ============================================================================ +// Packet Channel Types +// ============================================================================ + +/// A packet received from a transport. +#[derive(Clone, Debug)] +pub struct ReceivedPacket { + /// Which transport received this packet. + pub transport_id: TransportId, + /// Remote peer address. + pub remote_addr: TransportAddr, + /// Packet data. + pub data: Vec, + /// Receipt timestamp (Unix milliseconds). + pub timestamp_ms: u64, +} + +impl ReceivedPacket { + /// Create a new received packet with current timestamp. + pub fn new(transport_id: TransportId, remote_addr: TransportAddr, data: Vec) -> Self { + let timestamp_ms = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_millis() as u64) + .unwrap_or(0); + Self { + transport_id, + remote_addr, + data, + timestamp_ms, + } + } + + /// Create a received packet with explicit timestamp. + pub fn with_timestamp( + transport_id: TransportId, + remote_addr: TransportAddr, + data: Vec, + timestamp_ms: u64, + ) -> Self { + Self { + transport_id, + remote_addr, + data, + timestamp_ms, + } + } +} + +/// Channel sender for received packets. +pub type PacketTx = tokio::sync::mpsc::Sender; + +/// Channel receiver for received packets. +pub type PacketRx = tokio::sync::mpsc::Receiver; + +/// Create a packet channel with the given buffer size. +pub fn packet_channel(buffer: usize) -> (PacketTx, PacketRx) { + tokio::sync::mpsc::channel(buffer) +} + +// ============================================================================ +// Transport Identifiers +// ============================================================================ + /// Unique identifier for a transport instance. #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] pub struct TransportId(u32); @@ -53,6 +118,10 @@ impl fmt::Display for LinkId { } } +// ============================================================================ +// Errors +// ============================================================================ + /// Errors related to transport operations. #[derive(Debug, Error)] pub enum TransportError { @@ -96,6 +165,10 @@ pub enum TransportError { Io(#[from] std::io::Error), } +// ============================================================================ +// Transport Type Metadata +// ============================================================================ + /// Static metadata about a transport type. #[derive(Clone, Debug, PartialEq, Eq)] pub struct TransportType { @@ -162,6 +235,10 @@ impl fmt::Display for TransportType { } } +// ============================================================================ +// Transport State +// ============================================================================ + /// Transport lifecycle state. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum TransportState { @@ -210,6 +287,10 @@ impl fmt::Display for TransportState { } } +// ============================================================================ +// Link State +// ============================================================================ + /// Link lifecycle state. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum LinkState { @@ -266,6 +347,10 @@ impl fmt::Display for LinkDirection { } } +// ============================================================================ +// Transport Address +// ============================================================================ + /// Opaque transport-specific address. /// /// Each transport type interprets this differently: @@ -348,6 +433,10 @@ impl From for TransportAddr { } } +// ============================================================================ +// Link Statistics +// ============================================================================ + /// Statistics for a link. #[derive(Clone, Debug, Default)] pub struct LinkStats { @@ -425,6 +514,10 @@ impl LinkStats { } } +// ============================================================================ +// Link +// ============================================================================ + /// A link to a remote endpoint over a transport. #[derive(Clone, Debug)] pub struct Link { @@ -586,6 +679,10 @@ impl Link { } } +// ============================================================================ +// Discovered Peer +// ============================================================================ + /// A peer discovered via transport-layer discovery. #[derive(Clone, Debug)] pub struct DiscoveredPeer { @@ -621,6 +718,10 @@ impl DiscoveredPeer { } } +// ============================================================================ +// Transport Trait +// ============================================================================ + /// Transport trait defining the interface for transport drivers. /// /// This is a simplified synchronous trait. Actual implementations would @@ -651,6 +752,10 @@ pub trait Transport { fn discover(&self) -> Result, TransportError>; } +// ============================================================================ +// Tests +// ============================================================================ + #[cfg(test)] mod tests { use super::*; @@ -885,4 +990,45 @@ mod tests { assert_eq!(format!("{}", TransportState::Up), "up"); assert_eq!(format!("{}", TransportState::Failed), "failed"); } + + #[test] + fn test_received_packet() { + let packet = ReceivedPacket::new( + TransportId::new(1), + TransportAddr::from_string("192.168.1.1:4000"), + vec![1, 2, 3, 4], + ); + + assert_eq!(packet.transport_id, TransportId::new(1)); + assert_eq!(packet.data, vec![1, 2, 3, 4]); + assert!(packet.timestamp_ms > 0); + } + + #[test] + fn test_received_packet_with_timestamp() { + let packet = ReceivedPacket::with_timestamp( + TransportId::new(1), + TransportAddr::from_string("test"), + vec![5, 6], + 12345, + ); + + assert_eq!(packet.timestamp_ms, 12345); + } + + #[tokio::test] + async fn test_packet_channel() { + let (tx, mut rx) = packet_channel(10); + + let packet = ReceivedPacket::new( + TransportId::new(1), + TransportAddr::from_string("test"), + vec![1, 2, 3], + ); + + tx.send(packet.clone()).await.unwrap(); + + let received = rx.recv().await.unwrap(); + assert_eq!(received.data, vec![1, 2, 3]); + } } diff --git a/src/transport/udp.rs b/src/transport/udp.rs new file mode 100644 index 0000000..39483bf --- /dev/null +++ b/src/transport/udp.rs @@ -0,0 +1,494 @@ +//! UDP Transport Implementation +//! +//! Provides UDP-based transport for FIPS peer communication. + +use super::{ + DiscoveredPeer, PacketTx, ReceivedPacket, Transport, TransportAddr, TransportError, + TransportId, TransportState, TransportType, +}; +use crate::config::UdpConfig; +use std::net::SocketAddr; +use std::sync::Arc; +use tokio::net::UdpSocket; +use tokio::task::JoinHandle; +use tracing::{debug, info, warn}; + +/// UDP transport for FIPS. +/// +/// Provides connectionless, unreliable packet delivery over UDP/IP. +/// A single socket serves all peers; links are virtual tuples of +/// (transport_id, remote_addr). +pub struct UdpTransport { + /// Unique transport identifier. + transport_id: TransportId, + /// Configuration. + config: UdpConfig, + /// Current state. + state: TransportState, + /// Bound socket (None until started). + socket: Option>, + /// Channel for delivering received packets to Node. + packet_tx: PacketTx, + /// Receive loop task handle. + recv_task: Option>, + /// Local bound address (after start). + local_addr: Option, +} + +impl UdpTransport { + /// Create a new UDP transport. + pub fn new(transport_id: TransportId, config: UdpConfig, packet_tx: PacketTx) -> Self { + Self { + transport_id, + config, + state: TransportState::Configured, + socket: None, + packet_tx, + recv_task: None, + local_addr: None, + } + } + + /// Get the local bound address (only valid after start). + pub fn local_addr(&self) -> Option { + self.local_addr + } + + /// Get a reference to the socket (only valid after start). + pub fn socket(&self) -> Option<&Arc> { + self.socket.as_ref() + } + + /// Start the transport asynchronously. + /// + /// Binds the UDP socket and spawns the receive loop. + pub async fn start_async(&mut self) -> Result<(), TransportError> { + if !self.state.can_start() { + return Err(TransportError::AlreadyStarted); + } + + self.state = TransportState::Starting; + + // Parse bind address + let bind_addr: SocketAddr = self + .config + .bind_addr() + .parse() + .map_err(|e| TransportError::StartFailed(format!("invalid bind address: {}", e)))?; + + // Bind socket + let socket = UdpSocket::bind(bind_addr) + .await + .map_err(|e| TransportError::StartFailed(format!("bind failed: {}", e)))?; + + self.local_addr = Some( + socket + .local_addr() + .map_err(|e| TransportError::StartFailed(format!("get local addr: {}", e)))?, + ); + + let socket = Arc::new(socket); + self.socket = Some(socket.clone()); + + // Spawn receive loop + let transport_id = self.transport_id; + let packet_tx = self.packet_tx.clone(); + let mtu = self.config.mtu(); + + let recv_task = tokio::spawn(async move { + udp_receive_loop(socket, transport_id, packet_tx, mtu).await; + }); + + self.recv_task = Some(recv_task); + self.state = TransportState::Up; + + info!( + transport_id = %self.transport_id, + local_addr = %self.local_addr.unwrap(), + mtu = self.config.mtu(), + "UDP transport started" + ); + + Ok(()) + } + + /// Stop the transport asynchronously. + pub async fn stop_async(&mut self) -> Result<(), TransportError> { + if !self.state.is_operational() { + return Err(TransportError::NotStarted); + } + + // Abort receive task + if let Some(task) = self.recv_task.take() { + task.abort(); + let _ = task.await; // Ignore JoinError from abort + } + + // Drop socket + self.socket.take(); + self.local_addr = None; + + self.state = TransportState::Down; + + info!( + transport_id = %self.transport_id, + "UDP transport stopped" + ); + + Ok(()) + } + + /// Send a packet asynchronously. + pub async fn send_async( + &self, + addr: &TransportAddr, + data: &[u8], + ) -> Result { + if !self.state.is_operational() { + return Err(TransportError::NotStarted); + } + + if data.len() > self.config.mtu() as usize { + return Err(TransportError::MtuExceeded { + packet_size: data.len(), + mtu: self.config.mtu(), + }); + } + + let socket_addr = parse_socket_addr(addr)?; + let socket = self.socket.as_ref().ok_or(TransportError::NotStarted)?; + + let bytes_sent = socket + .send_to(data, socket_addr) + .await + .map_err(|e| TransportError::SendFailed(format!("{}", e)))?; + + debug!( + transport_id = %self.transport_id, + remote_addr = %socket_addr, + bytes = bytes_sent, + "UDP packet sent" + ); + + Ok(bytes_sent) + } +} + +impl Transport for UdpTransport { + fn transport_id(&self) -> TransportId { + self.transport_id + } + + fn transport_type(&self) -> &TransportType { + &TransportType::UDP + } + + fn state(&self) -> TransportState { + self.state + } + + fn mtu(&self) -> u16 { + self.config.mtu() + } + + fn start(&mut self) -> Result<(), TransportError> { + // Synchronous start not supported - use start_async() + Err(TransportError::NotSupported( + "use start_async() for UDP transport".into(), + )) + } + + fn stop(&mut self) -> Result<(), TransportError> { + // Synchronous stop not supported - use stop_async() + Err(TransportError::NotSupported( + "use stop_async() for UDP transport".into(), + )) + } + + fn send(&self, _addr: &TransportAddr, _data: &[u8]) -> Result<(), TransportError> { + // Synchronous send not supported - use send_async() + Err(TransportError::NotSupported( + "use send_async() for UDP transport".into(), + )) + } + + fn discover(&self) -> Result, TransportError> { + // UDP discovery not yet implemented (would use multicast/DNS-SD) + // Peer configuration is handled at the node level, not transport level + Ok(Vec::new()) + } +} + +/// Parse a TransportAddr as SocketAddr. +fn parse_socket_addr(addr: &TransportAddr) -> Result { + addr.as_str() + .ok_or_else(|| TransportError::InvalidAddress("not valid UTF-8".into()))? + .parse() + .map_err(|e| TransportError::InvalidAddress(format!("{}", e))) +} + +/// UDP receive loop - runs as a spawned task. +async fn udp_receive_loop( + socket: Arc, + transport_id: TransportId, + packet_tx: PacketTx, + mtu: u16, +) { + // Buffer with headroom for slightly oversized packets + let mut buf = vec![0u8; mtu as usize + 100]; + + debug!(transport_id = %transport_id, "UDP receive loop starting"); + + loop { + match socket.recv_from(&mut buf).await { + Ok((len, remote_addr)) => { + let data = buf[..len].to_vec(); + let addr = TransportAddr::from_string(&remote_addr.to_string()); + let packet = ReceivedPacket::new(transport_id, addr, data); + + debug!( + transport_id = %transport_id, + remote_addr = %remote_addr, + bytes = len, + "UDP packet received" + ); + + if packet_tx.send(packet).await.is_err() { + // Receiver dropped, exit loop + info!( + transport_id = %transport_id, + "Packet channel closed, stopping receive loop" + ); + break; + } + } + Err(e) => { + // Log error but continue - transient errors are expected + warn!( + transport_id = %transport_id, + error = %e, + "UDP receive error" + ); + } + } + } + + debug!(transport_id = %transport_id, "UDP receive loop stopped"); +} + +// ============================================================================ +// Tests +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + use crate::transport::packet_channel; + use tokio::time::{timeout, Duration}; + + fn make_config(port: u16) -> UdpConfig { + UdpConfig { + enabled: true, + bind_addr: Some(format!("127.0.0.1:{}", port)), + mtu: Some(1280), + } + } + + #[tokio::test] + async fn test_start_stop() { + let (tx, _rx) = packet_channel(100); + let mut transport = UdpTransport::new(TransportId::new(1), make_config(0), tx); + + assert_eq!(transport.state(), TransportState::Configured); + + transport.start_async().await.unwrap(); + assert_eq!(transport.state(), TransportState::Up); + assert!(transport.local_addr().is_some()); + + transport.stop_async().await.unwrap(); + assert_eq!(transport.state(), TransportState::Down); + } + + #[tokio::test] + async fn test_double_start_fails() { + let (tx, _rx) = packet_channel(100); + let mut transport = UdpTransport::new(TransportId::new(1), make_config(0), tx); + + transport.start_async().await.unwrap(); + + let result = transport.start_async().await; + assert!(matches!(result, Err(TransportError::AlreadyStarted))); + + transport.stop_async().await.unwrap(); + } + + #[tokio::test] + async fn test_stop_not_started_fails() { + let (tx, _rx) = packet_channel(100); + let mut transport = UdpTransport::new(TransportId::new(1), make_config(0), tx); + + let result = transport.stop_async().await; + assert!(matches!(result, Err(TransportError::NotStarted))); + } + + #[tokio::test] + async fn test_send_recv() { + let (tx1, _rx1) = packet_channel(100); + let (tx2, mut rx2) = packet_channel(100); + + let mut t1 = UdpTransport::new(TransportId::new(1), make_config(0), tx1); + let mut t2 = UdpTransport::new(TransportId::new(2), make_config(0), tx2); + + t1.start_async().await.unwrap(); + t2.start_async().await.unwrap(); + + let addr1 = t1.local_addr().unwrap(); + let addr2 = t2.local_addr().unwrap(); + + // Send from t1 to t2 + let data = b"hello world"; + let bytes_sent = t1 + .send_async(&TransportAddr::from_string(&addr2.to_string()), data) + .await + .unwrap(); + assert_eq!(bytes_sent, data.len()); + + // Receive on t2 + let packet = timeout(Duration::from_secs(1), rx2.recv()) + .await + .expect("timeout") + .expect("channel closed"); + + assert_eq!(packet.data, data); + assert_eq!(packet.remote_addr.as_str(), Some(addr1.to_string().as_str())); + + t1.stop_async().await.unwrap(); + t2.stop_async().await.unwrap(); + } + + #[tokio::test] + async fn test_bidirectional() { + let (tx1, mut rx1) = packet_channel(100); + let (tx2, mut rx2) = packet_channel(100); + + let mut t1 = UdpTransport::new(TransportId::new(1), make_config(0), tx1); + let mut t2 = UdpTransport::new(TransportId::new(2), make_config(0), tx2); + + t1.start_async().await.unwrap(); + t2.start_async().await.unwrap(); + + let addr1 = TransportAddr::from_string(&t1.local_addr().unwrap().to_string()); + let addr2 = TransportAddr::from_string(&t2.local_addr().unwrap().to_string()); + + // Send from t1 to t2 + t1.send_async(&addr2, b"ping").await.unwrap(); + + // Receive on t2 + let packet = timeout(Duration::from_secs(1), rx2.recv()) + .await + .expect("timeout") + .expect("channel closed"); + assert_eq!(packet.data, b"ping"); + + // Send from t2 to t1 + t2.send_async(&addr1, b"pong").await.unwrap(); + + // Receive on t1 + let packet = timeout(Duration::from_secs(1), rx1.recv()) + .await + .expect("timeout") + .expect("channel closed"); + assert_eq!(packet.data, b"pong"); + + t1.stop_async().await.unwrap(); + t2.stop_async().await.unwrap(); + } + + #[tokio::test] + async fn test_mtu_exceeded() { + let (tx, _rx) = packet_channel(100); + let mut transport = UdpTransport::new( + TransportId::new(1), + UdpConfig { + mtu: Some(100), + ..make_config(0) + }, + tx, + ); + + transport.start_async().await.unwrap(); + + let oversized = vec![0u8; 200]; + let result = transport + .send_async(&TransportAddr::from_string("127.0.0.1:9999"), &oversized) + .await; + + assert!(matches!(result, Err(TransportError::MtuExceeded { .. }))); + + transport.stop_async().await.unwrap(); + } + + #[tokio::test] + async fn test_send_not_started() { + let (tx, _rx) = packet_channel(100); + let transport = UdpTransport::new(TransportId::new(1), make_config(0), tx); + + let result = transport + .send_async(&TransportAddr::from_string("127.0.0.1:9999"), b"test") + .await; + + assert!(matches!(result, Err(TransportError::NotStarted))); + } + + #[tokio::test] + async fn test_discover_returns_empty() { + let (tx, _rx) = packet_channel(100); + let transport = UdpTransport::new(TransportId::new(1), make_config(0), tx); + + // Discovery returns empty until multicast/DNS-SD is implemented + let peers = transport.discover().unwrap(); + assert!(peers.is_empty()); + } + + #[test] + fn test_transport_type() { + let (tx, _rx) = packet_channel(100); + let transport = UdpTransport::new(TransportId::new(1), make_config(0), tx); + + assert_eq!(transport.transport_type().name, "udp"); + assert!(!transport.transport_type().connection_oriented); + assert!(!transport.transport_type().reliable); + } + + #[test] + fn test_sync_methods_return_not_supported() { + let (tx, _rx) = packet_channel(100); + let mut transport = UdpTransport::new(TransportId::new(1), make_config(0), tx); + + assert!(matches!( + transport.start(), + Err(TransportError::NotSupported(_)) + )); + assert!(matches!( + transport.stop(), + Err(TransportError::NotSupported(_)) + )); + assert!(matches!( + transport.send(&TransportAddr::from_string("test"), b"data"), + Err(TransportError::NotSupported(_)) + )); + } + + #[test] + fn test_parse_socket_addr() { + let addr = TransportAddr::from_string("192.168.1.1:4000"); + let result = parse_socket_addr(&addr).unwrap(); + assert_eq!(result.to_string(), "192.168.1.1:4000"); + + let invalid = TransportAddr::from_string("not_an_address"); + assert!(parse_socket_addr(&invalid).is_err()); + + let binary = TransportAddr::new(vec![0xff, 0x80]); + assert!(parse_socket_addr(&binary).is_err()); + } +}