Add UDP transport implementation

- Convert transport.rs to module directory (transport/mod.rs + transport/udp.rs)
- Add packet channel types for transport→Node communication:
  - ReceivedPacket struct with transport_id, remote_addr, data, timestamp
  - PacketTx/PacketRx type aliases for tokio mpsc channels
  - packet_channel() constructor function
- Add UdpConfig to config.rs (enabled, bind_addr, mtu)
- Implement UdpTransport with async lifecycle:
  - start_async(): binds socket, spawns receive loop
  - stop_async(): aborts receive task, closes socket
  - send_async(): sends packet with MTU validation
  - discover(): returns empty (peer config is node-level)
- Update lib.rs with new re-exports
- Add tokio net and time features to Cargo.toml

All 185 tests pass.
This commit is contained in:
Johnathan Corgan
2026-01-30 21:02:33 +00:00
parent c421dad525
commit d30865f60b
5 changed files with 695 additions and 5 deletions
+1 -1
View File
@@ -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]
+48
View File
@@ -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<String>,
/// UDP MTU (`udp.mtu`). Defaults to 1280 (IPv6 minimum).
#[serde(default, skip_serializing_if = "Option::is_none")]
pub mtu: Option<u16>,
}
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.
+5 -3
View File
@@ -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::{
+147 -1
View File
@@ -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<u8>,
/// 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<u8>) -> 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<u8>,
timestamp_ms: u64,
) -> Self {
Self {
transport_id,
remote_addr,
data,
timestamp_ms,
}
}
}
/// Channel sender for received packets.
pub type PacketTx = tokio::sync::mpsc::Sender<ReceivedPacket>;
/// Channel receiver for received packets.
pub type PacketRx = tokio::sync::mpsc::Receiver<ReceivedPacket>;
/// 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<String> 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<Vec<DiscoveredPeer>, 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]);
}
}
+494
View File
@@ -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<Arc<UdpSocket>>,
/// Channel for delivering received packets to Node.
packet_tx: PacketTx,
/// Receive loop task handle.
recv_task: Option<JoinHandle<()>>,
/// Local bound address (after start).
local_addr: Option<SocketAddr>,
}
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<SocketAddr> {
self.local_addr
}
/// Get a reference to the socket (only valid after start).
pub fn socket(&self) -> Option<&Arc<UdpSocket>> {
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<usize, TransportError> {
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<Vec<DiscoveredPeer>, 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<SocketAddr, TransportError> {
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<UdpSocket>,
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());
}
}