Add rustfmt formatting policy and reformat codebase

Add rustfmt.toml with stable defaults and apply cargo fmt to all
source files. This establishes a consistent formatting baseline
for CI enforcement.
This commit is contained in:
Johnathan Corgan
2026-04-10 08:27:07 +00:00
parent a859da7748
commit 13c0b70dc3
101 changed files with 3451 additions and 2227 deletions

1
rustfmt.toml Normal file
View File

@@ -0,0 +1 @@
# Use stable defaults

View File

@@ -3,12 +3,12 @@
//! Loads configuration and creates the top-level node instance. //! Loads configuration and creates the top-level node instance.
use clap::Parser; use clap::Parser;
use fips::config::{resolve_identity, IdentitySource}; use fips::config::{IdentitySource, resolve_identity};
use fips::version; use fips::version;
use fips::{Config, Node}; use fips::{Config, Node};
use std::path::PathBuf; use std::path::PathBuf;
use tracing::{error, info, warn, Level}; use tracing::{Level, error, info, warn};
use tracing_subscriber::{fmt, EnvFilter}; use tracing_subscriber::{EnvFilter, fmt};
/// FIPS mesh network daemon /// FIPS mesh network daemon
#[derive(Parser, Debug)] #[derive(Parser, Debug)]
@@ -31,10 +31,7 @@ async fn main() {
.with_default_directive(Level::INFO.into()) .with_default_directive(Level::INFO.into())
.from_env_lossy(); .from_env_lossy();
fmt() fmt().with_env_filter(filter).with_target(true).init();
.with_env_filter(filter)
.with_target(true)
.init();
let args = Args::parse(); let args = Args::parse();
@@ -47,7 +44,11 @@ async fn main() {
match Config::load_file(config_path) { match Config::load_file(config_path) {
Ok(config) => (config, vec![config_path.clone()]), Ok(config) => (config, vec![config_path.clone()]),
Err(e) => { Err(e) => {
error!("Failed to load configuration from {}: {}", config_path.display(), e); error!(
"Failed to load configuration from {}: {}",
config_path.display(),
e
);
std::process::exit(1); std::process::exit(1);
} }
} }
@@ -80,8 +81,12 @@ async fn main() {
}; };
match &resolved.source { match &resolved.source {
IdentitySource::Config => info!("Using identity from configuration"), IdentitySource::Config => info!("Using identity from configuration"),
IdentitySource::KeyFile(p) => info!(path = %p.display(), "Loaded persistent identity from key file"), IdentitySource::KeyFile(p) => {
IdentitySource::Generated(p) => info!(path = %p.display(), "Generated persistent identity, saved to key file"), info!(path = %p.display(), "Loaded persistent identity from key file")
}
IdentitySource::Generated(p) => {
info!(path = %p.display(), "Generated persistent identity, saved to key file")
}
IdentitySource::Ephemeral => info!("Using ephemeral identity (new keypair each start)"), IdentitySource::Ephemeral => info!("Using ephemeral identity (new keypair each start)"),
} }

View File

@@ -7,7 +7,7 @@ use clap::{Parser, Subcommand};
use fips::config::{write_key_file, write_pub_file}; use fips::config::{write_key_file, write_pub_file};
use fips::upper::hosts::HostMap; use fips::upper::hosts::HostMap;
use fips::version; use fips::version;
use fips::{encode_nsec, Identity}; use fips::{Identity, encode_nsec};
use std::io::{BufRead, BufReader, Write}; use std::io::{BufRead, BufReader, Write};
use std::os::unix::net::UnixStream; use std::os::unix::net::UnixStream;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
@@ -223,12 +223,7 @@ fn main() {
let cli = Cli::parse(); let cli = Cli::parse();
// Commands that don't require a running daemon // Commands that don't require a running daemon
if let Commands::Keygen { if let Commands::Keygen { dir, force, stdout } = &cli.command {
dir,
force,
stdout,
} = &cli.command
{
let identity = Identity::generate(); let identity = Identity::generate();
let nsec = encode_nsec(&identity.keypair().secret_key()); let nsec = encode_nsec(&identity.keypair().secret_key());
let npub = identity.npub(); let npub = identity.npub();

View File

@@ -94,10 +94,7 @@ impl Tab {
/// Whether this tab has a table view with row selection. /// Whether this tab has a table view with row selection.
pub fn has_table(&self) -> bool { pub fn has_table(&self) -> bool {
matches!( matches!(self, Tab::Peers | Tab::Sessions | Tab::Transports)
self,
Tab::Peers | Tab::Sessions | Tab::Transports
)
} }
} }
@@ -172,7 +169,10 @@ impl App {
return; return;
} }
let state = self.table_states.entry(self.active_tab).or_default(); let state = self.table_states.entry(self.active_tab).or_default();
let i = state.selected().map(|s| (s + 1).min(count - 1)).unwrap_or(0); let i = state
.selected()
.map(|s| (s + 1).min(count - 1))
.unwrap_or(0);
state.select(Some(i)); state.select(Some(i));
} }

View File

@@ -17,27 +17,29 @@ impl EventHandler {
pub fn new(tick_rate: Duration) -> Self { pub fn new(tick_rate: Duration) -> Self {
let (tx, rx) = mpsc::channel(); let (tx, rx) = mpsc::channel();
thread::spawn(move || loop { thread::spawn(move || {
if event::poll(tick_rate).unwrap_or(false) { loop {
if let Ok(evt) = event::read() { if event::poll(tick_rate).unwrap_or(false) {
match evt { if let Ok(evt) = event::read() {
CrosstermEvent::Key(key) => { match evt {
if tx.send(Event::Key(key)).is_err() { CrosstermEvent::Key(key) => {
return; if tx.send(Event::Key(key)).is_err() {
return;
}
} }
} CrosstermEvent::Resize(..) => {
CrosstermEvent::Resize(..) => { if tx.send(Event::Resize).is_err() {
if tx.send(Event::Resize).is_err() { return;
return; }
} }
_ => {}
} }
_ => {}
} }
} } else {
} else { // Poll timed out — send a tick
// Poll timed out — send a tick if tx.send(Event::Tick).is_err() {
if tx.send(Event::Tick).is_err() { return;
return; }
} }
} }
}); });

View File

@@ -12,8 +12,8 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let data = match app.data.get(&Tab::Bloom) { let data = match app.data.get(&Tab::Bloom) {
Some(d) => d, Some(d) => d,
None => { None => {
let msg = Paragraph::new(" Waiting for data...") let msg =
.style(Style::default().fg(Color::DarkGray)); Paragraph::new(" Waiting for data...").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, area); frame.render_widget(msg, area);
return; return;
} }
@@ -22,7 +22,7 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let chunks = Layout::vertical([ let chunks = Layout::vertical([
Constraint::Length(7), // Bloom Filter State Constraint::Length(7), // Bloom Filter State
Constraint::Length(15), // Bloom Announce Stats Constraint::Length(15), // Bloom Announce Stats
Constraint::Min(3), // Peer Filters Constraint::Min(3), // Peer Filters
]) ])
.split(area); .split(area);
@@ -64,16 +64,28 @@ fn draw_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
helpers::section_header("Inbound"), helpers::section_header("Inbound"),
helpers::kv_line("Received", &helpers::nested_u64(data, "stats", "received")), helpers::kv_line("Received", &helpers::nested_u64(data, "stats", "received")),
helpers::kv_line("Accepted", &helpers::nested_u64(data, "stats", "accepted")), helpers::kv_line("Accepted", &helpers::nested_u64(data, "stats", "accepted")),
helpers::kv_line("Decode Error", &helpers::nested_u64(data, "stats", "decode_error")), helpers::kv_line(
"Decode Error",
&helpers::nested_u64(data, "stats", "decode_error"),
),
helpers::kv_line("Invalid", &helpers::nested_u64(data, "stats", "invalid")), helpers::kv_line("Invalid", &helpers::nested_u64(data, "stats", "invalid")),
helpers::kv_line("Non-V1", &helpers::nested_u64(data, "stats", "non_v1")), helpers::kv_line("Non-V1", &helpers::nested_u64(data, "stats", "non_v1")),
helpers::kv_line("Unknown Peer", &helpers::nested_u64(data, "stats", "unknown_peer")), helpers::kv_line(
"Unknown Peer",
&helpers::nested_u64(data, "stats", "unknown_peer"),
),
helpers::kv_line("Stale", &helpers::nested_u64(data, "stats", "stale")), helpers::kv_line("Stale", &helpers::nested_u64(data, "stats", "stale")),
Line::from(""), Line::from(""),
helpers::section_header("Outbound"), helpers::section_header("Outbound"),
helpers::kv_line("Sent", &helpers::nested_u64(data, "stats", "sent")), helpers::kv_line("Sent", &helpers::nested_u64(data, "stats", "sent")),
helpers::kv_line("Debounce Suppressed", &helpers::nested_u64(data, "stats", "debounce_suppressed")), helpers::kv_line(
helpers::kv_line("Send Failed", &helpers::nested_u64(data, "stats", "send_failed")), "Debounce Suppressed",
&helpers::nested_u64(data, "stats", "debounce_suppressed"),
),
helpers::kv_line(
"Send Failed",
&helpers::nested_u64(data, "stats", "send_failed"),
),
]; ];
let max_lines = inner.height as usize; let max_lines = inner.height as usize;
@@ -97,8 +109,7 @@ fn draw_peer_filters(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
frame.render_widget(block, area); frame.render_widget(block, area);
if filters.is_empty() { if filters.is_empty() {
let msg = let msg = Paragraph::new(" No peers").style(Style::default().fg(Color::DarkGray));
Paragraph::new(" No peers").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, inner); frame.render_widget(msg, inner);
return; return;
} }

View File

@@ -12,18 +12,18 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let data = match app.data.get(&crate::app::Tab::Node) { let data = match app.data.get(&crate::app::Tab::Node) {
Some(d) => d, Some(d) => d,
None => { None => {
let msg = Paragraph::new(" Waiting for data...") let msg =
.style(Style::default().fg(Color::DarkGray)); Paragraph::new(" Waiting for data...").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, area); frame.render_widget(msg, area);
return; return;
} }
}; };
let chunks = Layout::vertical([ let chunks = Layout::vertical([
Constraint::Length(7), // Runtime Constraint::Length(7), // Runtime
Constraint::Length(7), // Identity Constraint::Length(7), // Identity
Constraint::Length(5), // State Constraint::Length(5), // State
Constraint::Length(9), // Traffic Constraint::Length(9), // Traffic
Constraint::Min(0), // remaining Constraint::Min(0), // remaining
]) ])
.split(area); .split(area);
@@ -35,16 +35,17 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
} }
fn draw_runtime(frame: &mut Frame, data: &serde_json::Value, area: Rect) { fn draw_runtime(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
let block = Block::default() let block = Block::default().borders(Borders::ALL).title(" Runtime ");
.borders(Borders::ALL)
.title(" Runtime ");
let inner = block.inner(area); let inner = block.inner(area);
frame.render_widget(block, area); frame.render_widget(block, area);
let version = helpers::str_field(data, "version"); let version = helpers::str_field(data, "version");
let pid = helpers::u64_field(data, "pid"); let pid = helpers::u64_field(data, "pid");
let exe = helpers::str_field(data, "exe_path"); let exe = helpers::str_field(data, "exe_path");
let uptime_secs = data.get("uptime_secs").and_then(|v| v.as_u64()).unwrap_or(0); let uptime_secs = data
.get("uptime_secs")
.and_then(|v| v.as_u64())
.unwrap_or(0);
let uptime = format_uptime(uptime_secs); let uptime = format_uptime(uptime_secs);
let socket = helpers::str_field(data, "control_socket"); let socket = helpers::str_field(data, "control_socket");
let tun_name = helpers::str_field(data, "tun_name"); let tun_name = helpers::str_field(data, "tun_name");
@@ -78,9 +79,7 @@ fn draw_runtime(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
} }
fn draw_identity(frame: &mut Frame, data: &serde_json::Value, area: Rect) { fn draw_identity(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
let block = Block::default() let block = Block::default().borders(Borders::ALL).title(" Identity ");
.borders(Borders::ALL)
.title(" Identity ");
let inner = block.inner(area); let inner = block.inner(area);
frame.render_widget(block, area); frame.render_widget(block, area);
@@ -109,9 +108,7 @@ fn draw_identity(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
} }
fn draw_state(frame: &mut Frame, data: &serde_json::Value, area: Rect) { fn draw_state(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
let block = Block::default() let block = Block::default().borders(Borders::ALL).title(" State ");
.borders(Borders::ALL)
.title(" State ");
let inner = block.inner(area); let inner = block.inner(area);
frame.render_widget(block, area); frame.render_widget(block, area);
@@ -167,9 +164,7 @@ fn draw_state(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
} }
fn draw_node_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) { fn draw_node_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
let block = Block::default() let block = Block::default().borders(Borders::ALL).title(" Traffic ");
.borders(Borders::ALL)
.title(" Traffic ");
let inner = block.inner(area); let inner = block.inner(area);
frame.render_widget(block, area); frame.render_widget(block, area);
@@ -197,7 +192,10 @@ fn fwd_line(data: &serde_json::Value, label: &str, pkt_key: &str, byte_key: &str
.and_then(|f| f.get(byte_key)) .and_then(|f| f.get(byte_key))
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(0); .unwrap_or(0);
helpers::kv_line(label, &format!("{} pkts ({})", pkts, helpers::format_bytes(bytes))) helpers::kv_line(
label,
&format!("{} pkts ({})", pkts, helpers::format_bytes(bytes)),
)
} }
/// Format seconds as human-readable uptime (e.g., "3d 2h 15m 4s"). /// Format seconds as human-readable uptime (e.g., "3d 2h 15m 4s").

View File

@@ -119,7 +119,13 @@ pub fn nested_f64(data: &Value, outer: &str, inner: &str, decimals: usize) -> St
} }
/// Get a nested f64 field, preferring `preferred` key with fallback to `fallback` key. /// Get a nested f64 field, preferring `preferred` key with fallback to `fallback` key.
pub fn nested_f64_prefer(data: &Value, outer: &str, preferred: &str, fallback: &str, decimals: usize) -> String { pub fn nested_f64_prefer(
data: &Value,
outer: &str,
preferred: &str,
fallback: &str,
decimals: usize,
) -> String {
data.get(outer) data.get(outer)
.and_then(|o| o.get(preferred).or_else(|| o.get(fallback))) .and_then(|o| o.get(preferred).or_else(|| o.get(fallback)))
.and_then(|v| v.as_f64()) .and_then(|v| v.as_f64())
@@ -161,11 +167,7 @@ pub fn section_header(title: &str) -> Line<'static> {
/// Key-value line for detail views. /// Key-value line for detail views.
pub fn kv_line(key: &str, value: &str) -> Line<'static> { pub fn kv_line(key: &str, value: &str) -> Line<'static> {
Line::from(vec![ Line::from(vec![
Span::styled( Span::styled(format!(" {key}: "), Style::default().fg(Color::DarkGray)),
format!(" {key}: "),
Style::default().fg(Color::DarkGray),
),
Span::raw(value.to_string()), Span::raw(value.to_string()),
]) ])
} }

View File

@@ -12,18 +12,15 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let data = match app.data.get(&Tab::Mmp) { let data = match app.data.get(&Tab::Mmp) {
Some(d) => d, Some(d) => d,
None => { None => {
let msg = Paragraph::new(" Waiting for data...") let msg =
.style(Style::default().fg(Color::DarkGray)); Paragraph::new(" Waiting for data...").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, area); frame.render_widget(msg, area);
return; return;
} }
}; };
let chunks = Layout::vertical([ let chunks =
Constraint::Percentage(60), Layout::vertical([Constraint::Percentage(60), Constraint::Percentage(40)]).split(area);
Constraint::Percentage(40),
])
.split(area);
draw_link_mmp(frame, data, chunks[0]); draw_link_mmp(frame, data, chunks[0]);
draw_session_mmp(frame, data, chunks[1]); draw_session_mmp(frame, data, chunks[1]);
@@ -44,8 +41,7 @@ fn draw_link_mmp(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
frame.render_widget(block, area); frame.render_widget(block, area);
if peers.is_empty() { if peers.is_empty() {
let msg = let msg = Paragraph::new(" No peers").style(Style::default().fg(Color::DarkGray));
Paragraph::new(" No peers").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, inner); frame.render_widget(msg, inner);
return; return;
} }
@@ -152,8 +148,7 @@ fn draw_session_mmp(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
frame.render_widget(block, area); frame.render_widget(block, area);
if sessions.is_empty() { if sessions.is_empty() {
let msg = Paragraph::new(" No sessions") let msg = Paragraph::new(" No sessions").style(Style::default().fg(Color::DarkGray));
.style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, inner); frame.render_widget(msg, inner);
return; return;
} }
@@ -234,4 +229,3 @@ fn trend_color(trend: &str, bad_rising: bool) -> Color {
_ => Color::DarkGray, // "stable" _ => Color::DarkGray, // "stable"
} }
} }

View File

@@ -19,7 +19,7 @@ use crate::app::{App, ConnectionState, Tab};
pub fn draw(frame: &mut Frame, app: &mut App) { pub fn draw(frame: &mut Frame, app: &mut App) {
let chunks = Layout::vertical([ let chunks = Layout::vertical([
Constraint::Length(3), // tab bar Constraint::Length(3), // tab bar
Constraint::Min(1), // content Constraint::Min(1), // content
Constraint::Length(1), // status bar Constraint::Length(1), // status bar
]) ])
.split(frame.area()); .split(frame.area());
@@ -51,7 +51,11 @@ fn draw_tab_bar(frame: &mut Frame, app: &App, area: Rect) {
spans.push(Span::styled(" | ", divider)); spans.push(Span::styled(" | ", divider));
} }
} }
let style = if *tab == app.active_tab { highlight } else { normal }; let style = if *tab == app.active_tab {
highlight
} else {
normal
};
spans.push(Span::styled(tab.label(), style)); spans.push(Span::styled(tab.label(), style));
} }

View File

@@ -2,7 +2,9 @@ use ratatui::Frame;
use ratatui::layout::{Constraint, Layout, Rect}; use ratatui::layout::{Constraint, Layout, Rect};
use ratatui::style::{Color, Modifier, Style}; use ratatui::style::{Color, Modifier, Style};
use ratatui::text::Line; use ratatui::text::Line;
use ratatui::widgets::{Block, Borders, Cell, Paragraph, Row, Scrollbar, ScrollbarOrientation, ScrollbarState, Table}; use ratatui::widgets::{
Block, Borders, Cell, Paragraph, Row, Scrollbar, ScrollbarOrientation, ScrollbarState, Table,
};
use crate::app::{App, Tab}; use crate::app::{App, Tab};
@@ -14,11 +16,8 @@ pub fn draw(frame: &mut Frame, app: &mut App, area: Rect) {
if app.detail_view.is_some() { if app.detail_view.is_some() {
// Split: left 40% table, right 60% detail // Split: left 40% table, right 60% detail
let chunks = Layout::horizontal([ let chunks = Layout::horizontal([Constraint::Percentage(40), Constraint::Percentage(60)])
Constraint::Percentage(40), .split(area);
Constraint::Percentage(60),
])
.split(area);
draw_table(frame, app, chunks[0], &peers, row_count); draw_table(frame, app, chunks[0], &peers, row_count);
draw_detail(frame, app, chunks[1], &peers); draw_detail(frame, app, chunks[1], &peers);
@@ -29,7 +28,8 @@ pub fn draw(frame: &mut Frame, app: &mut App, area: Rect) {
/// Get peers sorted by LQI ascending (best first). Peers without LQI sort last. /// Get peers sorted by LQI ascending (best first). Peers without LQI sort last.
fn get_peers_sorted(app: &App) -> Vec<serde_json::Value> { fn get_peers_sorted(app: &App) -> Vec<serde_json::Value> {
let mut peers = app.data let mut peers = app
.data
.get(&Tab::Peers) .get(&Tab::Peers)
.and_then(|v| v.get("peers")) .and_then(|v| v.get("peers"))
.and_then(|v| v.as_array()) .and_then(|v| v.as_array())
@@ -37,8 +37,14 @@ fn get_peers_sorted(app: &App) -> Vec<serde_json::Value> {
.unwrap_or_default(); .unwrap_or_default();
peers.sort_by(|a, b| { peers.sort_by(|a, b| {
let lqi_a = a.get("mmp").and_then(|m| m.get("lqi")).and_then(|v| v.as_f64()); let lqi_a = a
let lqi_b = b.get("mmp").and_then(|m| m.get("lqi")).and_then(|v| v.as_f64()); .get("mmp")
.and_then(|m| m.get("lqi"))
.and_then(|v| v.as_f64());
let lqi_b = b
.get("mmp")
.and_then(|m| m.get("lqi"))
.and_then(|v| v.as_f64());
match (lqi_a, lqi_b) { match (lqi_a, lqi_b) {
(Some(a), Some(b)) => a.partial_cmp(&b).unwrap_or(std::cmp::Ordering::Equal), (Some(a), Some(b)) => a.partial_cmp(&b).unwrap_or(std::cmp::Ordering::Equal),
(Some(_), None) => std::cmp::Ordering::Less, (Some(_), None) => std::cmp::Ordering::Less,
@@ -50,7 +56,13 @@ fn get_peers_sorted(app: &App) -> Vec<serde_json::Value> {
peers peers
} }
fn draw_table(frame: &mut Frame, app: &mut App, area: Rect, peers: &[serde_json::Value], row_count: usize) { fn draw_table(
frame: &mut Frame,
app: &mut App,
area: Rect,
peers: &[serde_json::Value],
row_count: usize,
) {
let header = Row::new(vec![ let header = Row::new(vec![
Cell::from("Name"), Cell::from("Name"),
Cell::from("Npub"), Cell::from("Npub"),
@@ -74,13 +86,25 @@ fn draw_table(frame: &mut Frame, app: &mut App, area: Rect, peers: &[serde_json:
.map(|peer| { .map(|peer| {
let name = helpers::str_field(peer, "display_name"); let name = helpers::str_field(peer, "display_name");
let npub = helpers::str_field(peer, "npub"); let npub = helpers::str_field(peer, "npub");
let is_parent = peer.get("is_parent").and_then(|v| v.as_bool()).unwrap_or(false); let is_parent = peer
let is_child = peer.get("is_child").and_then(|v| v.as_bool()).unwrap_or(false); .get("is_parent")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let is_child = peer
.get("is_child")
.and_then(|v| v.as_bool())
.unwrap_or(false);
// Transport: "type addr" (e.g., "udp 1.2.3.4:2121") // Transport: "type addr" (e.g., "udp 1.2.3.4:2121")
let transport = { let transport = {
let t_type = peer.get("transport_type").and_then(|v| v.as_str()).unwrap_or(""); let t_type = peer
let t_addr = peer.get("transport_addr").and_then(|v| v.as_str()).unwrap_or(""); .get("transport_type")
.and_then(|v| v.as_str())
.unwrap_or("");
let t_addr = peer
.get("transport_addr")
.and_then(|v| v.as_str())
.unwrap_or("");
if t_type.is_empty() && t_addr.is_empty() { if t_type.is_empty() && t_addr.is_empty() {
"-".to_string() "-".to_string()
} else if t_type.is_empty() { } else if t_type.is_empty() {
@@ -92,11 +116,15 @@ fn draw_table(frame: &mut Frame, app: &mut App, area: Rect, peers: &[serde_json:
} }
}; };
let dir = peer.get("direction").and_then(|v| v.as_str()).map(|d| match d { let dir = peer
"inbound" => "in", .get("direction")
"outbound" => "out", .and_then(|v| v.as_str())
other => other, .map(|d| match d {
}).unwrap_or("-"); "inbound" => "in",
"outbound" => "out",
other => other,
})
.unwrap_or("-");
let srtt = helpers::nested_f64(peer, "mmp", "srtt_ms", 1); let srtt = helpers::nested_f64(peer, "mmp", "srtt_ms", 1);
let loss = helpers::nested_f64_prefer(peer, "mmp", "smoothed_loss", "loss_rate", 3); let loss = helpers::nested_f64_prefer(peer, "mmp", "smoothed_loss", "loss_rate", 3);
let lqi = helpers::nested_f64(peer, "mmp", "lqi", 2); let lqi = helpers::nested_f64(peer, "mmp", "lqi", 2);
@@ -130,16 +158,16 @@ fn draw_table(frame: &mut Frame, app: &mut App, area: Rect, peers: &[serde_json:
.collect(); .collect();
let widths = [ let widths = [
Constraint::Min(12), // Name Constraint::Min(12), // Name
Constraint::Length(67), // Npub (full bech32) Constraint::Length(67), // Npub (full bech32)
Constraint::Min(20), // Transport Constraint::Min(20), // Transport
Constraint::Length(4), // Dir Constraint::Length(4), // Dir
Constraint::Length(8), // SRTT Constraint::Length(8), // SRTT
Constraint::Length(7), // Loss Constraint::Length(7), // Loss
Constraint::Length(6), // LQI Constraint::Length(6), // LQI
Constraint::Length(10), // Goodput Constraint::Length(10), // Goodput
Constraint::Length(9), // Pkts Tx Constraint::Length(9), // Pkts Tx
Constraint::Length(9), // Pkts Rx Constraint::Length(9), // Pkts Rx
]; ];
let table = Table::new(rows, widths) let table = Table::new(rows, widths)
@@ -189,10 +217,22 @@ fn draw_detail(frame: &mut Frame, app: &App, area: Rect, peers: &[serde_json::Va
return; return;
}; };
let has_tree = peer.get("has_tree_position").and_then(|v| v.as_bool()).unwrap_or(false); let has_tree = peer
let has_bloom = peer.get("has_bloom_filter").and_then(|v| v.as_bool()).unwrap_or(false); .get("has_tree_position")
let is_parent = peer.get("is_parent").and_then(|v| v.as_bool()).unwrap_or(false); .and_then(|v| v.as_bool())
let is_child = peer.get("is_child").and_then(|v| v.as_bool()).unwrap_or(false); .unwrap_or(false);
let has_bloom = peer
.get("has_bloom_filter")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let is_parent = peer
.get("is_parent")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let is_child = peer
.get("is_child")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let tree_role = if is_parent { let tree_role = if is_parent {
"parent" "parent"
@@ -215,7 +255,12 @@ fn draw_detail(frame: &mut Frame, app: &App, area: Rect, peers: &[serde_json::Va
helpers::section_header("Connection"), helpers::section_header("Connection"),
helpers::kv_line("Connectivity", helpers::str_field(peer, "connectivity")), helpers::kv_line("Connectivity", helpers::str_field(peer, "connectivity")),
helpers::kv_line("Link ID", &helpers::u64_field(peer, "link_id")), helpers::kv_line("Link ID", &helpers::u64_field(peer, "link_id")),
helpers::kv_line("Direction", peer.get("direction").and_then(|v| v.as_str()).unwrap_or("-")), helpers::kv_line(
"Direction",
peer.get("direction")
.and_then(|v| v.as_str())
.unwrap_or("-"),
),
]; ];
if let Some(addr) = peer.get("transport_addr").and_then(|v| v.as_str()) { if let Some(addr) = peer.get("transport_addr").and_then(|v| v.as_str()) {
lines.push(helpers::kv_line("Transport Addr", addr)); lines.push(helpers::kv_line("Transport Addr", addr));
@@ -227,30 +272,52 @@ fn draw_detail(frame: &mut Frame, app: &App, area: Rect, peers: &[serde_json::Va
let link_id = peer.get("link_id").and_then(|v| v.as_u64()); let link_id = peer.get("link_id").and_then(|v| v.as_u64());
let link = lookup_link(app, link_id); let link = lookup_link(app, link_id);
if let Some(ref link) = link { if let Some(ref link) = link {
lines.push(helpers::kv_line("Link State", helpers::str_field(link, "state"))); lines.push(helpers::kv_line(
"Link State",
helpers::str_field(link, "state"),
));
} }
lines.extend([ lines.extend([
helpers::kv_line("Authenticated", &helpers::format_elapsed_ms( helpers::kv_line(
peer.get("authenticated_at_ms").and_then(|v| v.as_u64()).unwrap_or(0), "Authenticated",
)), &helpers::format_elapsed_ms(
helpers::kv_line("Last Seen", &helpers::format_elapsed_ms( peer.get("authenticated_at_ms")
peer.get("last_seen_ms").and_then(|v| v.as_u64()).unwrap_or(0), .and_then(|v| v.as_u64())
)), .unwrap_or(0),
),
),
helpers::kv_line(
"Last Seen",
&helpers::format_elapsed_ms(
peer.get("last_seen_ms")
.and_then(|v| v.as_u64())
.unwrap_or(0),
),
),
Line::from(""), Line::from(""),
]); ]);
// Transport info (cross-referenced from link -> transport) // Transport info (cross-referenced from link -> transport)
if let Some(transport) = link.as_ref().and_then(|l| lookup_transport(app, l)) { if let Some(transport) = link.as_ref().and_then(|l| lookup_transport(app, l)) {
lines.push(helpers::section_header("Transport")); lines.push(helpers::section_header("Transport"));
lines.push(helpers::kv_line("Type", helpers::str_field(&transport, "type"))); lines.push(helpers::kv_line(
"Type",
helpers::str_field(&transport, "type"),
));
if let Some(name) = transport.get("name").and_then(|v| v.as_str()) { if let Some(name) = transport.get("name").and_then(|v| v.as_str()) {
lines.push(helpers::kv_line("Name", name)); lines.push(helpers::kv_line("Name", name));
} }
lines.push(helpers::kv_line("MTU", &helpers::u64_field(&transport, "mtu"))); lines.push(helpers::kv_line(
"MTU",
&helpers::u64_field(&transport, "mtu"),
));
if let Some(addr) = transport.get("local_addr").and_then(|v| v.as_str()) { if let Some(addr) = transport.get("local_addr").and_then(|v| v.as_str()) {
lines.push(helpers::kv_line("Local Addr", addr)); lines.push(helpers::kv_line("Local Addr", addr));
} }
lines.push(helpers::kv_line("State", helpers::str_field(&transport, "state"))); lines.push(helpers::kv_line(
"State",
helpers::str_field(&transport, "state"),
));
lines.push(Line::from("")); lines.push(Line::from(""));
} }
@@ -268,28 +335,70 @@ fn draw_detail(frame: &mut Frame, app: &App, area: Rect, peers: &[serde_json::Va
Line::from(""), Line::from(""),
// Stats // Stats
helpers::section_header("Link Stats"), helpers::section_header("Link Stats"),
helpers::kv_line("Pkts Sent", &helpers::nested_u64(peer, "stats", "packets_sent")), helpers::kv_line(
helpers::kv_line("Pkts Recv", &helpers::nested_u64(peer, "stats", "packets_recv")), "Pkts Sent",
helpers::kv_line("Bytes Sent", &helpers::format_bytes( &helpers::nested_u64(peer, "stats", "packets_sent"),
peer.get("stats").and_then(|s| s.get("bytes_sent")).and_then(|v| v.as_u64()).unwrap_or(0), ),
)), helpers::kv_line(
helpers::kv_line("Bytes Recv", &helpers::format_bytes( "Pkts Recv",
peer.get("stats").and_then(|s| s.get("bytes_recv")).and_then(|v| v.as_u64()).unwrap_or(0), &helpers::nested_u64(peer, "stats", "packets_recv"),
)), ),
helpers::kv_line(
"Bytes Sent",
&helpers::format_bytes(
peer.get("stats")
.and_then(|s| s.get("bytes_sent"))
.and_then(|v| v.as_u64())
.unwrap_or(0),
),
),
helpers::kv_line(
"Bytes Recv",
&helpers::format_bytes(
peer.get("stats")
.and_then(|s| s.get("bytes_recv"))
.and_then(|v| v.as_u64())
.unwrap_or(0),
),
),
Line::from(""), Line::from(""),
]); ]);
// MMP (if present) // MMP (if present)
if peer.get("mmp").is_some() { if peer.get("mmp").is_some() {
lines.push(helpers::section_header("MMP Metrics")); lines.push(helpers::section_header("MMP Metrics"));
lines.push(helpers::kv_line("Mode", &helpers::nested_str(peer, "mmp", "mode"))); lines.push(helpers::kv_line(
lines.push(helpers::kv_line("SRTT", &format!("{}ms", helpers::nested_f64(peer, "mmp", "srtt_ms", 1)))); "Mode",
lines.push(helpers::kv_line("Loss Rate", &helpers::nested_f64_prefer(peer, "mmp", "smoothed_loss", "loss_rate", 4))); &helpers::nested_str(peer, "mmp", "mode"),
lines.push(helpers::kv_line("ETX", &helpers::nested_f64_prefer(peer, "mmp", "smoothed_etx", "etx", 2))); ));
lines.push(helpers::kv_line("LQI", &helpers::nested_f64(peer, "mmp", "lqi", 2))); lines.push(helpers::kv_line(
lines.push(helpers::kv_line("Goodput", &helpers::nested_throughput(peer, "mmp", "goodput_bps"))); "SRTT",
lines.push(helpers::kv_line("Delivery Fwd", &helpers::nested_f64(peer, "mmp", "delivery_ratio_forward", 3))); &format!("{}ms", helpers::nested_f64(peer, "mmp", "srtt_ms", 1)),
lines.push(helpers::kv_line("Delivery Rev", &helpers::nested_f64(peer, "mmp", "delivery_ratio_reverse", 3))); ));
lines.push(helpers::kv_line(
"Loss Rate",
&helpers::nested_f64_prefer(peer, "mmp", "smoothed_loss", "loss_rate", 4),
));
lines.push(helpers::kv_line(
"ETX",
&helpers::nested_f64_prefer(peer, "mmp", "smoothed_etx", "etx", 2),
));
lines.push(helpers::kv_line(
"LQI",
&helpers::nested_f64(peer, "mmp", "lqi", 2),
));
lines.push(helpers::kv_line(
"Goodput",
&helpers::nested_throughput(peer, "mmp", "goodput_bps"),
));
lines.push(helpers::kv_line(
"Delivery Fwd",
&helpers::nested_f64(peer, "mmp", "delivery_ratio_forward", 3),
));
lines.push(helpers::kv_line(
"Delivery Rev",
&helpers::nested_f64(peer, "mmp", "delivery_ratio_reverse", 3),
));
} }
let detail_scroll = app.detail_view.as_ref().map(|d| d.scroll).unwrap_or(0); let detail_scroll = app.detail_view.as_ref().map(|d| d.scroll).unwrap_or(0);
@@ -320,4 +429,3 @@ fn lookup_transport(app: &App, link: &serde_json::Value) -> Option<serde_json::V
.find(|t| t.get("transport_id").and_then(|v| v.as_u64()) == Some(transport_id)) .find(|t| t.get("transport_id").and_then(|v| v.as_u64()) == Some(transport_id))
.cloned() .cloned()
} }

View File

@@ -12,16 +12,16 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let data = match app.data.get(&Tab::Routing) { let data = match app.data.get(&Tab::Routing) {
Some(d) => d, Some(d) => d,
None => { None => {
let msg = Paragraph::new(" Waiting for data...") let msg =
.style(Style::default().fg(Color::DarkGray)); Paragraph::new(" Waiting for data...").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, area); frame.render_widget(msg, area);
return; return;
} }
}; };
let chunks = Layout::vertical([ let chunks = Layout::vertical([
Constraint::Length(7), // Routing State Constraint::Length(7), // Routing State
Constraint::Length(8), // Coordinate Cache Constraint::Length(8), // Coordinate Cache
Constraint::Min(3), // Routing Statistics Constraint::Min(3), // Routing Statistics
]) ])
.split(area); .split(area);
@@ -33,7 +33,10 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
fn draw_routing_state(frame: &mut Frame, data: &serde_json::Value, area: Rect) { fn draw_routing_state(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
let lines = vec![ let lines = vec![
helpers::kv_line("Coord Cache", &helpers::u64_field(data, "coord_cache_entries")), helpers::kv_line(
"Coord Cache",
&helpers::u64_field(data, "coord_cache_entries"),
),
helpers::kv_line( helpers::kv_line(
"Identity Cache", "Identity Cache",
&helpers::u64_field(data, "identity_cache_entries"), &helpers::u64_field(data, "identity_cache_entries"),
@@ -68,7 +71,10 @@ fn fwd_line(data: &serde_json::Value, label: &str, pkt_key: &str, byte_key: &str
.and_then(|f| f.get(byte_key)) .and_then(|f| f.get(byte_key))
.and_then(|v| v.as_u64()) .and_then(|v| v.as_u64())
.unwrap_or(0); .unwrap_or(0);
helpers::kv_line(label, &format!("{} pkts ({})", pkts, helpers::format_bytes(bytes))) helpers::kv_line(
label,
&format!("{} pkts ({})", pkts, helpers::format_bytes(bytes)),
)
} }
fn draw_routing_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) { fn draw_routing_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
@@ -78,11 +84,8 @@ fn draw_routing_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
let inner = block.inner(area); let inner = block.inner(area);
frame.render_widget(block, area); frame.render_widget(block, area);
let cols = Layout::horizontal([ let cols =
Constraint::Percentage(50), Layout::horizontal([Constraint::Percentage(50), Constraint::Percentage(50)]).split(inner);
Constraint::Percentage(50),
])
.split(inner);
// Left column: Forwarding + Discovery // Left column: Forwarding + Discovery
let mut left = vec![ let mut left = vec![
@@ -91,47 +94,147 @@ fn draw_routing_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
fwd_line(data, "Delivered", "delivered_packets", "delivered_bytes"), fwd_line(data, "Delivered", "delivered_packets", "delivered_bytes"),
fwd_line(data, "Forwarded", "forwarded_packets", "forwarded_bytes"), fwd_line(data, "Forwarded", "forwarded_packets", "forwarded_bytes"),
fwd_line(data, "Originated", "originated_packets", "originated_bytes"), fwd_line(data, "Originated", "originated_packets", "originated_bytes"),
fwd_line(data, "Decode Error", "decode_error_packets", "decode_error_bytes"), fwd_line(
fwd_line(data, "TTL Exhausted", "ttl_exhausted_packets", "ttl_exhausted_bytes"), data,
fwd_line(data, "No Route", "drop_no_route_packets", "drop_no_route_bytes"), "Decode Error",
fwd_line(data, "MTU Exceeded", "drop_mtu_exceeded_packets", "drop_mtu_exceeded_bytes"), "decode_error_packets",
fwd_line(data, "Send Error", "drop_send_error_packets", "drop_send_error_bytes"), "decode_error_bytes",
),
fwd_line(
data,
"TTL Exhausted",
"ttl_exhausted_packets",
"ttl_exhausted_bytes",
),
fwd_line(
data,
"No Route",
"drop_no_route_packets",
"drop_no_route_bytes",
),
fwd_line(
data,
"MTU Exceeded",
"drop_mtu_exceeded_packets",
"drop_mtu_exceeded_bytes",
),
fwd_line(
data,
"Send Error",
"drop_send_error_packets",
"drop_send_error_bytes",
),
Line::from(""), Line::from(""),
helpers::section_header("Discovery Requests"), helpers::section_header("Discovery Requests"),
helpers::kv_line("Received", &helpers::nested_u64(data, "discovery", "req_received")), helpers::kv_line(
helpers::kv_line("Forwarded", &helpers::nested_u64(data, "discovery", "req_forwarded")), "Received",
helpers::kv_line("Initiated", &helpers::nested_u64(data, "discovery", "req_initiated")), &helpers::nested_u64(data, "discovery", "req_received"),
helpers::kv_line("Deduplicated", &helpers::nested_u64(data, "discovery", "req_deduplicated")), ),
helpers::kv_line("Target Is Us", &helpers::nested_u64(data, "discovery", "req_target_is_us")), helpers::kv_line(
helpers::kv_line("Duplicate", &helpers::nested_u64(data, "discovery", "req_duplicate")), "Forwarded",
helpers::kv_line("Bloom Miss", &helpers::nested_u64(data, "discovery", "req_bloom_miss")), &helpers::nested_u64(data, "discovery", "req_forwarded"),
helpers::kv_line("Backoff Suppressed", &helpers::nested_u64(data, "discovery", "req_backoff_suppressed")), ),
helpers::kv_line("Fwd Rate Limited", &helpers::nested_u64(data, "discovery", "req_forward_rate_limited")), helpers::kv_line(
helpers::kv_line("TTL Exhausted", &helpers::nested_u64(data, "discovery", "req_ttl_exhausted")), "Initiated",
helpers::kv_line("Decode Error", &helpers::nested_u64(data, "discovery", "req_decode_error")), &helpers::nested_u64(data, "discovery", "req_initiated"),
),
helpers::kv_line(
"Deduplicated",
&helpers::nested_u64(data, "discovery", "req_deduplicated"),
),
helpers::kv_line(
"Target Is Us",
&helpers::nested_u64(data, "discovery", "req_target_is_us"),
),
helpers::kv_line(
"Duplicate",
&helpers::nested_u64(data, "discovery", "req_duplicate"),
),
helpers::kv_line(
"Bloom Miss",
&helpers::nested_u64(data, "discovery", "req_bloom_miss"),
),
helpers::kv_line(
"Backoff Suppressed",
&helpers::nested_u64(data, "discovery", "req_backoff_suppressed"),
),
helpers::kv_line(
"Fwd Rate Limited",
&helpers::nested_u64(data, "discovery", "req_forward_rate_limited"),
),
helpers::kv_line(
"TTL Exhausted",
&helpers::nested_u64(data, "discovery", "req_ttl_exhausted"),
),
helpers::kv_line(
"Decode Error",
&helpers::nested_u64(data, "discovery", "req_decode_error"),
),
Line::from(""), Line::from(""),
helpers::section_header("Discovery Responses"), helpers::section_header("Discovery Responses"),
helpers::kv_line("Received", &helpers::nested_u64(data, "discovery", "resp_received")), helpers::kv_line(
helpers::kv_line("Accepted", &helpers::nested_u64(data, "discovery", "resp_accepted")), "Received",
helpers::kv_line("Forwarded", &helpers::nested_u64(data, "discovery", "resp_forwarded")), &helpers::nested_u64(data, "discovery", "resp_received"),
helpers::kv_line("Timed Out", &helpers::nested_u64(data, "discovery", "resp_timed_out")), ),
helpers::kv_line("Identity Miss", &helpers::nested_u64(data, "discovery", "resp_identity_miss")), helpers::kv_line(
helpers::kv_line("Proof Failed", &helpers::nested_u64(data, "discovery", "resp_proof_failed")), "Accepted",
helpers::kv_line("Decode Error", &helpers::nested_u64(data, "discovery", "resp_decode_error")), &helpers::nested_u64(data, "discovery", "resp_accepted"),
),
helpers::kv_line(
"Forwarded",
&helpers::nested_u64(data, "discovery", "resp_forwarded"),
),
helpers::kv_line(
"Timed Out",
&helpers::nested_u64(data, "discovery", "resp_timed_out"),
),
helpers::kv_line(
"Identity Miss",
&helpers::nested_u64(data, "discovery", "resp_identity_miss"),
),
helpers::kv_line(
"Proof Failed",
&helpers::nested_u64(data, "discovery", "resp_proof_failed"),
),
helpers::kv_line(
"Decode Error",
&helpers::nested_u64(data, "discovery", "resp_decode_error"),
),
]; ];
// Right column: Error Signals + Congestion // Right column: Error Signals + Congestion
let mut right = vec![ let mut right = vec![
helpers::section_header("Error Signals"), helpers::section_header("Error Signals"),
helpers::kv_line("Coords Required", &helpers::nested_u64(data, "error_signals", "coords_required")), helpers::kv_line(
helpers::kv_line("Path Broken", &helpers::nested_u64(data, "error_signals", "path_broken")), "Coords Required",
helpers::kv_line("MTU Exceeded", &helpers::nested_u64(data, "error_signals", "mtu_exceeded")), &helpers::nested_u64(data, "error_signals", "coords_required"),
),
helpers::kv_line(
"Path Broken",
&helpers::nested_u64(data, "error_signals", "path_broken"),
),
helpers::kv_line(
"MTU Exceeded",
&helpers::nested_u64(data, "error_signals", "mtu_exceeded"),
),
Line::from(""), Line::from(""),
helpers::section_header("Congestion"), helpers::section_header("Congestion"),
helpers::kv_line("CE Forwarded", &helpers::nested_u64(data, "congestion", "ce_forwarded")), helpers::kv_line(
helpers::kv_line("CE Received", &helpers::nested_u64(data, "congestion", "ce_received")), "CE Forwarded",
helpers::kv_line("Congestion Detected", &helpers::nested_u64(data, "congestion", "congestion_detected")), &helpers::nested_u64(data, "congestion", "ce_forwarded"),
helpers::kv_line("Kernel Drops", &helpers::nested_u64(data, "congestion", "kernel_drop_events")), ),
helpers::kv_line(
"CE Received",
&helpers::nested_u64(data, "congestion", "ce_received"),
),
helpers::kv_line(
"Congestion Detected",
&helpers::nested_u64(data, "congestion", "congestion_detected"),
),
helpers::kv_line(
"Kernel Drops",
&helpers::nested_u64(data, "congestion", "kernel_drop_events"),
),
]; ];
let max_lines = cols[0].height as usize; let max_lines = cols[0].height as usize;

View File

@@ -15,11 +15,8 @@ pub fn draw(frame: &mut Frame, app: &mut App, area: Rect) {
let row_count = sessions.len(); let row_count = sessions.len();
if app.detail_view.is_some() { if app.detail_view.is_some() {
let chunks = Layout::horizontal([ let chunks = Layout::horizontal([Constraint::Percentage(40), Constraint::Percentage(60)])
Constraint::Percentage(40), .split(area);
Constraint::Percentage(60),
])
.split(area);
draw_table(frame, app, chunks[0], &sessions, row_count); draw_table(frame, app, chunks[0], &sessions, row_count);
draw_detail(frame, app, chunks[1], &sessions); draw_detail(frame, app, chunks[1], &sessions);
@@ -223,10 +220,7 @@ fn draw_detail(frame: &mut Frame, app: &App, area: Rect, sessions: &[serde_json:
)); ));
lines.push(helpers::kv_line( lines.push(helpers::kv_line(
"SRTT", "SRTT",
&format!( &format!("{}ms", helpers::nested_f64(session, "mmp", "srtt_ms", 1)),
"{}ms",
helpers::nested_f64(session, "mmp", "srtt_ms", 1)
),
)); ));
lines.push(helpers::kv_line( lines.push(helpers::kv_line(
"Loss Rate", "Loss Rate",
@@ -271,4 +265,3 @@ fn state_styled(state: &str) -> Span<'static> {
}; };
Span::styled(state.to_string(), Style::default().fg(color)) Span::styled(state.to_string(), Style::default().fg(color))
} }

View File

@@ -33,11 +33,8 @@ pub fn draw(frame: &mut Frame, app: &mut App, area: Rect) {
update_selected_tree_item(app, &tree_rows); update_selected_tree_item(app, &tree_rows);
if app.detail_view.is_some() { if app.detail_view.is_some() {
let chunks = Layout::horizontal([ let chunks = Layout::horizontal([Constraint::Percentage(40), Constraint::Percentage(60)])
Constraint::Percentage(40), .split(area);
Constraint::Percentage(60),
])
.split(area);
draw_table(frame, app, chunks[0], &transports, &links, &tree_rows); draw_table(frame, app, chunks[0], &transports, &links, &tree_rows);
draw_detail(frame, app, chunks[1], &transports, &links, &tree_rows); draw_detail(frame, app, chunks[1], &transports, &links, &tree_rows);
@@ -110,9 +107,7 @@ fn update_selected_tree_item(app: &mut App, tree_rows: &[TreeRow]) {
.unwrap_or(0); .unwrap_or(0);
app.selected_tree_item = match tree_rows.get(selected) { app.selected_tree_item = match tree_rows.get(selected) {
Some(TreeRow::Transport { transport_id, .. }) => { Some(TreeRow::Transport { transport_id, .. }) => SelectedTreeItem::Transport(*transport_id),
SelectedTreeItem::Transport(*transport_id)
}
Some(TreeRow::Link { .. }) => SelectedTreeItem::Link, Some(TreeRow::Link { .. }) => SelectedTreeItem::Link,
None => SelectedTreeItem::None, None => SelectedTreeItem::None,
}; };
@@ -169,11 +164,7 @@ fn draw_table(
.get("onion_address") .get("onion_address")
.and_then(|v| v.as_str()) .and_then(|v| v.as_str())
.map(|a| { .map(|a| {
let short = if a.len() > 16 { let short = if a.len() > 16 { &a[..16] } else { a };
&a[..16]
} else {
a
};
format!(" {short}..") format!(" {short}..")
}) })
.unwrap_or_default(); .unwrap_or_default();
@@ -209,7 +200,11 @@ fn draw_table(
} }
TreeRow::Link { index, is_last } => { TreeRow::Link { index, is_last } => {
let link = &links[*index]; let link = &links[*index];
let tree_char = if *is_last { "\u{2514}\u{2500}" } else { "\u{251C}\u{2500}" }; // └─ or ├─ let tree_char = if *is_last {
"\u{2514}\u{2500}"
} else {
"\u{251C}\u{2500}"
}; // └─ or ├─
let dir = helpers::str_field(link, "direction"); let dir = helpers::str_field(link, "direction");
let dir_short = match dir { let dir_short = match dir {
"Outbound" => "Out", "Outbound" => "Out",
@@ -303,13 +298,10 @@ fn draw_detail(
.unwrap_or(0); .unwrap_or(0);
let Some(tree_row) = tree_rows.get(selected) else { let Some(tree_row) = tree_rows.get(selected) else {
let block = Block::default() let block = Block::default().borders(Borders::ALL).title(" Detail ");
.borders(Borders::ALL)
.title(" Detail ");
let inner = block.inner(area); let inner = block.inner(area);
frame.render_widget(block, area); frame.render_widget(block, area);
let msg = Paragraph::new(" No item selected") let msg = Paragraph::new(" No item selected").style(Style::default().fg(Color::DarkGray));
.style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, inner); frame.render_widget(msg, inner);
return; return;
}; };
@@ -539,7 +531,10 @@ fn draw_transport_detail(frame: &mut Frame, app: &App, area: Rect, t: &serde_jso
"Network", "Network",
&helpers::nested_str(t, "tor_monitoring", "network_liveness"), &helpers::nested_str(t, "tor_monitoring", "network_liveness"),
)); ));
lines.push(helpers::kv_line("Dormant", helpers::bool_field(mon, "dormant"))); lines.push(helpers::kv_line(
"Dormant",
helpers::bool_field(mon, "dormant"),
));
let tor_read = mon let tor_read = mon
.get("traffic_read") .get("traffic_read")

View File

@@ -12,8 +12,8 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let data = match app.data.get(&Tab::Tree) { let data = match app.data.get(&Tab::Tree) {
Some(d) => d, Some(d) => d,
None => { None => {
let msg = Paragraph::new(" Waiting for data...") let msg =
.style(Style::default().fg(Color::DarkGray)); Paragraph::new(" Waiting for data...").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, area); frame.render_widget(msg, area);
return; return;
} }
@@ -22,7 +22,7 @@ pub fn draw(frame: &mut Frame, app: &App, area: Rect) {
let chunks = Layout::vertical([ let chunks = Layout::vertical([
Constraint::Length(10), // Tree Position Constraint::Length(10), // Tree Position
Constraint::Length(22), // Tree Announce Stats Constraint::Length(22), // Tree Announce Stats
Constraint::Min(3), // Tree Peers Constraint::Min(3), // Tree Peers
]) ])
.split(area); .split(area);
@@ -70,16 +70,12 @@ fn draw_position(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
)]; )];
if coords.is_empty() { if coords.is_empty() {
path_parts.push(Span::styled( path_parts.push(Span::styled("[root]", Style::default().fg(Color::Yellow)));
"[root]",
Style::default().fg(Color::Yellow),
));
} else { } else {
// Reverse: root first, self last // Reverse: root first, self last
for (i, entry) in coords.iter().rev().enumerate() { for (i, entry) in coords.iter().rev().enumerate() {
if i > 0 { if i > 0 {
path_parts path_parts.push(Span::styled(" > ", Style::default().fg(Color::DarkGray)));
.push(Span::styled(" > ", Style::default().fg(Color::DarkGray)));
} }
let hex = entry.as_str().unwrap_or("-"); let hex = entry.as_str().unwrap_or("-");
let color = if i == 0 { let color = if i == 0 {
@@ -87,8 +83,10 @@ fn draw_position(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
} else { } else {
Color::White Color::White
}; };
path_parts path_parts.push(Span::styled(
.push(Span::styled(helpers::truncate_hex(hex, 8), Style::default().fg(color))); helpers::truncate_hex(hex, 8),
Style::default().fg(color),
));
} }
path_parts.push(Span::styled(" > ", Style::default().fg(Color::DarkGray))); path_parts.push(Span::styled(" > ", Style::default().fg(Color::DarkGray)));
path_parts.push(Span::styled( path_parts.push(Span::styled(
@@ -121,24 +119,60 @@ fn draw_stats(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
helpers::section_header("Inbound"), helpers::section_header("Inbound"),
helpers::kv_line("Received", &helpers::nested_u64(data, "stats", "received")), helpers::kv_line("Received", &helpers::nested_u64(data, "stats", "received")),
helpers::kv_line("Accepted", &helpers::nested_u64(data, "stats", "accepted")), helpers::kv_line("Accepted", &helpers::nested_u64(data, "stats", "accepted")),
helpers::kv_line("Decode Error", &helpers::nested_u64(data, "stats", "decode_error")), helpers::kv_line(
helpers::kv_line("Unknown Peer", &helpers::nested_u64(data, "stats", "unknown_peer")), "Decode Error",
helpers::kv_line("Addr Mismatch", &helpers::nested_u64(data, "stats", "addr_mismatch")), &helpers::nested_u64(data, "stats", "decode_error"),
helpers::kv_line("Sig Failed", &helpers::nested_u64(data, "stats", "sig_failed")), ),
helpers::kv_line(
"Unknown Peer",
&helpers::nested_u64(data, "stats", "unknown_peer"),
),
helpers::kv_line(
"Addr Mismatch",
&helpers::nested_u64(data, "stats", "addr_mismatch"),
),
helpers::kv_line(
"Sig Failed",
&helpers::nested_u64(data, "stats", "sig_failed"),
),
helpers::kv_line("Stale", &helpers::nested_u64(data, "stats", "stale")), helpers::kv_line("Stale", &helpers::nested_u64(data, "stats", "stale")),
helpers::kv_line("Parent Switched", &helpers::nested_u64(data, "stats", "parent_switched")), helpers::kv_line(
helpers::kv_line("Loop Detected", &helpers::nested_u64(data, "stats", "loop_detected")), "Parent Switched",
helpers::kv_line("Ancestry Changed", &helpers::nested_u64(data, "stats", "ancestry_changed")), &helpers::nested_u64(data, "stats", "parent_switched"),
),
helpers::kv_line(
"Loop Detected",
&helpers::nested_u64(data, "stats", "loop_detected"),
),
helpers::kv_line(
"Ancestry Changed",
&helpers::nested_u64(data, "stats", "ancestry_changed"),
),
Line::from(""), Line::from(""),
helpers::section_header("Outbound"), helpers::section_header("Outbound"),
helpers::kv_line("Sent", &helpers::nested_u64(data, "stats", "sent")), helpers::kv_line("Sent", &helpers::nested_u64(data, "stats", "sent")),
helpers::kv_line("Rate Limited", &helpers::nested_u64(data, "stats", "rate_limited")), helpers::kv_line(
helpers::kv_line("Send Failed", &helpers::nested_u64(data, "stats", "send_failed")), "Rate Limited",
&helpers::nested_u64(data, "stats", "rate_limited"),
),
helpers::kv_line(
"Send Failed",
&helpers::nested_u64(data, "stats", "send_failed"),
),
Line::from(""), Line::from(""),
helpers::section_header("Cumulative"), helpers::section_header("Cumulative"),
helpers::kv_line("Parent Switches", &helpers::nested_u64(data, "stats", "parent_switches")), helpers::kv_line(
helpers::kv_line("Parent Losses", &helpers::nested_u64(data, "stats", "parent_losses")), "Parent Switches",
helpers::kv_line("Flap Dampened", &helpers::nested_u64(data, "stats", "flap_dampened")), &helpers::nested_u64(data, "stats", "parent_switches"),
),
helpers::kv_line(
"Parent Losses",
&helpers::nested_u64(data, "stats", "parent_losses"),
),
helpers::kv_line(
"Flap Dampened",
&helpers::nested_u64(data, "stats", "flap_dampened"),
),
]; ];
// Trim to fit available height // Trim to fit available height
@@ -164,8 +198,7 @@ fn draw_peers(frame: &mut Frame, data: &serde_json::Value, area: Rect) {
frame.render_widget(block, area); frame.render_widget(block, area);
if peers.is_empty() { if peers.is_empty() {
let msg = let msg = Paragraph::new(" No peers").style(Style::default().fg(Color::DarkGray));
Paragraph::new(" No peers").style(Style::default().fg(Color::DarkGray));
frame.render_widget(msg, inner); frame.render_widget(msg, inner);
return; return;
} }

View File

@@ -149,8 +149,7 @@ fn test_bloom_filter_from_bytes() {
let original = BloomFilter::new(); let original = BloomFilter::new();
let bytes = original.as_bytes().to_vec(); let bytes = original.as_bytes().to_vec();
let restored = let restored = BloomFilter::from_bytes(bytes, original.hash_count()).unwrap();
BloomFilter::from_bytes(bytes, original.hash_count()).unwrap();
assert_eq!(original, restored); assert_eq!(original, restored);
} }

View File

@@ -6,10 +6,10 @@
use std::collections::HashMap; use std::collections::HashMap;
use super::entry::CacheEntry;
use super::CacheStats; use super::CacheStats;
use crate::tree::TreeCoordinate; use super::entry::CacheEntry;
use crate::NodeAddr; use crate::NodeAddr;
use crate::tree::TreeCoordinate;
/// Default maximum entries in coordinate cache. /// Default maximum entries in coordinate cache.
pub const DEFAULT_COORD_CACHE_SIZE: usize = 50_000; pub const DEFAULT_COORD_CACHE_SIZE: usize = 50_000;

View File

@@ -34,7 +34,10 @@ pub use node::{
TreeConfig, TreeConfig,
}; };
pub use peer::{ConnectPolicy, PeerAddress, PeerConfig}; pub use peer::{ConnectPolicy, PeerAddress, PeerConfig};
pub use transport::{DirectoryServiceConfig, EthernetConfig, TcpConfig, TorConfig, TransportInstances, TransportsConfig, UdpConfig}; pub use transport::{
DirectoryServiceConfig, EthernetConfig, TcpConfig, TorConfig, TransportInstances,
TransportsConfig, UdpConfig,
};
/// Default config filename. /// Default config filename.
const CONFIG_FILENAME: &str = "fips.yaml"; const CONFIG_FILENAME: &str = "fips.yaml";
@@ -539,10 +542,7 @@ node:
override_config.node.identity.nsec = Some("override_nsec".to_string()); override_config.node.identity.nsec = Some("override_nsec".to_string());
base.merge(override_config); base.merge(override_config);
assert_eq!( assert_eq!(base.node.identity.nsec, Some("override_nsec".to_string()));
base.node.identity.nsec,
Some("override_nsec".to_string())
);
} }
#[test] #[test]
@@ -559,9 +559,8 @@ node:
#[test] #[test]
fn test_create_identity_from_nsec() { fn test_create_identity_from_nsec() {
let mut config = Config::new(); let mut config = Config::new();
config.node.identity.nsec = Some( config.node.identity.nsec =
"0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20".to_string(), Some("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20".to_string());
);
let identity = config.create_identity().unwrap(); let identity = config.create_identity().unwrap();
assert!(!identity.npub().is_empty()); assert!(!identity.npub().is_empty());
@@ -660,9 +659,11 @@ node:
assert!(paths.iter().any(|p| p.ends_with("fips.yaml"))); assert!(paths.iter().any(|p| p.ends_with("fips.yaml")));
// Should include /etc/fips // Should include /etc/fips
assert!(paths assert!(
.iter() paths
.any(|p| p.starts_with("/etc/fips") && p.ends_with("fips.yaml"))); .iter()
.any(|p| p.starts_with("/etc/fips") && p.ends_with("fips.yaml"))
);
} }
#[test] #[test]
@@ -747,16 +748,21 @@ node:
#[test] #[test]
fn test_key_file_path_derivation() { fn test_key_file_path_derivation() {
let config_path = PathBuf::from("/etc/fips/fips.yaml"); let config_path = PathBuf::from("/etc/fips/fips.yaml");
assert_eq!(key_file_path(&config_path), PathBuf::from("/etc/fips/fips.key")); assert_eq!(
assert_eq!(pub_file_path(&config_path), PathBuf::from("/etc/fips/fips.pub")); key_file_path(&config_path),
PathBuf::from("/etc/fips/fips.key")
);
assert_eq!(
pub_file_path(&config_path),
PathBuf::from("/etc/fips/fips.pub")
);
} }
#[test] #[test]
fn test_resolve_identity_from_config() { fn test_resolve_identity_from_config() {
let mut config = Config::new(); let mut config = Config::new();
config.node.identity.nsec = Some( config.node.identity.nsec =
"0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20".to_string(), Some("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20".to_string());
);
let resolved = resolve_identity(&config, &[]).unwrap(); let resolved = resolve_identity(&config, &[]).unwrap();
assert!(matches!(resolved.source, IdentitySource::Config)); assert!(matches!(resolved.source, IdentitySource::Config));
@@ -803,11 +809,7 @@ node:
let config_path = temp_dir.path().join("fips.yaml"); let config_path = temp_dir.path().join("fips.yaml");
let key_path = temp_dir.path().join("fips.key"); let key_path = temp_dir.path().join("fips.key");
fs::write( fs::write(&config_path, "node:\n identity:\n persistent: true\n").unwrap();
&config_path,
"node:\n identity:\n persistent: true\n",
)
.unwrap();
// Write a key file // Write a key file
let identity = crate::Identity::generate(); let identity = crate::Identity::generate();
@@ -827,11 +829,7 @@ node:
let temp_dir = TempDir::new().unwrap(); let temp_dir = TempDir::new().unwrap();
let config_path = temp_dir.path().join("fips.yaml"); let config_path = temp_dir.path().join("fips.yaml");
fs::write( fs::write(&config_path, "node:\n identity:\n persistent: true\n").unwrap();
&config_path,
"node:\n identity:\n persistent: true\n",
)
.unwrap();
let config = Config::load_file(&config_path).unwrap(); let config = Config::load_file(&config_path).unwrap();
let resolved = resolve_identity(&config, std::slice::from_ref(&config_path)).unwrap(); let resolved = resolve_identity(&config, std::slice::from_ref(&config_path)).unwrap();
@@ -892,8 +890,7 @@ transports:
assert_eq!(config.transports.udp.len(), 2); assert_eq!(config.transports.udp.len(), 2);
let instances: std::collections::HashMap<_, _> = let instances: std::collections::HashMap<_, _> = config.transports.udp.iter().collect();
config.transports.udp.iter().collect();
// Named instances have Some(name) // Named instances have Some(name)
assert!(instances.contains_key(&Some("main"))); assert!(instances.contains_key(&Some("main")));

View File

@@ -7,7 +7,7 @@
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use super::IdentityConfig; use super::IdentityConfig;
use crate::mmp::{MmpConfig, MmpMode, DEFAULT_LOG_INTERVAL_SECS, DEFAULT_OWD_WINDOW_SIZE}; use crate::mmp::{DEFAULT_LOG_INTERVAL_SECS, DEFAULT_OWD_WINDOW_SIZE, MmpConfig, MmpMode};
// ============================================================================ // ============================================================================
// Node Configuration Subsections // Node Configuration Subsections
@@ -42,10 +42,18 @@ impl Default for LimitsConfig {
} }
impl LimitsConfig { impl LimitsConfig {
fn default_max_connections() -> usize { 256 } fn default_max_connections() -> usize {
fn default_max_peers() -> usize { 128 } 256
fn default_max_links() -> usize { 256 } }
fn default_max_pending_inbound() -> usize { 1000 } fn default_max_peers() -> usize {
128
}
fn default_max_links() -> usize {
256
}
fn default_max_pending_inbound() -> usize {
1000
}
} }
/// Rate limiting (`node.rate_limit.*`). /// Rate limiting (`node.rate_limit.*`).
@@ -86,12 +94,24 @@ impl Default for RateLimitConfig {
} }
impl RateLimitConfig { impl RateLimitConfig {
fn default_handshake_burst() -> u32 { 100 } fn default_handshake_burst() -> u32 {
fn default_handshake_rate() -> f64 { 10.0 } 100
fn default_handshake_timeout_secs() -> u64 { 30 } }
fn default_handshake_resend_interval_ms() -> u64 { 1000 } fn default_handshake_rate() -> f64 {
fn default_handshake_resend_backoff() -> f64 { 2.0 } 10.0
fn default_handshake_max_resends() -> u32 { 5 } }
fn default_handshake_timeout_secs() -> u64 {
30
}
fn default_handshake_resend_interval_ms() -> u64 {
1000
}
fn default_handshake_resend_backoff() -> f64 {
2.0
}
fn default_handshake_max_resends() -> u32 {
5
}
} }
/// Retry/backoff configuration (`node.retry.*`). /// Retry/backoff configuration (`node.retry.*`).
@@ -119,9 +139,15 @@ impl Default for RetryConfig {
} }
impl RetryConfig { impl RetryConfig {
fn default_max_retries() -> u32 { 5 } fn default_max_retries() -> u32 {
fn default_base_interval_secs() -> u64 { 5 } 5
fn default_max_backoff_secs() -> u64 { 300 } }
fn default_base_interval_secs() -> u64 {
5
}
fn default_max_backoff_secs() -> u64 {
300
}
} }
/// Cache parameters (`node.cache.*`). /// Cache parameters (`node.cache.*`).
@@ -149,9 +175,15 @@ impl Default for CacheConfig {
} }
impl CacheConfig { impl CacheConfig {
fn default_coord_size() -> usize { 50_000 } fn default_coord_size() -> usize {
fn default_coord_ttl_secs() -> u64 { 300 } 50_000
fn default_identity_size() -> usize { 10_000 } }
fn default_coord_ttl_secs() -> u64 {
300
}
fn default_identity_size() -> usize {
10_000
}
} }
/// Discovery protocol (`node.discovery.*`). /// Discovery protocol (`node.discovery.*`).
@@ -205,14 +237,30 @@ impl Default for DiscoveryConfig {
} }
impl DiscoveryConfig { impl DiscoveryConfig {
fn default_ttl() -> u8 { 64 } fn default_ttl() -> u8 {
fn default_timeout_secs() -> u64 { 10 } 64
fn default_recent_expiry_secs() -> u64 { 10 } }
fn default_backoff_base_secs() -> u64 { 30 } fn default_timeout_secs() -> u64 {
fn default_backoff_max_secs() -> u64 { 300 } 10
fn default_forward_min_interval_secs() -> u64 { 2 } }
fn default_retry_interval_secs() -> u64 { 5 } fn default_recent_expiry_secs() -> u64 {
fn default_max_attempts() -> u8 { 2 } 10
}
fn default_backoff_base_secs() -> u64 {
30
}
fn default_backoff_max_secs() -> u64 {
300
}
fn default_forward_min_interval_secs() -> u64 {
2
}
fn default_retry_interval_secs() -> u64 {
5
}
fn default_max_attempts() -> u8 {
2
}
} }
/// Spanning tree (`node.tree.*`). /// Spanning tree (`node.tree.*`).
@@ -267,13 +315,27 @@ impl Default for TreeConfig {
} }
impl TreeConfig { impl TreeConfig {
fn default_announce_min_interval_ms() -> u64 { 500 } fn default_announce_min_interval_ms() -> u64 {
fn default_parent_hysteresis() -> f64 { 0.2 } 500
fn default_hold_down_secs() -> u64 { 30 } }
fn default_reeval_interval_secs() -> u64 { 60 } fn default_parent_hysteresis() -> f64 {
fn default_flap_threshold() -> u32 { 4 } 0.2
fn default_flap_window_secs() -> u64 { 60 } }
fn default_flap_dampening_secs() -> u64 { 120 } fn default_hold_down_secs() -> u64 {
30
}
fn default_reeval_interval_secs() -> u64 {
60
}
fn default_flap_threshold() -> u32 {
4
}
fn default_flap_window_secs() -> u64 {
60
}
fn default_flap_dampening_secs() -> u64 {
120
}
} }
/// Bloom filter (`node.bloom.*`). /// Bloom filter (`node.bloom.*`).
@@ -286,12 +348,16 @@ pub struct BloomConfig {
impl Default for BloomConfig { impl Default for BloomConfig {
fn default() -> Self { fn default() -> Self {
Self { update_debounce_ms: 500 } Self {
update_debounce_ms: 500,
}
} }
} }
impl BloomConfig { impl BloomConfig {
fn default_update_debounce_ms() -> u64 { 500 } fn default_update_debounce_ms() -> u64 {
500
}
} }
/// Session/data plane (`node.session.*`). /// Session/data plane (`node.session.*`).
@@ -337,12 +403,24 @@ impl Default for SessionConfig {
} }
impl SessionConfig { impl SessionConfig {
fn default_ttl() -> u8 { 64 } fn default_ttl() -> u8 {
fn default_pending_packets_per_dest() -> usize { 16 } 64
fn default_pending_max_destinations() -> usize { 256 } }
fn default_idle_timeout_secs() -> u64 { 90 } fn default_pending_packets_per_dest() -> usize {
fn default_coords_warmup_packets() -> u8 { 5 } 16
fn default_coords_response_interval_ms() -> u64 { 2000 } }
fn default_pending_max_destinations() -> usize {
256
}
fn default_idle_timeout_secs() -> u64 {
90
}
fn default_coords_warmup_packets() -> u8 {
5
}
fn default_coords_response_interval_ms() -> u64 {
2000
}
} }
/// Session-layer Metrics Measurement Protocol (`node.session_mmp.*`). /// Session-layer Metrics Measurement Protocol (`node.session_mmp.*`).
@@ -377,8 +455,12 @@ impl Default for SessionMmpConfig {
} }
impl SessionMmpConfig { impl SessionMmpConfig {
fn default_log_interval_secs() -> u64 { DEFAULT_LOG_INTERVAL_SECS } fn default_log_interval_secs() -> u64 {
fn default_owd_window_size() -> usize { DEFAULT_OWD_WINDOW_SIZE } DEFAULT_LOG_INTERVAL_SECS
}
fn default_owd_window_size() -> usize {
DEFAULT_OWD_WINDOW_SIZE
}
} }
/// Control socket configuration (`node.control.*`). /// Control socket configuration (`node.control.*`).
@@ -402,7 +484,9 @@ impl Default for ControlConfig {
} }
impl ControlConfig { impl ControlConfig {
fn default_enabled() -> bool { true } fn default_enabled() -> bool {
true
}
fn default_socket_path() -> String { fn default_socket_path() -> String {
if let Ok(runtime_dir) = std::env::var("XDG_RUNTIME_DIR") { if let Ok(runtime_dir) = std::env::var("XDG_RUNTIME_DIR") {
@@ -440,9 +524,15 @@ impl Default for BuffersConfig {
} }
impl BuffersConfig { impl BuffersConfig {
fn default_packet_channel() -> usize { 1024 } fn default_packet_channel() -> usize {
fn default_tun_channel() -> usize { 1024 } 1024
fn default_dns_channel() -> usize { 64 } }
fn default_tun_channel() -> usize {
1024
}
fn default_dns_channel() -> usize {
64
}
} }
// ============================================================================ // ============================================================================
@@ -480,9 +570,15 @@ impl Default for RekeyConfig {
} }
impl RekeyConfig { impl RekeyConfig {
fn default_enabled() -> bool { true } fn default_enabled() -> bool {
fn default_after_secs() -> u64 { 120 } true
fn default_after_messages() -> u64 { 1 << 16 } }
fn default_after_secs() -> u64 {
120
}
fn default_after_messages() -> u64 {
1 << 16
}
} }
/// ECN congestion signaling configuration (`node.ecn.*`). /// ECN congestion signaling configuration (`node.ecn.*`).
@@ -520,9 +616,15 @@ impl Default for EcnConfig {
} }
impl EcnConfig { impl EcnConfig {
fn default_enabled() -> bool { true } fn default_enabled() -> bool {
fn default_loss_threshold() -> f64 { 0.05 } true
fn default_etx_threshold() -> f64 { 3.0 } }
fn default_loss_threshold() -> f64 {
0.05
}
fn default_etx_threshold() -> f64 {
3.0
}
} }
// ============================================================================ // ============================================================================
@@ -642,10 +744,18 @@ impl Default for NodeConfig {
} }
impl NodeConfig { impl NodeConfig {
fn default_tick_interval_secs() -> u64 { 1 } fn default_tick_interval_secs() -> u64 {
fn default_base_rtt_ms() -> u64 { 100 } 1
fn default_heartbeat_interval_secs() -> u64 { 10 } }
fn default_link_dead_timeout_secs() -> u64 { 30 } fn default_base_rtt_ms() -> u64 {
100
}
fn default_heartbeat_interval_secs() -> u64 {
10
}
fn default_link_dead_timeout_secs() -> u64 {
30
}
} }
#[cfg(test)] #[cfg(test)]

View File

@@ -66,7 +66,11 @@ impl PeerAddress {
} }
/// Create a new peer address with priority. /// Create a new peer address with priority.
pub fn with_priority(transport: impl Into<String>, addr: impl Into<String>, priority: u8) -> Self { pub fn with_priority(
transport: impl Into<String>,
addr: impl Into<String>,
priority: u8,
) -> Self {
Self { Self {
transport: transport.into(), transport: transport.into(),
addr: addr.into(), addr: addr.into(),
@@ -118,7 +122,11 @@ impl Default for PeerConfig {
impl PeerConfig { impl PeerConfig {
/// Create a new peer config with a single address. /// Create a new peer config with a single address.
pub fn new(npub: impl Into<String>, transport: impl Into<String>, addr: impl Into<String>) -> Self { pub fn new(
npub: impl Into<String>,
transport: impl Into<String>,
addr: impl Into<String>,
) -> Self {
Self { Self {
npub: npub.into(), npub: npub.into(),
alias: None, alias: None,

View File

@@ -112,15 +112,12 @@ impl<T> TransportInstances<T> {
/// Named instances have `Some(name)`. /// Named instances have `Some(name)`.
pub fn iter(&self) -> impl Iterator<Item = (Option<&str>, &T)> { pub fn iter(&self) -> impl Iterator<Item = (Option<&str>, &T)> {
match self { match self {
TransportInstances::Single(config) => { TransportInstances::Single(config) => vec![(None, config)].into_iter(),
vec![(None, config)].into_iter() TransportInstances::Named(map) => map
} .iter()
TransportInstances::Named(map) => { .map(|(k, v)| (Some(k.as_str()), v))
map.iter() .collect::<Vec<_>>()
.map(|(k, v)| (Some(k.as_str()), v)) .into_iter(),
.collect::<Vec<_>>()
.into_iter()
}
} }
} }
} }
@@ -306,7 +303,8 @@ impl TcpConfig {
/// Get the connect timeout in milliseconds. /// Get the connect timeout in milliseconds.
pub fn connect_timeout_ms(&self) -> u64 { pub fn connect_timeout_ms(&self) -> u64 {
self.connect_timeout_ms.unwrap_or(DEFAULT_TCP_CONNECT_TIMEOUT_MS) self.connect_timeout_ms
.unwrap_or(DEFAULT_TCP_CONNECT_TIMEOUT_MS)
} }
/// Whether TCP_NODELAY is enabled. Default: true. /// Whether TCP_NODELAY is enabled. Default: true.
@@ -331,7 +329,8 @@ impl TcpConfig {
/// Get the maximum number of inbound connections. Default: 256. /// Get the maximum number of inbound connections. Default: 256.
pub fn max_inbound_connections(&self) -> usize { pub fn max_inbound_connections(&self) -> usize {
self.max_inbound_connections.unwrap_or(DEFAULT_TCP_MAX_INBOUND) self.max_inbound_connections
.unwrap_or(DEFAULT_TCP_MAX_INBOUND)
} }
} }
@@ -447,12 +446,16 @@ pub struct DirectoryServiceConfig {
impl DirectoryServiceConfig { impl DirectoryServiceConfig {
/// Path to the hostname file. Default: "/var/lib/tor/fips_onion_service/hostname". /// Path to the hostname file. Default: "/var/lib/tor/fips_onion_service/hostname".
pub fn hostname_file(&self) -> &str { pub fn hostname_file(&self) -> &str {
self.hostname_file.as_deref().unwrap_or(DEFAULT_HOSTNAME_FILE) self.hostname_file
.as_deref()
.unwrap_or(DEFAULT_HOSTNAME_FILE)
} }
/// Local bind address for the listener. Default: "127.0.0.1:8443". /// Local bind address for the listener. Default: "127.0.0.1:8443".
pub fn bind_addr(&self) -> &str { pub fn bind_addr(&self) -> &str {
self.bind_addr.as_deref().unwrap_or(DEFAULT_DIRECTORY_BIND_ADDR) self.bind_addr
.as_deref()
.unwrap_or(DEFAULT_DIRECTORY_BIND_ADDR)
} }
} }
@@ -464,12 +467,16 @@ impl TorConfig {
/// Get the SOCKS5 proxy address. Default: "127.0.0.1:9050". /// Get the SOCKS5 proxy address. Default: "127.0.0.1:9050".
pub fn socks5_addr(&self) -> &str { pub fn socks5_addr(&self) -> &str {
self.socks5_addr.as_deref().unwrap_or(DEFAULT_TOR_SOCKS5_ADDR) self.socks5_addr
.as_deref()
.unwrap_or(DEFAULT_TOR_SOCKS5_ADDR)
} }
/// Get the control port address. Default: "/run/tor/control". /// Get the control port address. Default: "/run/tor/control".
pub fn control_addr(&self) -> &str { pub fn control_addr(&self) -> &str {
self.control_addr.as_deref().unwrap_or(DEFAULT_TOR_CONTROL_ADDR) self.control_addr
.as_deref()
.unwrap_or(DEFAULT_TOR_CONTROL_ADDR)
} }
/// Get the control auth string. Default: "cookie". /// Get the control auth string. Default: "cookie".
@@ -479,12 +486,15 @@ impl TorConfig {
/// Get the cookie file path. Default: "/var/run/tor/control.authcookie". /// Get the cookie file path. Default: "/var/run/tor/control.authcookie".
pub fn cookie_path(&self) -> &str { pub fn cookie_path(&self) -> &str {
self.cookie_path.as_deref().unwrap_or(DEFAULT_TOR_COOKIE_PATH) self.cookie_path
.as_deref()
.unwrap_or(DEFAULT_TOR_COOKIE_PATH)
} }
/// Get the connect timeout in milliseconds. Default: 120000. /// Get the connect timeout in milliseconds. Default: 120000.
pub fn connect_timeout_ms(&self) -> u64 { pub fn connect_timeout_ms(&self) -> u64 {
self.connect_timeout_ms.unwrap_or(DEFAULT_TOR_CONNECT_TIMEOUT_MS) self.connect_timeout_ms
.unwrap_or(DEFAULT_TOR_CONNECT_TIMEOUT_MS)
} }
/// Get the default MTU. Default: 1400. /// Get the default MTU. Default: 1400.
@@ -494,7 +504,8 @@ impl TorConfig {
/// Get the max inbound connections. Default: 64. /// Get the max inbound connections. Default: 64.
pub fn max_inbound_connections(&self) -> usize { pub fn max_inbound_connections(&self) -> usize {
self.max_inbound_connections.unwrap_or(DEFAULT_TOR_MAX_INBOUND) self.max_inbound_connections
.unwrap_or(DEFAULT_TOR_MAX_INBOUND)
} }
} }
@@ -533,7 +544,10 @@ fn is_transport_empty<T>(instances: &TransportInstances<T>) -> bool {
impl TransportsConfig { impl TransportsConfig {
/// Check if any transports are configured. /// Check if any transports are configured.
pub fn is_empty(&self) -> bool { pub fn is_empty(&self) -> bool {
self.udp.is_empty() && self.ethernet.is_empty() && self.tcp.is_empty() && self.tor.is_empty() self.udp.is_empty()
&& self.ethernet.is_empty()
&& self.tcp.is_empty()
&& self.tor.is_empty()
} }
/// Merge another TransportsConfig into this one. /// Merge another TransportsConfig into this one.

View File

@@ -105,7 +105,10 @@ impl ControlSocket {
let group_name = CString::new("fips").unwrap(); let group_name = CString::new("fips").unwrap();
let grp = unsafe { libc::getgrnam(group_name.as_ptr()) }; let grp = unsafe { libc::getgrnam(group_name.as_ptr()) };
if grp.is_null() { if grp.is_null() {
debug!("'fips' group not found, skipping chown for {}", path.display()); debug!(
"'fips' group not found, skipping chown for {}",
path.display()
);
return; return;
} }
let gid = unsafe { (*grp).gr_gid }; let gid = unsafe { (*grp).gr_gid };
@@ -184,9 +187,7 @@ impl ControlSocket {
.await; .await;
let response = match read_result { let response = match read_result {
Ok(Ok(())) if line.is_empty() => { Ok(Ok(())) if line.is_empty() => Response::error("empty request"),
Response::error("empty request")
}
Ok(Ok(())) => { Ok(Ok(())) => {
// Parse the request // Parse the request
match serde_json::from_str::<Request>(line.trim()) { match serde_json::from_str::<Request>(line.trim()) {

View File

@@ -5,7 +5,7 @@
use crate::identity::encode_npub; use crate::identity::encode_npub;
use crate::node::Node; use crate::node::Node;
use serde_json::{json, Value}; use serde_json::{Value, json};
/// Helper: get current Unix time in milliseconds. /// Helper: get current Unix time in milliseconds.
fn now_ms() -> u64 { fn now_ms() -> u64 {
@@ -70,113 +70,120 @@ pub fn show_peers(node: &Node) -> Value {
let parent_id = *tree.my_declaration().parent_id(); let parent_id = *tree.my_declaration().parent_id();
let is_root = tree.is_root(); let is_root = tree.is_root();
let peers: Vec<Value> = node.peers().map(|peer| { let peers: Vec<Value> = node
let node_addr = *peer.node_addr(); .peers()
let addr_hex = hex::encode(node_addr.as_bytes()); .map(|peer| {
let node_addr = *peer.node_addr();
let addr_hex = hex::encode(node_addr.as_bytes());
// Determine tree relationship // Determine tree relationship
let is_parent = !is_root && node_addr == parent_id; let is_parent = !is_root && node_addr == parent_id;
let is_child = tree.peer_declaration(&node_addr) let is_child = tree
.is_some_and(|decl| *decl.parent_id() == my_addr); .peer_declaration(&node_addr)
.is_some_and(|decl| *decl.parent_id() == my_addr);
let mut peer_json = json!({ let mut peer_json = json!({
"node_addr": addr_hex, "node_addr": addr_hex,
"npub": peer.npub(), "npub": peer.npub(),
"display_name": node.peer_display_name(&node_addr), "display_name": node.peer_display_name(&node_addr),
"ipv6_addr": format!("{}", peer.address()), "ipv6_addr": format!("{}", peer.address()),
"connectivity": format!("{}", peer.connectivity()), "connectivity": format!("{}", peer.connectivity()),
"link_id": peer.link_id().as_u64(), "link_id": peer.link_id().as_u64(),
"authenticated_at_ms": peer.authenticated_at(), "authenticated_at_ms": peer.authenticated_at(),
"last_seen_ms": peer.last_seen(), "last_seen_ms": peer.last_seen(),
"has_tree_position": peer.has_tree_position(), "has_tree_position": peer.has_tree_position(),
"has_bloom_filter": peer.filter_sequence() > 0, "has_bloom_filter": peer.filter_sequence() > 0,
"filter_sequence": peer.filter_sequence(), "filter_sequence": peer.filter_sequence(),
"is_parent": is_parent, "is_parent": is_parent,
"is_child": is_child, "is_child": is_child,
});
// Add transport address if available
if let Some(addr) = peer.current_addr() {
peer_json["transport_addr"] = json!(format!("{}", addr));
}
// Add link info (direction, transport type)
let link_id = peer.link_id();
if let Some(link) = node.get_link(&link_id) {
peer_json["direction"] = json!(format!("{}", link.direction()));
let transport_id = link.transport_id();
if let Some(handle) = node.get_transport(&transport_id) {
peer_json["transport_type"] = json!(handle.transport_type().name);
}
}
// Add tree depth if available
if let Some(coords) = peer.coords() {
peer_json["tree_depth"] = json!(coords.depth());
}
// Add link stats
let stats = peer.link_stats();
peer_json["stats"] = json!({
"packets_sent": stats.packets_sent,
"packets_recv": stats.packets_recv,
"bytes_sent": stats.bytes_sent,
"bytes_recv": stats.bytes_recv,
});
// Add MMP metrics if available
if let Some(mmp) = peer.mmp() {
let mut mmp_json = json!({
"mode": format!("{}", mmp.mode()),
}); });
if let Some(srtt) = mmp.metrics.srtt_ms() {
mmp_json["srtt_ms"] = json!(srtt);
}
mmp_json["loss_rate"] = json!(mmp.metrics.loss_rate());
mmp_json["etx"] = json!(mmp.metrics.etx);
mmp_json["goodput_bps"] = json!(mmp.metrics.goodput_bps);
mmp_json["delivery_ratio_forward"] = json!(mmp.metrics.delivery_ratio_forward);
mmp_json["delivery_ratio_reverse"] = json!(mmp.metrics.delivery_ratio_reverse);
if let Some(smoothed_loss) = mmp.metrics.smoothed_loss() {
mmp_json["smoothed_loss"] = json!(smoothed_loss);
}
if let Some(smoothed_etx) = mmp.metrics.smoothed_etx() {
mmp_json["smoothed_etx"] = json!(smoothed_etx);
}
if let Some(srtt) = mmp.metrics.srtt_ms()
&& let Some(setx) = mmp.metrics.smoothed_etx()
{
mmp_json["lqi"] = json!(setx * (1.0 + srtt / 100.0));
}
peer_json["mmp"] = mmp_json;
}
peer_json // Add transport address if available
}).collect(); if let Some(addr) = peer.current_addr() {
peer_json["transport_addr"] = json!(format!("{}", addr));
}
// Add link info (direction, transport type)
let link_id = peer.link_id();
if let Some(link) = node.get_link(&link_id) {
peer_json["direction"] = json!(format!("{}", link.direction()));
let transport_id = link.transport_id();
if let Some(handle) = node.get_transport(&transport_id) {
peer_json["transport_type"] = json!(handle.transport_type().name);
}
}
// Add tree depth if available
if let Some(coords) = peer.coords() {
peer_json["tree_depth"] = json!(coords.depth());
}
// Add link stats
let stats = peer.link_stats();
peer_json["stats"] = json!({
"packets_sent": stats.packets_sent,
"packets_recv": stats.packets_recv,
"bytes_sent": stats.bytes_sent,
"bytes_recv": stats.bytes_recv,
});
// Add MMP metrics if available
if let Some(mmp) = peer.mmp() {
let mut mmp_json = json!({
"mode": format!("{}", mmp.mode()),
});
if let Some(srtt) = mmp.metrics.srtt_ms() {
mmp_json["srtt_ms"] = json!(srtt);
}
mmp_json["loss_rate"] = json!(mmp.metrics.loss_rate());
mmp_json["etx"] = json!(mmp.metrics.etx);
mmp_json["goodput_bps"] = json!(mmp.metrics.goodput_bps);
mmp_json["delivery_ratio_forward"] = json!(mmp.metrics.delivery_ratio_forward);
mmp_json["delivery_ratio_reverse"] = json!(mmp.metrics.delivery_ratio_reverse);
if let Some(smoothed_loss) = mmp.metrics.smoothed_loss() {
mmp_json["smoothed_loss"] = json!(smoothed_loss);
}
if let Some(smoothed_etx) = mmp.metrics.smoothed_etx() {
mmp_json["smoothed_etx"] = json!(smoothed_etx);
}
if let Some(srtt) = mmp.metrics.srtt_ms()
&& let Some(setx) = mmp.metrics.smoothed_etx()
{
mmp_json["lqi"] = json!(setx * (1.0 + srtt / 100.0));
}
peer_json["mmp"] = mmp_json;
}
peer_json
})
.collect();
json!({ "peers": peers }) json!({ "peers": peers })
} }
/// `show_links` — Active links. /// `show_links` — Active links.
pub fn show_links(node: &Node) -> Value { pub fn show_links(node: &Node) -> Value {
let links: Vec<Value> = node.links().map(|link| { let links: Vec<Value> = node
let stats = link.stats(); .links()
json!({ .map(|link| {
"link_id": link.link_id().as_u64(), let stats = link.stats();
"transport_id": link.transport_id().as_u32(), json!({
"remote_addr": format!("{}", link.remote_addr()), "link_id": link.link_id().as_u64(),
"direction": format!("{}", link.direction()), "transport_id": link.transport_id().as_u32(),
"state": format!("{}", link.state()), "remote_addr": format!("{}", link.remote_addr()),
"created_at_ms": link.created_at(), "direction": format!("{}", link.direction()),
"stats": { "state": format!("{}", link.state()),
"packets_sent": stats.packets_sent, "created_at_ms": link.created_at(),
"packets_recv": stats.packets_recv, "stats": {
"bytes_sent": stats.bytes_sent, "packets_sent": stats.packets_sent,
"bytes_recv": stats.bytes_recv, "packets_recv": stats.packets_recv,
"last_recv_ms": stats.last_recv_ms, "bytes_sent": stats.bytes_sent,
}, "bytes_recv": stats.bytes_recv,
"last_recv_ms": stats.last_recv_ms,
},
})
}) })
}).collect(); .collect();
json!({ "links": links }) json!({ "links": links })
} }
@@ -188,29 +195,34 @@ pub fn show_tree(node: &Node) -> Value {
let decl = tree.my_declaration(); let decl = tree.my_declaration();
// Build coords array as hex strings // Build coords array as hex strings
let coords: Vec<String> = my_coords.entries() let coords: Vec<String> = my_coords
.entries()
.iter() .iter()
.map(|e| hex::encode(e.node_addr.as_bytes())) .map(|e| hex::encode(e.node_addr.as_bytes()))
.collect(); .collect();
// Build peer tree data // Build peer tree data
let peers: Vec<Value> = tree.peer_ids().map(|peer_id| { let peers: Vec<Value> = tree
let mut peer_json = json!({ .peer_ids()
"node_addr": hex::encode(peer_id.as_bytes()), .map(|peer_id| {
"display_name": node.peer_display_name(peer_id), let mut peer_json = json!({
}); "node_addr": hex::encode(peer_id.as_bytes()),
if let Some(coords) = tree.peer_coords(peer_id) { "display_name": node.peer_display_name(peer_id),
let coord_path: Vec<String> = coords.entries() });
.iter() if let Some(coords) = tree.peer_coords(peer_id) {
.map(|e| hex::encode(e.node_addr.as_bytes())) let coord_path: Vec<String> = coords
.collect(); .entries()
peer_json["depth"] = json!(coords.depth()); .iter()
peer_json["root"] = json!(hex::encode(coords.root_id().as_bytes())); .map(|e| hex::encode(e.node_addr.as_bytes()))
peer_json["coords"] = json!(coord_path); .collect();
peer_json["distance_to_us"] = json!(my_coords.distance_to(coords)); peer_json["depth"] = json!(coords.depth());
} peer_json["root"] = json!(hex::encode(coords.root_id().as_bytes()));
peer_json peer_json["coords"] = json!(coord_path);
}).collect(); peer_json["distance_to_us"] = json!(my_coords.distance_to(coords));
}
peer_json
})
.collect();
// Determine parent display name // Determine parent display name
let parent_addr = my_coords.parent_id(); let parent_addr = my_coords.parent_id();
@@ -237,68 +249,71 @@ pub fn show_tree(node: &Node) -> Value {
/// `show_sessions` — End-to-end sessions. /// `show_sessions` — End-to-end sessions.
pub fn show_sessions(node: &Node) -> Value { pub fn show_sessions(node: &Node) -> Value {
let sessions: Vec<Value> = node.session_entries().map(|(addr, entry)| { let sessions: Vec<Value> = node
let state_str = if entry.is_established() { .session_entries()
"established" .map(|(addr, entry)| {
} else if entry.is_initiating() { let state_str = if entry.is_established() {
"initiating" "established"
} else if entry.is_awaiting_msg3() { } else if entry.is_initiating() {
"awaiting_msg3" "initiating"
} else { } else if entry.is_awaiting_msg3() {
"unknown" "awaiting_msg3"
}; } else {
"unknown"
};
let mut session_json = json!({ let mut session_json = json!({
"remote_addr": hex::encode(addr.as_bytes()), "remote_addr": hex::encode(addr.as_bytes()),
"display_name": node.peer_display_name(addr), "display_name": node.peer_display_name(addr),
"state": state_str, "state": state_str,
"is_initiator": entry.is_initiator(), "is_initiator": entry.is_initiator(),
"last_activity_ms": entry.last_activity(), "last_activity_ms": entry.last_activity(),
});
// Derive npub from session's remote public key
let (xonly, _parity) = entry.remote_pubkey().x_only_public_key();
session_json["npub"] = json!(encode_npub(&xonly));
// Traffic counters
let (pkts_tx, pkts_rx, bytes_tx, bytes_rx) = entry.traffic_counters();
session_json["stats"] = json!({
"packets_sent": pkts_tx,
"packets_recv": pkts_rx,
"bytes_sent": bytes_tx,
"bytes_recv": bytes_rx,
});
// Add session MMP if available
if let Some(mmp) = entry.mmp() {
let mut mmp_json = json!({
"mode": format!("{}", mmp.mode()),
"loss_rate": mmp.metrics.loss_rate(),
"etx": mmp.metrics.etx,
"goodput_bps": mmp.metrics.goodput_bps,
"delivery_ratio_forward": mmp.metrics.delivery_ratio_forward,
"delivery_ratio_reverse": mmp.metrics.delivery_ratio_reverse,
"path_mtu": mmp.path_mtu.current_mtu(),
}); });
if let Some(srtt) = mmp.metrics.srtt_ms() {
mmp_json["srtt_ms"] = json!(srtt);
}
if let Some(smoothed_loss) = mmp.metrics.smoothed_loss() {
mmp_json["smoothed_loss"] = json!(smoothed_loss);
}
if let Some(smoothed_etx) = mmp.metrics.smoothed_etx() {
mmp_json["smoothed_etx"] = json!(smoothed_etx);
}
if let Some(srtt) = mmp.metrics.srtt_ms()
&& let Some(setx) = mmp.metrics.smoothed_etx()
{
mmp_json["sqi"] = json!(setx * (1.0 + srtt / 100.0));
}
session_json["mmp"] = mmp_json;
}
session_json // Derive npub from session's remote public key
}).collect(); let (xonly, _parity) = entry.remote_pubkey().x_only_public_key();
session_json["npub"] = json!(encode_npub(&xonly));
// Traffic counters
let (pkts_tx, pkts_rx, bytes_tx, bytes_rx) = entry.traffic_counters();
session_json["stats"] = json!({
"packets_sent": pkts_tx,
"packets_recv": pkts_rx,
"bytes_sent": bytes_tx,
"bytes_recv": bytes_rx,
});
// Add session MMP if available
if let Some(mmp) = entry.mmp() {
let mut mmp_json = json!({
"mode": format!("{}", mmp.mode()),
"loss_rate": mmp.metrics.loss_rate(),
"etx": mmp.metrics.etx,
"goodput_bps": mmp.metrics.goodput_bps,
"delivery_ratio_forward": mmp.metrics.delivery_ratio_forward,
"delivery_ratio_reverse": mmp.metrics.delivery_ratio_reverse,
"path_mtu": mmp.path_mtu.current_mtu(),
});
if let Some(srtt) = mmp.metrics.srtt_ms() {
mmp_json["srtt_ms"] = json!(srtt);
}
if let Some(smoothed_loss) = mmp.metrics.smoothed_loss() {
mmp_json["smoothed_loss"] = json!(smoothed_loss);
}
if let Some(smoothed_etx) = mmp.metrics.smoothed_etx() {
mmp_json["smoothed_etx"] = json!(smoothed_etx);
}
if let Some(srtt) = mmp.metrics.srtt_ms()
&& let Some(setx) = mmp.metrics.smoothed_etx()
{
mmp_json["sqi"] = json!(setx * (1.0 + srtt / 100.0));
}
session_json["mmp"] = mmp_json;
}
session_json
})
.collect();
json!({ "sessions": sessions }) json!({ "sessions": sessions })
} }
@@ -307,27 +322,31 @@ pub fn show_sessions(node: &Node) -> Value {
pub fn show_bloom(node: &Node) -> Value { pub fn show_bloom(node: &Node) -> Value {
let bloom = node.bloom_state(); let bloom = node.bloom_state();
let leaf_deps: Vec<String> = bloom.leaf_dependents() let leaf_deps: Vec<String> = bloom
.leaf_dependents()
.iter() .iter()
.map(|addr| hex::encode(addr.as_bytes())) .map(|addr| hex::encode(addr.as_bytes()))
.collect(); .collect();
// Build per-peer filter info // Build per-peer filter info
let peer_filters: Vec<Value> = node.peers().map(|peer| { let peer_filters: Vec<Value> = node
let addr = *peer.node_addr(); .peers()
let mut pf = json!({ .map(|peer| {
"peer": hex::encode(addr.as_bytes()), let addr = *peer.node_addr();
"display_name": node.peer_display_name(&addr), let mut pf = json!({
"has_filter": peer.filter_sequence() > 0, "peer": hex::encode(addr.as_bytes()),
"filter_sequence": peer.filter_sequence(), "display_name": node.peer_display_name(&addr),
}); "has_filter": peer.filter_sequence() > 0,
if let Some(filter) = peer.inbound_filter() { "filter_sequence": peer.filter_sequence(),
pf["estimated_count"] = json!(filter.estimated_count()); });
pf["set_bits"] = json!(filter.count_ones()); if let Some(filter) = peer.inbound_filter() {
pf["fill_ratio"] = json!(filter.fill_ratio()); pf["estimated_count"] = json!(filter.estimated_count());
} pf["set_bits"] = json!(filter.count_ones());
pf pf["fill_ratio"] = json!(filter.fill_ratio());
}).collect(); }
pf
})
.collect();
let bloom_stats = node.stats().snapshot().bloom; let bloom_stats = node.stats().snapshot().bloom;
@@ -397,36 +416,39 @@ pub fn show_mmp(node: &Node) -> Value {
}).collect(); }).collect();
// Session-layer MMP // Session-layer MMP
let sessions: Vec<Value> = node.session_entries().filter_map(|(addr, entry)| { let sessions: Vec<Value> = node
let mmp = entry.mmp()?; .session_entries()
let metrics = &mmp.metrics; .filter_map(|(addr, entry)| {
let mmp = entry.mmp()?;
let metrics = &mmp.metrics;
let mut session_layer = json!({ let mut session_layer = json!({
"loss_rate": metrics.loss_rate(), "loss_rate": metrics.loss_rate(),
"etx": metrics.etx, "etx": metrics.etx,
"path_mtu": mmp.path_mtu.current_mtu(), "path_mtu": mmp.path_mtu.current_mtu(),
}); });
if let Some(smoothed_loss) = metrics.smoothed_loss() { if let Some(smoothed_loss) = metrics.smoothed_loss() {
session_layer["smoothed_loss"] = json!(smoothed_loss); session_layer["smoothed_loss"] = json!(smoothed_loss);
} }
if let Some(smoothed_etx) = metrics.smoothed_etx() { if let Some(smoothed_etx) = metrics.smoothed_etx() {
session_layer["smoothed_etx"] = json!(smoothed_etx); session_layer["smoothed_etx"] = json!(smoothed_etx);
} }
if let Some(srtt) = metrics.srtt_ms() { if let Some(srtt) = metrics.srtt_ms() {
session_layer["srtt_ms"] = json!(srtt); session_layer["srtt_ms"] = json!(srtt);
if let Some(setx) = metrics.smoothed_etx() { if let Some(setx) = metrics.smoothed_etx() {
session_layer["sqi"] = json!(setx * (1.0 + srtt / 100.0)); session_layer["sqi"] = json!(setx * (1.0 + srtt / 100.0));
}
} }
}
Some(json!({ Some(json!({
"remote": hex::encode(addr.as_bytes()), "remote": hex::encode(addr.as_bytes()),
"display_name": node.peer_display_name(addr), "display_name": node.peer_display_name(addr),
"mode": format!("{}", mmp.mode()), "mode": format!("{}", mmp.mode()),
"session_layer": session_layer, "session_layer": session_layer,
})) }))
}).collect(); })
.collect();
json!({ json!({
"peers": peers, "peers": peers,
@@ -452,60 +474,65 @@ pub fn show_cache(node: &Node) -> Value {
/// `show_connections` — Pending handshakes. /// `show_connections` — Pending handshakes.
pub fn show_connections(node: &Node) -> Value { pub fn show_connections(node: &Node) -> Value {
let now = now_ms(); let now = now_ms();
let connections: Vec<Value> = node.connections().map(|conn| { let connections: Vec<Value> = node
let mut conn_json = json!({ .connections()
"link_id": conn.link_id().as_u64(), .map(|conn| {
"direction": format!("{}", conn.direction()), let mut conn_json = json!({
"handshake_state": format!("{}", conn.handshake_state()), "link_id": conn.link_id().as_u64(),
"started_at_ms": conn.started_at(), "direction": format!("{}", conn.direction()),
"idle_ms": now.saturating_sub(conn.last_activity()), "handshake_state": format!("{}", conn.handshake_state()),
"resend_count": conn.resend_count(), "started_at_ms": conn.started_at(),
}); "idle_ms": now.saturating_sub(conn.last_activity()),
"resend_count": conn.resend_count(),
});
if let Some(identity) = conn.expected_identity() { if let Some(identity) = conn.expected_identity() {
conn_json["expected_peer"] = json!(identity.npub()); conn_json["expected_peer"] = json!(identity.npub());
} }
conn_json conn_json
}).collect(); })
.collect();
json!({ "connections": connections }) json!({ "connections": connections })
} }
/// `show_transports` — Transport instances. /// `show_transports` — Transport instances.
pub fn show_transports(node: &Node) -> Value { pub fn show_transports(node: &Node) -> Value {
let transports: Vec<Value> = node.transport_ids().map(|id| { let transports: Vec<Value> = node
let handle = node.get_transport(id).unwrap(); .transport_ids()
let mut t_json = json!({ .map(|id| {
"transport_id": id.as_u32(), let handle = node.get_transport(id).unwrap();
"type": handle.transport_type().name, let mut t_json = json!({
"state": format!("{}", handle.state()), "transport_id": id.as_u32(),
"mtu": handle.mtu(), "type": handle.transport_type().name,
}); "state": format!("{}", handle.state()),
"mtu": handle.mtu(),
});
if let Some(name) = handle.name() { if let Some(name) = handle.name() {
t_json["name"] = json!(name); t_json["name"] = json!(name);
} }
if let Some(addr) = handle.local_addr() { if let Some(addr) = handle.local_addr() {
t_json["local_addr"] = json!(format!("{}", addr)); t_json["local_addr"] = json!(format!("{}", addr));
} }
// Tor-specific fields // Tor-specific fields
if let Some(mode) = handle.tor_mode() { if let Some(mode) = handle.tor_mode() {
t_json["tor_mode"] = json!(mode); t_json["tor_mode"] = json!(mode);
} }
if let Some(onion) = handle.onion_address() { if let Some(onion) = handle.onion_address() {
t_json["onion_address"] = json!(onion); t_json["onion_address"] = json!(onion);
} }
if let Some(monitoring) = handle.tor_monitoring() { if let Some(monitoring) = handle.tor_monitoring() {
t_json["tor_monitoring"] = t_json["tor_monitoring"] = serde_json::to_value(&monitoring).unwrap_or_default();
serde_json::to_value(&monitoring).unwrap_or_default(); }
}
t_json["stats"] = handle.transport_stats(); t_json["stats"] = handle.transport_stats();
t_json t_json
}).collect(); })
.collect();
json!({ "transports": transports }) json!({ "transports": transports })
} }

View File

@@ -3,7 +3,7 @@
use std::fmt; use std::fmt;
use std::net::Ipv6Addr; use std::net::Ipv6Addr;
use super::{IdentityError, NodeAddr, FIPS_ADDRESS_PREFIX}; use super::{FIPS_ADDRESS_PREFIX, IdentityError, NodeAddr};
/// 128-bit FIPS address with IPv6-compatible format. /// 128-bit FIPS address with IPv6-compatible format.
/// ///

View File

@@ -3,9 +3,9 @@
use secp256k1::{Keypair, PublicKey, Secp256k1, SecretKey, XOnlyPublicKey}; use secp256k1::{Keypair, PublicKey, Secp256k1, SecretKey, XOnlyPublicKey};
use std::fmt; use std::fmt;
use super::auth::{auth_challenge_digest, AuthResponse}; use super::auth::{AuthResponse, auth_challenge_digest};
use super::encoding::{decode_secret, encode_npub}; use super::encoding::{decode_secret, encode_npub};
use super::{sha256, FipsAddress, IdentityError, NodeAddr}; use super::{FipsAddress, IdentityError, NodeAddr, sha256};
/// A FIPS node identity consisting of a keypair and derived identifiers. /// A FIPS node identity consisting of a keypair and derived identifiers.
/// ///
@@ -22,8 +22,8 @@ impl Identity {
pub fn generate() -> Self { pub fn generate() -> Self {
let mut secret_bytes = [0u8; 32]; let mut secret_bytes = [0u8; 32];
rand::Rng::fill_bytes(&mut rand::rng(), &mut secret_bytes); rand::Rng::fill_bytes(&mut rand::rng(), &mut secret_bytes);
let secret_key = SecretKey::from_slice(&secret_bytes) let secret_key =
.expect("32 random bytes is a valid secret key"); SecretKey::from_slice(&secret_bytes).expect("32 random bytes is a valid secret key");
Self::from_secret_key(secret_key) Self::from_secret_key(secret_key)
} }

View File

@@ -4,7 +4,7 @@ use secp256k1::XOnlyPublicKey;
use sha2::{Digest, Sha256}; use sha2::{Digest, Sha256};
use std::fmt; use std::fmt;
use super::{hex_encode, IdentityError}; use super::{IdentityError, hex_encode};
/// 16-byte node identifier derived from truncated SHA-256(pubkey). /// 16-byte node identifier derived from truncated SHA-256(pubkey).
/// ///

View File

@@ -4,7 +4,7 @@ use secp256k1::{Parity, PublicKey, Secp256k1, XOnlyPublicKey};
use std::fmt; use std::fmt;
use super::encoding::{decode_npub, encode_npub}; use super::encoding::{decode_npub, encode_npub};
use super::{sha256, FipsAddress, IdentityError, NodeAddr}; use super::{FipsAddress, IdentityError, NodeAddr, sha256};
/// A known peer's identity (public key only, no signing capability). /// A known peer's identity (public key only, no signing capability).
/// ///
@@ -99,7 +99,8 @@ impl PeerIdentity {
pub fn verify(&self, data: &[u8], signature: &secp256k1::schnorr::Signature) -> bool { pub fn verify(&self, data: &[u8], signature: &secp256k1::schnorr::Signature) -> bool {
let secp = Secp256k1::new(); let secp = Secp256k1::new();
let digest = sha256(data); let digest = sha256(data);
secp.verify_schnorr(signature, &digest, &self.pubkey).is_ok() secp.verify_schnorr(signature, &digest, &self.pubkey)
.is_ok()
} }
} }

View File

@@ -110,9 +110,9 @@ fn test_node_addr_ordering() {
fn test_identity_from_secret_bytes() { fn test_identity_from_secret_bytes() {
// A known secret key (32 bytes) // A known secret key (32 bytes)
let secret_bytes: [u8; 32] = [ let secret_bytes: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e,
0x1d, 0x1e, 0x1f, 0x20, 0x1f, 0x20,
]; ];
let identity1 = Identity::from_secret_bytes(&secret_bytes).unwrap(); let identity1 = Identity::from_secret_bytes(&secret_bytes).unwrap();
@@ -163,9 +163,10 @@ fn test_identity_sign() {
// Verify the signature manually // Verify the signature manually
let secp = secp256k1::Secp256k1::new(); let secp = secp256k1::Secp256k1::new();
let digest = super::sha256(data); let digest = super::sha256(data);
assert!(secp assert!(
.verify_schnorr(&sig, &digest, &identity.pubkey()) secp.verify_schnorr(&sig, &digest, &identity.pubkey())
.is_ok()); .is_ok()
);
} }
#[test] #[test]
@@ -193,9 +194,9 @@ fn test_npub_roundtrip() {
fn test_npub_known_vector() { fn test_npub_known_vector() {
// Test against a known npub (from NIP-19 test vectors or generated externally) // Test against a known npub (from NIP-19 test vectors or generated externally)
let secret_bytes: [u8; 32] = [ let secret_bytes: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e,
0x1d, 0x1e, 0x1f, 0x20, 0x1f, 0x20,
]; ];
let identity = Identity::from_secret_bytes(&secret_bytes).unwrap(); let identity = Identity::from_secret_bytes(&secret_bytes).unwrap();
@@ -274,9 +275,9 @@ fn test_peer_identity_display() {
#[test] #[test]
fn test_nsec_roundtrip() { fn test_nsec_roundtrip() {
let secret_bytes: [u8; 32] = [ let secret_bytes: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e,
0x1d, 0x1e, 0x1f, 0x20, 0x1f, 0x20,
]; ];
let secret_key = SecretKey::from_slice(&secret_bytes).unwrap(); let secret_key = SecretKey::from_slice(&secret_bytes).unwrap();
@@ -301,9 +302,9 @@ fn test_decode_nsec_invalid_prefix() {
#[test] #[test]
fn test_decode_secret_nsec() { fn test_decode_secret_nsec() {
let secret_bytes: [u8; 32] = [ let secret_bytes: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e,
0x1d, 0x1e, 0x1f, 0x20, 0x1f, 0x20,
]; ];
let secret_key = SecretKey::from_slice(&secret_bytes).unwrap(); let secret_key = SecretKey::from_slice(&secret_bytes).unwrap();
@@ -319,9 +320,9 @@ fn test_decode_secret_hex() {
let decoded = decode_secret(hex_str).unwrap(); let decoded = decode_secret(hex_str).unwrap();
let expected: [u8; 32] = [ let expected: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e,
0x1d, 0x1e, 0x1f, 0x20, 0x1f, 0x20,
]; ];
assert_eq!(decoded.secret_bytes(), expected); assert_eq!(decoded.secret_bytes(), expected);
} }
@@ -329,9 +330,9 @@ fn test_decode_secret_hex() {
#[test] #[test]
fn test_identity_from_secret_str_nsec() { fn test_identity_from_secret_str_nsec() {
let secret_bytes: [u8; 32] = [ let secret_bytes: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e,
0x1d, 0x1e, 0x1f, 0x20, 0x1f, 0x20,
]; ];
let secret_key = SecretKey::from_slice(&secret_bytes).unwrap(); let secret_key = SecretKey::from_slice(&secret_bytes).unwrap();
@@ -347,9 +348,9 @@ fn test_identity_from_secret_str_nsec() {
fn test_identity_from_secret_str_hex() { fn test_identity_from_secret_str_hex() {
let hex_str = "0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20"; let hex_str = "0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20";
let secret_bytes: [u8; 32] = [ let secret_bytes: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e,
0x1d, 0x1e, 0x1f, 0x20, 0x1f, 0x20,
]; ];
let identity = Identity::from_secret_str(hex_str).unwrap(); let identity = Identity::from_secret_str(hex_str).unwrap();
@@ -379,11 +380,8 @@ fn test_hex_conversion_case2() {
#[test] #[test]
fn test_decode_npub_invalid_length() { fn test_decode_npub_invalid_length() {
// Encode 16 bytes (too short) as bech32 with npub prefix // Encode 16 bytes (too short) as bech32 with npub prefix
let short = bech32::encode::<bech32::Bech32>( let short =
bech32::Hrp::parse_unchecked("npub"), bech32::encode::<bech32::Bech32>(bech32::Hrp::parse_unchecked("npub"), &[0u8; 16]).unwrap();
&[0u8; 16],
)
.unwrap();
let result = decode_npub(&short); let result = decode_npub(&short);
assert!(matches!(result, Err(IdentityError::InvalidNpubLength(16)))); assert!(matches!(result, Err(IdentityError::InvalidNpubLength(16))));
} }
@@ -391,11 +389,8 @@ fn test_decode_npub_invalid_length() {
#[test] #[test]
fn test_decode_nsec_invalid_length() { fn test_decode_nsec_invalid_length() {
// Encode 16 bytes (too short) as bech32 with nsec prefix // Encode 16 bytes (too short) as bech32 with nsec prefix
let short = bech32::encode::<bech32::Bech32>( let short =
bech32::Hrp::parse_unchecked("nsec"), bech32::encode::<bech32::Bech32>(bech32::Hrp::parse_unchecked("nsec"), &[0u8; 16]).unwrap();
&[0u8; 16],
)
.unwrap();
let result = decode_nsec(&short); let result = decode_nsec(&short);
assert!(matches!(result, Err(IdentityError::InvalidNsecLength(16)))); assert!(matches!(result, Err(IdentityError::InvalidNsecLength(16))));
} }
@@ -418,8 +413,8 @@ fn test_decode_secret_hex_invalid_chars() {
#[test] #[test]
fn test_node_addr_debug() { fn test_node_addr_debug() {
let bytes = [ let bytes = [
0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
0x00, 0x00, 0x00,
]; ];
let node_addr = NodeAddr::from_bytes(bytes); let node_addr = NodeAddr::from_bytes(bytes);
let debug = format!("{:?}", node_addr); let debug = format!("{:?}", node_addr);
@@ -429,8 +424,8 @@ fn test_node_addr_debug() {
#[test] #[test]
fn test_node_addr_display() { fn test_node_addr_display() {
let bytes = [ let bytes = [
0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0xfe, 0xdc, 0xba, 0x98, 0x76, 0x54, 0x01, 0x23, 0x45, 0x67, 0x89, 0xab, 0xcd, 0xef, 0xfe, 0xdc, 0xba, 0x98, 0x76, 0x54, 0x32,
0x32, 0x10, 0x10,
]; ];
let node_addr = NodeAddr::from_bytes(bytes); let node_addr = NodeAddr::from_bytes(bytes);
let display = format!("{}", node_addr); let display = format!("{}", node_addr);
@@ -587,9 +582,9 @@ fn test_peer_identity_pubkey_full_preserved_parity() {
// Create two identities and find one with odd parity to make this test meaningful // Create two identities and find one with odd parity to make this test meaningful
let secp = Secp256k1::new(); let secp = Secp256k1::new();
let secret_bytes: [u8; 32] = [ let secret_bytes: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f,
0x0f, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e,
0x1d, 0x1e, 0x1f, 0x20, 0x1f, 0x20,
]; ];
let keypair = Keypair::from_seckey_slice(&secp, &secret_bytes).unwrap(); let keypair = Keypair::from_seckey_slice(&secp, &secret_bytes).unwrap();
let full_pubkey = keypair.public_key(); let full_pubkey = keypair.public_key();

View File

@@ -3,26 +3,26 @@
//! A distributed, decentralized network routing protocol for mesh nodes //! A distributed, decentralized network routing protocol for mesh nodes
//! connecting over arbitrary transports. //! connecting over arbitrary transports.
pub mod version;
pub mod bloom; pub mod bloom;
pub mod cache; pub mod cache;
pub mod config; pub mod config;
pub mod control; pub mod control;
pub mod identity; pub mod identity;
pub mod mmp; pub mod mmp;
pub mod noise;
pub mod utils;
pub mod node; pub mod node;
pub mod noise;
pub mod peer; pub mod peer;
pub mod protocol; pub mod protocol;
pub mod transport; pub mod transport;
pub mod tree; pub mod tree;
pub mod upper; pub mod upper;
pub mod utils;
pub mod version;
// Re-export identity types // Re-export identity types
pub use identity::{ pub use identity::{
decode_npub, decode_nsec, decode_secret, encode_npub, encode_nsec, AuthChallenge, AuthResponse, AuthChallenge, AuthResponse, FipsAddress, Identity, IdentityError, NodeAddr, PeerIdentity,
FipsAddress, Identity, IdentityError, NodeAddr, PeerIdentity, decode_npub, decode_nsec, decode_secret, encode_npub, encode_nsec,
}; };
// Re-export config types // Re-export config types
@@ -36,18 +36,18 @@ pub use tree::{CoordEntry, ParentDeclaration, TreeCoordinate, TreeError, TreeSta
pub use bloom::{BloomError, BloomFilter, BloomState}; pub use bloom::{BloomError, BloomFilter, BloomState};
// Re-export transport types // Re-export transport types
pub use transport::{
packet_channel, DiscoveredPeer, Link, LinkDirection, LinkId, LinkState, LinkStats, PacketRx,
PacketTx, ReceivedPacket, Transport, TransportAddr, TransportError, TransportHandle,
TransportId, TransportState, TransportType,
};
pub use transport::udp::UdpTransport; pub use transport::udp::UdpTransport;
pub use transport::{
DiscoveredPeer, Link, LinkDirection, LinkId, LinkState, LinkStats, PacketRx, PacketTx,
ReceivedPacket, Transport, TransportAddr, TransportError, TransportHandle, TransportId,
TransportState, TransportType, packet_channel,
};
// Re-export protocol types // Re-export protocol types
pub use protocol::{ pub use protocol::{
CoordsRequired, FilterAnnounce, HandshakeMessageType, LinkMessageType, CoordsRequired, FilterAnnounce, HandshakeMessageType, LinkMessageType, LookupRequest,
LookupRequest, LookupResponse, PathBroken, ProtocolError, SessionAck, SessionDatagram, LookupResponse, PathBroken, ProtocolError, SessionAck, SessionDatagram, SessionFlags,
SessionFlags, SessionMessageType, SessionSetup, TreeAnnounce, SessionMessageType, SessionSetup, TreeAnnounce,
}; };
// Re-export cache types // Re-export cache types
@@ -55,10 +55,9 @@ pub use cache::{CacheEntry, CacheError, CacheStats, CoordCache};
// Re-export peer types // Re-export peer types
pub use peer::{ pub use peer::{
cross_connection_winner, ActivePeer, ConnectivityState, HandshakeState, PeerConnection, ActivePeer, ConnectivityState, HandshakeState, PeerConnection, PeerError, PeerSlot,
PeerError, PeerSlot, PromotionResult, PromotionResult, cross_connection_winner,
}; };
// Re-export node types // Re-export node types
pub use node::{Node, NodeError, NodeState}; pub use node::{Node, NodeError, NodeState};

View File

@@ -371,7 +371,10 @@ mod tests {
} }
// Should converge near 1000µs // Should converge near 1000µs
let jitter = j.jitter_us(); let jitter = j.jitter_us();
assert!(jitter > 900 && jitter < 1100, "jitter={jitter}, expected ~1000"); assert!(
jitter > 900 && jitter < 1100,
"jitter={jitter}, expected ~1000"
);
} }
#[test] #[test]
@@ -391,10 +394,7 @@ mod tests {
s.update(50_000); s.update(50_000);
} }
let srtt = s.srtt_us(); let srtt = s.srtt_us();
assert!( assert!((srtt - 50_000).abs() < 1000, "srtt={srtt}, expected ~50000");
(srtt - 50_000).abs() < 1000,
"srtt={srtt}, expected ~50000"
);
} }
#[test] #[test]
@@ -417,7 +417,12 @@ mod tests {
e.update(100.0); e.update(100.0);
} }
// Short should be closer to 100 than long // Short should be closer to 100 than long
assert!(e.short() > e.long(), "short={} long={}", e.short(), e.long()); assert!(
e.short() > e.long(),
"short={} long={}",
e.short(),
e.long()
);
} }
#[test] #[test]
@@ -437,7 +442,10 @@ mod tests {
d.push(i, 5000 + (i as i64) * 100); // increasing by 100µs per packet d.push(i, 5000 + (i as i64) * 100); // increasing by 100µs per packet
} }
let trend = d.trend_us_per_sec(); let trend = d.trend_us_per_sec();
assert!(trend > 0, "increasing OWD should have positive trend, got {trend}"); assert!(
trend > 0,
"increasing OWD should have positive trend, got {trend}"
);
} }
#[test] #[test]

View File

@@ -110,7 +110,12 @@ impl MmpMetrics {
/// ///
/// Returns `true` if this report produced the first SRTT measurement /// Returns `true` if this report produced the first SRTT measurement
/// (transition from uninitialized to initialized). /// (transition from uninitialized to initialized).
pub fn process_receiver_report(&mut self, rr: &ReceiverReport, our_timestamp_ms: u32, now: Instant) -> bool { pub fn process_receiver_report(
&mut self,
rr: &ReceiverReport,
our_timestamp_ms: u32,
now: Instant,
) -> bool {
let had_srtt = self.srtt.initialized(); let had_srtt = self.srtt.initialized();
// --- RTT from timestamp echo --- // --- RTT from timestamp echo ---
@@ -138,8 +143,12 @@ impl MmpMetrics {
// --- Loss rate from cumulative counters --- // --- Loss rate from cumulative counters ---
// Delta: frames the peer should have received vs. actually received // Delta: frames the peer should have received vs. actually received
if self.has_prev_rr { if self.has_prev_rr {
let counter_span = rr.highest_counter.saturating_sub(self.prev_rr_highest_counter); let counter_span = rr
let packets_delta = rr.cumulative_packets_recv.saturating_sub(self.prev_rr_cum_packets); .highest_counter
.saturating_sub(self.prev_rr_highest_counter);
let packets_delta = rr
.cumulative_packets_recv
.saturating_sub(self.prev_rr_cum_packets);
if counter_span > 0 { if counter_span > 0 {
let delivery = (packets_delta as f64) / (counter_span as f64); let delivery = (packets_delta as f64) / (counter_span as f64);
@@ -153,7 +162,9 @@ impl MmpMetrics {
// --- Goodput from cumulative bytes + time delta --- // --- Goodput from cumulative bytes + time delta ---
if self.has_prev_rr { if self.has_prev_rr {
let bytes_delta = rr.cumulative_bytes_recv.saturating_sub(self.prev_rr_cum_bytes); let bytes_delta = rr
.cumulative_bytes_recv
.saturating_sub(self.prev_rr_cum_bytes);
self.goodput_trend.update(bytes_delta as f64); self.goodput_trend.update(bytes_delta as f64);
// Compute bytes/sec if we have a time reference // Compute bytes/sec if we have a time reference
@@ -379,8 +390,16 @@ mod tests {
// Second report 1s later: 150KB total (100KB delta in 1s = 100KB/s) // Second report 1s later: 150KB total (100KB delta in 1s = 100KB/s)
let rr2 = make_rr(300, 290, 150_000, 0, 0, 0); let rr2 = make_rr(300, 290, 150_000, 0, 0, 0);
m.process_receiver_report(&rr2, 0, t0 + Duration::from_secs(1)); m.process_receiver_report(&rr2, 0, t0 + Duration::from_secs(1));
assert!(m.goodput_bps() > 90_000.0, "goodput={}, expected ~100000", m.goodput_bps()); assert!(
assert!(m.goodput_bps() < 110_000.0, "goodput={}, expected ~100000", m.goodput_bps()); m.goodput_bps() > 90_000.0,
"goodput={}, expected ~100000",
m.goodput_bps()
);
assert!(
m.goodput_bps() < 110_000.0,
"goodput={}, expected ~100000",
m.goodput_bps()
);
} }
#[test] #[test]
@@ -397,8 +416,11 @@ mod tests {
// Third call: 50% loss (100 frames sent, 50 received) // Third call: 50% loss (100 frames sent, 50 received)
m.update_reverse_delivery(350, 400); m.update_reverse_delivery(350, 400);
assert!((m.delivery_ratio_reverse - 0.5).abs() < 0.001, assert!(
"reverse={}, expected 0.5", m.delivery_ratio_reverse); (m.delivery_ratio_reverse - 0.5).abs() < 0.001,
"reverse={}, expected 0.5",
m.delivery_ratio_reverse
);
} }
#[test] #[test]
@@ -422,7 +444,10 @@ mod tests {
// Second call after rekey: 80% delivery // Second call after rekey: 80% delivery
m.update_reverse_delivery(90, 100); m.update_reverse_delivery(90, 100);
assert!((m.delivery_ratio_reverse - 0.8).abs() < 0.001, assert!(
"reverse={}, expected 0.8", m.delivery_ratio_reverse); (m.delivery_ratio_reverse - 0.8).abs() < 0.001,
"reverse={}, expected 0.8",
m.delivery_ratio_reverse
);
} }
} }

View File

@@ -7,8 +7,10 @@ use std::time::{Duration, Instant};
use crate::mmp::algorithms::{JitterEstimator, OwdTrendDetector}; use crate::mmp::algorithms::{JitterEstimator, OwdTrendDetector};
use crate::mmp::report::ReceiverReport; use crate::mmp::report::ReceiverReport;
use crate::mmp::{DEFAULT_COLD_START_INTERVAL_MS, DEFAULT_OWD_WINDOW_SIZE, use crate::mmp::{
MAX_REPORT_INTERVAL_MS, MIN_REPORT_INTERVAL_MS}; DEFAULT_COLD_START_INTERVAL_MS, DEFAULT_OWD_WINDOW_SIZE, MAX_REPORT_INTERVAL_MS,
MIN_REPORT_INTERVAL_MS,
};
/// Grace period after rekey before resuming jitter calculation. /// Grace period after rekey before resuming jitter calculation.
/// ///
@@ -228,8 +230,7 @@ impl ReceiverState {
self.owd_seq = 0; self.owd_seq = 0;
self.last_sender_timestamp = 0; self.last_sender_timestamp = 0;
self.last_recv_time = None; self.last_recv_time = None;
self.rekey_jitter_grace_until = self.rekey_jitter_grace_until = Some(now + Duration::from_secs(REKEY_JITTER_GRACE_SECS));
Some(now + Duration::from_secs(REKEY_JITTER_GRACE_SECS));
self.ecn_ce_count = 0; self.ecn_ce_count = 0;
self.interval_has_data = false; self.interval_has_data = false;
// Keep cumulative_packets_recv, cumulative_bytes_recv (lifetime stats) // Keep cumulative_packets_recv, cumulative_bytes_recv (lifetime stats)
@@ -281,14 +282,14 @@ impl ReceiverState {
// We can't get absolute µs from Instant, but we can compute the delta // We can't get absolute µs from Instant, but we can compute the delta
// between consecutive transits using relative Instant differences. // between consecutive transits using relative Instant differences.
// Skip during post-rekey grace period to avoid drain-window spikes. // Skip during post-rekey grace period to avoid drain-window spikes.
let in_grace = self.rekey_jitter_grace_until let in_grace = self
.rekey_jitter_grace_until
.is_some_and(|deadline| now < deadline); .is_some_and(|deadline| now < deadline);
if !in_grace { if !in_grace {
self.rekey_jitter_grace_until = None; // clear expired grace self.rekey_jitter_grace_until = None; // clear expired grace
if let Some(prev_recv) = self.last_recv_time { if let Some(prev_recv) = self.last_recv_time {
let recv_delta_us = now.duration_since(prev_recv).as_micros() as i64; let recv_delta_us = now.duration_since(prev_recv).as_micros() as i64;
let send_delta_us = let send_delta_us = sender_us - (self.last_sender_timestamp as i64 * 1000);
sender_us - (self.last_sender_timestamp as i64 * 1000);
let transit_delta = (recv_delta_us - send_delta_us) as i32; let transit_delta = (recv_delta_us - send_delta_us) as i32;
self.jitter.update(transit_delta); self.jitter.update(transit_delta);
} }
@@ -318,7 +319,8 @@ impl ReceiverState {
} }
// Dwell time: ms between last frame reception and report generation // Dwell time: ms between last frame reception and report generation
let dwell_time = self.last_recv_time let dwell_time = self
.last_recv_time
.map(|t| now.duration_since(t).as_millis() as u16) .map(|t| now.duration_since(t).as_millis() as u16)
.unwrap_or(0); .unwrap_or(0);
@@ -365,7 +367,11 @@ impl ReceiverState {
/// ///
/// Receiver reports at 1× SRTT, clamped to [MIN, MAX]. /// Receiver reports at 1× SRTT, clamped to [MIN, MAX].
pub fn update_report_interval_from_srtt(&mut self, srtt_us: i64) { pub fn update_report_interval_from_srtt(&mut self, srtt_us: i64) {
self.update_report_interval_with_bounds(srtt_us, MIN_REPORT_INTERVAL_MS, MAX_REPORT_INTERVAL_MS); self.update_report_interval_with_bounds(
srtt_us,
MIN_REPORT_INTERVAL_MS,
MAX_REPORT_INTERVAL_MS,
);
} }
/// Update the report interval based on SRTT with custom bounds. /// Update the report interval based on SRTT with custom bounds.
@@ -600,8 +606,8 @@ mod tests {
assert_eq!(r.jitter_us(), 0); assert_eq!(r.jitter_us(), 0);
// After grace expires, jitter updates resume // After grace expires, jitter updates resume
let after_grace = t0 + Duration::from_secs(2) let after_grace =
+ Duration::from_secs(REKEY_JITTER_GRACE_SECS + 1); t0 + Duration::from_secs(2) + Duration::from_secs(REKEY_JITTER_GRACE_SECS + 1);
r.record_recv(2, 200, 100, false, after_grace); r.record_recv(2, 200, 100, false, after_grace);
r.record_recv(3, 300, 100, false, after_grace + Duration::from_millis(100)); r.record_recv(3, 300, 100, false, after_grace + Duration::from_millis(100));
// Now jitter should be updating (non-zero or zero depending on timing) // Now jitter should be updating (non-zero or zero depending on timing)

View File

@@ -117,7 +117,9 @@ impl SenderState {
match self.last_report_time { match self.last_report_time {
None => true, // Never sent a report — send immediately None => true, // Never sent a report — send immediately
Some(last) => { Some(last) => {
let effective = self.report_interval.mul_f64(self.send_failure_backoff_multiplier()); let effective = self
.report_interval
.mul_f64(self.send_failure_backoff_multiplier());
now.duration_since(last) >= effective now.duration_since(last) >= effective
} }
} }
@@ -152,7 +154,11 @@ impl SenderState {
/// ///
/// Sender reports at 2× SRTT clamped to [MIN, MAX]. /// Sender reports at 2× SRTT clamped to [MIN, MAX].
pub fn update_report_interval_from_srtt(&mut self, srtt_us: i64) { pub fn update_report_interval_from_srtt(&mut self, srtt_us: i64) {
self.update_report_interval_with_bounds(srtt_us, MIN_REPORT_INTERVAL_MS, MAX_REPORT_INTERVAL_MS); self.update_report_interval_with_bounds(
srtt_us,
MIN_REPORT_INTERVAL_MS,
MAX_REPORT_INTERVAL_MS,
);
} }
/// Update the report interval based on SRTT with custom bounds. /// Update the report interval based on SRTT with custom bounds.
@@ -301,7 +307,10 @@ mod tests {
// 2s RTT → 4s, clamped to max 2s // 2s RTT → 4s, clamped to max 2s
s.update_report_interval_from_srtt(2_000_000); s.update_report_interval_from_srtt(2_000_000);
assert_eq!(s.report_interval(), Duration::from_millis(MAX_REPORT_INTERVAL_MS)); assert_eq!(
s.report_interval(),
Duration::from_millis(MAX_REPORT_INTERVAL_MS)
);
} }
#[test] #[test]

View File

@@ -3,9 +3,9 @@
//! Handles building, sending, and receiving FilterAnnounce messages, //! Handles building, sending, and receiving FilterAnnounce messages,
//! including debounced propagation to peers. //! including debounced propagation to peers.
use crate::NodeAddr;
use crate::bloom::BloomFilter; use crate::bloom::BloomFilter;
use crate::protocol::FilterAnnounce; use crate::protocol::FilterAnnounce;
use crate::NodeAddr;
use super::{Node, NodeError}; use super::{Node, NodeError};
use std::collections::HashMap; use std::collections::HashMap;

View File

@@ -81,11 +81,7 @@ impl DiscoveryBackoff {
/// window using exponential backoff. /// window using exponential backoff.
pub fn record_failure(&mut self, target: &NodeAddr) { pub fn record_failure(&mut self, target: &NodeAddr) {
let now = Instant::now(); let now = Instant::now();
let failures = self let failures = self.entries.get(target).map_or(0, |e| e.failures) + 1;
.entries
.get(target)
.map_or(0, |e| e.failures)
+ 1;
let backoff_secs = self let backoff_secs = self
.base .base
@@ -345,8 +341,7 @@ mod tests {
#[test] #[test]
fn test_forward_allowed_after_interval() { fn test_forward_allowed_after_interval() {
let mut limiter = let mut limiter = DiscoveryForwardRateLimiter::with_interval(Duration::from_millis(100));
DiscoveryForwardRateLimiter::with_interval(Duration::from_millis(100));
assert!(limiter.should_forward(&addr(1))); assert!(limiter.should_forward(&addr(1)));
thread::sleep(Duration::from_millis(110)); thread::sleep(Duration::from_millis(110));

View File

@@ -20,11 +20,7 @@ impl Node {
/// 4. Lazy purge expired entries /// 4. Lazy purge expired entries
/// 5. If we're the target, generate and send response /// 5. If we're the target, generate and send response
/// 6. If TTL > 0, forward to tree peers whose bloom filter matches /// 6. If TTL > 0, forward to tree peers whose bloom filter matches
pub(in crate::node) async fn handle_lookup_request( pub(in crate::node) async fn handle_lookup_request(&mut self, from: &NodeAddr, payload: &[u8]) {
&mut self,
from: &NodeAddr,
payload: &[u8],
) {
self.stats_mut().discovery.req_received += 1; self.stats_mut().discovery.req_received += 1;
let request = match LookupRequest::decode(payload) { let request = match LookupRequest::decode(payload) {
@@ -52,10 +48,8 @@ impl Node {
} }
// Record for reverse-path forwarding and dedup // Record for reverse-path forwarding and dedup
self.recent_requests.insert( self.recent_requests
request.request_id, .insert(request.request_id, RecentRequest::new(*from, now_ms));
RecentRequest::new(*from, now_ms),
);
// Lazy purge expired entries // Lazy purge expired entries
self.purge_expired_requests(now_ms); self.purge_expired_requests(now_ms);
@@ -76,7 +70,10 @@ impl Node {
if request.can_forward() { if request.can_forward() {
// Transit-side rate limit: collapse rapid-fire lookups for the // Transit-side rate limit: collapse rapid-fire lookups for the
// same target from misbehaving nodes generating fresh request_ids. // same target from misbehaving nodes generating fresh request_ids.
if !self.discovery_forward_limiter.should_forward(&request.target) { if !self
.discovery_forward_limiter
.should_forward(&request.target)
{
self.stats_mut().discovery.req_forward_rate_limited += 1; self.stats_mut().discovery.req_forward_rate_limited += 1;
debug!( debug!(
request_id = request.request_id, request_id = request.request_id,
@@ -192,11 +189,8 @@ impl Node {
// Verify the proof signature // Verify the proof signature
let (xonly, _parity) = target_pubkey.x_only_public_key(); let (xonly, _parity) = target_pubkey.x_only_public_key();
let peer_id = PeerIdentity::from_pubkey(xonly); let peer_id = PeerIdentity::from_pubkey(xonly);
let proof_data = LookupResponse::proof_bytes( let proof_data =
response.request_id, LookupResponse::proof_bytes(response.request_id, &target, &response.target_coords);
&target,
&response.target_coords,
);
if !peer_id.verify(&proof_data, &response.proof) { if !peer_id.verify(&proof_data, &response.proof) {
self.stats_mut().discovery.resp_proof_failed += 1; self.stats_mut().discovery.resp_proof_failed += 1;
warn!( warn!(
@@ -220,12 +214,8 @@ impl Node {
"Discovery succeeded, proof verified, route cached" "Discovery succeeded, proof verified, route cached"
); );
self.coord_cache.insert_with_path_mtu( self.coord_cache
target, .insert_with_path_mtu(target, response.target_coords, now_ms, path_mtu);
response.target_coords,
now_ms,
path_mtu,
);
// Clean up pending lookup tracking // Clean up pending lookup tracking
self.pending_lookups.remove(&target); self.pending_lookups.remove(&target);
@@ -262,15 +252,11 @@ impl Node {
let our_coords = self.tree_state().my_coords().clone(); let our_coords = self.tree_state().my_coords().clone();
// Sign proof: Identity::sign hashes with SHA-256 internally // Sign proof: Identity::sign hashes with SHA-256 internally
let proof_data = LookupResponse::proof_bytes(request.request_id, &request.target, &our_coords); let proof_data =
LookupResponse::proof_bytes(request.request_id, &request.target, &our_coords);
let proof = self.identity().sign(&proof_data); let proof = self.identity().sign(&proof_data);
let response = LookupResponse::new( let response = LookupResponse::new(request.request_id, request.target, our_coords, proof);
request.request_id,
request.target,
our_coords,
proof,
);
// Route toward origin via reverse path. // Route toward origin via reverse path.
let next_hop_addr = if let Some(recent) = self.recent_requests.get(&request.request_id) { let next_hop_addr = if let Some(recent) = self.recent_requests.get(&request.request_id) {
@@ -297,7 +283,10 @@ impl Node {
); );
let encoded = response.encode(); let encoded = response.encode();
if let Err(e) = self.send_encrypted_link_message(&next_hop_addr, &encoded).await { if let Err(e) = self
.send_encrypted_link_message(&next_hop_addr, &encoded)
.await
{
debug!( debug!(
next_hop = %self.peer_display_name(&next_hop_addr), next_hop = %self.peer_display_name(&next_hop_addr),
error = %e, error = %e,
@@ -324,9 +313,7 @@ impl Node {
let forward_to: Vec<NodeAddr> = self let forward_to: Vec<NodeAddr> = self
.peers .peers
.iter() .iter()
.filter(|(addr, peer)| { .filter(|(addr, peer)| self.is_tree_peer(addr) && peer.may_reach(&request.target))
self.is_tree_peer(addr) && peer.may_reach(&request.target)
})
.map(|(addr, _)| *addr) .map(|(addr, _)| *addr)
.collect(); .collect();
@@ -335,9 +322,7 @@ impl Node {
let fallback: Vec<NodeAddr> = self let fallback: Vec<NodeAddr> = self
.peers .peers
.iter() .iter()
.filter(|(addr, peer)| { .filter(|(addr, peer)| !self.is_tree_peer(addr) && peer.may_reach(&request.target))
!self.is_tree_peer(addr) && peer.may_reach(&request.target)
})
.map(|(addr, _)| *addr) .map(|(addr, _)| *addr)
.collect(); .collect();
if fallback.is_empty() { if fallback.is_empty() {
@@ -402,9 +387,7 @@ impl Node {
let peer_addrs: Vec<NodeAddr> = self let peer_addrs: Vec<NodeAddr> = self
.peers .peers
.iter() .iter()
.filter(|(addr, peer)| { .filter(|(addr, peer)| self.is_tree_peer(addr) && peer.may_reach(target))
self.is_tree_peer(addr) && peer.may_reach(target)
})
.map(|(addr, _)| *addr) .map(|(addr, _)| *addr)
.collect(); .collect();
@@ -488,7 +471,8 @@ impl Node {
return; return;
} }
self.pending_lookups.insert(*dest, PendingLookup::new(now_ms)); self.pending_lookups
.insert(*dest, PendingLookup::new(now_ms));
let ttl = self.config.node.discovery.ttl; let ttl = self.config.node.discovery.ttl;
let sent = self.initiate_lookup(dest, ttl).await; let sent = self.initiate_lookup(dest, ttl).await;

View File

@@ -1,14 +1,19 @@
//! Link message dispatch and peer removal. //! Link message dispatch and peer removal.
use crate::node::Node;
use crate::NodeAddr; use crate::NodeAddr;
use crate::node::Node;
use tracing::{debug, info, trace}; use tracing::{debug, info, trace};
impl Node { impl Node {
/// Dispatch a decrypted link message to the appropriate handler. /// Dispatch a decrypted link message to the appropriate handler.
/// ///
/// Link messages are protocol messages exchanged between authenticated peers. /// Link messages are protocol messages exchanged between authenticated peers.
pub(in crate::node) async fn dispatch_link_message(&mut self, from: &NodeAddr, plaintext: &[u8], ce_flag: bool) { pub(in crate::node) async fn dispatch_link_message(
&mut self,
from: &NodeAddr,
plaintext: &[u8],
ce_flag: bool,
) {
if plaintext.is_empty() { if plaintext.is_empty() {
return; return;
} }
@@ -109,7 +114,9 @@ impl Node {
} }
// MMP teardown log (before we drop the peer) // MMP teardown log (before we drop the peer)
let peer_name = self.peer_aliases.get(node_addr) let peer_name = self
.peer_aliases
.get(node_addr)
.cloned() .cloned()
.unwrap_or_else(|| peer.identity().short_npub()); .unwrap_or_else(|| peer.identity().short_npub());
if let Some(mmp) = peer.mmp() { if let Some(mmp) = peer.mmp() {

View File

@@ -1,8 +1,8 @@
//! Encrypted frame handling (hot path). //! Encrypted frame handling (hot path).
use crate::noise::NoiseError;
use crate::node::Node; use crate::node::Node;
use crate::node::wire::{EncryptedHeader, strip_inner_header, FLAG_CE, FLAG_KEY_EPOCH, FLAG_SP}; use crate::node::wire::{EncryptedHeader, FLAG_CE, FLAG_KEY_EPOCH, FLAG_SP, strip_inner_header};
use crate::noise::NoiseError;
use crate::transport::ReceivedPacket; use crate::transport::ReceivedPacket;
use std::time::Instant; use std::time::Instant;
use tracing::{debug, info, trace, warn}; use tracing::{debug, info, trace, warn};
@@ -53,8 +53,8 @@ impl Node {
// Check and perform cutover in a scoped borrow. // Check and perform cutover in a scoped borrow.
{ {
let peer = self.peers.get(&node_addr).unwrap(); let peer = self.peers.get(&node_addr).unwrap();
let k_bit_flipped = received_k_bit != peer.current_k_bit() let k_bit_flipped =
&& peer.pending_new_session().is_some(); received_k_bit != peer.current_k_bit() && peer.pending_new_session().is_some();
if k_bit_flipped { if k_bit_flipped {
let display_name = self.peer_display_name(&node_addr); let display_name = self.peer_display_name(&node_addr);
@@ -70,9 +70,10 @@ impl Node {
debug_assert!( debug_assert!(
peer.transport_id().is_some() peer.transport_id().is_some()
&& peer.our_index().is_some() && peer.our_index().is_some()
&& self.peers_by_index.contains_key( && self.peers_by_index.contains_key(&(
&(peer.transport_id().unwrap(), peer.our_index().unwrap().as_u32()) peer.transport_id().unwrap(),
), peer.our_index().unwrap().as_u32()
)),
"peers_by_index should contain pre-registered new index after K-bit flip" "peers_by_index should contain pre-registered new index after K-bit flip"
); );
} }
@@ -162,12 +163,14 @@ impl Node {
let _spin_rtt = mmp.spin_bit.rx_observe(sp_flag, header.counter, now); let _spin_rtt = mmp.spin_bit.rx_observe(sp_flag, header.counter, now);
} }
peer.set_current_addr(packet.transport_id, packet.remote_addr.clone()); peer.set_current_addr(packet.transport_id, packet.remote_addr.clone());
peer.link_stats_mut().record_recv(packet.data.len(), packet.timestamp_ms); peer.link_stats_mut()
.record_recv(packet.data.len(), packet.timestamp_ms);
peer.touch(packet.timestamp_ms); peer.touch(packet.timestamp_ms);
} }
// Dispatch to link message handler // Dispatch to link message handler
self.dispatch_link_message(&node_addr, link_message, ce_flag).await; self.dispatch_link_message(&node_addr, link_message, ce_flag)
.await;
} }
/// Log a decryption failure with replay suppression. /// Log a decryption failure with replay suppression.

View File

@@ -5,15 +5,15 @@
//! plaintext session-layer headers, routes to the next hop or delivers //! plaintext session-layer headers, routes to the next hop or delivers
//! locally, and generates error signals on routing failure. //! locally, and generates error signals on routing failure.
use crate::node::{Node, NodeError}; use crate::NodeAddr;
use crate::node::session_wire::{ use crate::node::session_wire::{
parse_encrypted_coords, FspCommonPrefix, FSP_COMMON_PREFIX_SIZE, FSP_HEADER_SIZE, FSP_COMMON_PREFIX_SIZE, FSP_HEADER_SIZE, FSP_PHASE_ESTABLISHED, FSP_PHASE_MSG1, FSP_PHASE_MSG2,
FSP_PHASE_ESTABLISHED, FSP_PHASE_MSG1, FSP_PHASE_MSG2, FspCommonPrefix, parse_encrypted_coords,
}; };
use crate::node::{Node, NodeError};
use crate::protocol::{ use crate::protocol::{
CoordsRequired, MtuExceeded, PathBroken, SessionAck, SessionDatagram, SessionSetup, CoordsRequired, MtuExceeded, PathBroken, SessionAck, SessionDatagram, SessionSetup,
}; };
use crate::NodeAddr;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use tracing::{debug, warn}; use tracing::{debug, warn};
@@ -22,13 +22,20 @@ impl Node {
/// ///
/// Called by `dispatch_link_message` for msg_type 0x00. The payload /// Called by `dispatch_link_message` for msg_type 0x00. The payload
/// has already had its msg_type byte stripped by dispatch. /// has already had its msg_type byte stripped by dispatch.
pub(in crate::node) async fn handle_session_datagram(&mut self, _from: &NodeAddr, payload: &[u8], incoming_ce: bool) { pub(in crate::node) async fn handle_session_datagram(
&mut self,
_from: &NodeAddr,
payload: &[u8],
incoming_ce: bool,
) {
self.stats_mut().forwarding.record_received(payload.len()); self.stats_mut().forwarding.record_received(payload.len());
let mut datagram = match SessionDatagram::decode(payload) { let mut datagram = match SessionDatagram::decode(payload) {
Ok(dg) => dg, Ok(dg) => dg,
Err(e) => { Err(e) => {
self.stats_mut().forwarding.record_decode_error(payload.len()); self.stats_mut()
.forwarding
.record_decode_error(payload.len());
debug!(error = %e, "Malformed SessionDatagram"); debug!(error = %e, "Malformed SessionDatagram");
return; return;
} }
@@ -36,7 +43,9 @@ impl Node {
// TTL enforcement: decrement and drop if exhausted // TTL enforcement: decrement and drop if exhausted
if !datagram.decrement_ttl() { if !datagram.decrement_ttl() {
self.stats_mut().forwarding.record_ttl_exhausted(payload.len()); self.stats_mut()
.forwarding
.record_ttl_exhausted(payload.len());
debug!( debug!(
src = %datagram.src_addr, src = %datagram.src_addr,
dest = %datagram.dest_addr, dest = %datagram.dest_addr,
@@ -51,8 +60,13 @@ impl Node {
// Local delivery: dispatch to session layer handlers // Local delivery: dispatch to session layer handlers
if datagram.dest_addr == *self.node_addr() { if datagram.dest_addr == *self.node_addr() {
self.stats_mut().forwarding.record_delivered(payload.len()); self.stats_mut().forwarding.record_delivered(payload.len());
self.handle_session_payload(&datagram.src_addr, &datagram.payload, datagram.path_mtu, incoming_ce) self.handle_session_payload(
.await; &datagram.src_addr,
&datagram.payload,
datagram.path_mtu,
incoming_ce,
)
.await;
return; return;
} }
@@ -60,7 +74,9 @@ impl Node {
let next_hop_addr = match self.find_next_hop(&datagram.dest_addr) { let next_hop_addr = match self.find_next_hop(&datagram.dest_addr) {
Some(peer) => *peer.node_addr(), Some(peer) => *peer.node_addr(),
None => { None => {
self.stats_mut().forwarding.record_drop_no_route(payload.len()); self.stats_mut()
.forwarding
.record_drop_no_route(payload.len());
self.send_routing_error(&datagram).await; self.send_routing_error(&datagram).await;
return; return;
} }
@@ -84,7 +100,8 @@ impl Node {
if local_congestion { if local_congestion {
self.stats_mut().congestion.record_congestion_detected(); self.stats_mut().congestion.record_congestion_detected();
let now = Instant::now(); let now = Instant::now();
let should_log = self.last_congestion_log let should_log = self
.last_congestion_log
.map(|t| now.duration_since(t) >= Duration::from_secs(5)) .map(|t| now.duration_since(t) >= Duration::from_secs(5))
.unwrap_or(true); .unwrap_or(true);
if should_log { if should_log {
@@ -101,11 +118,15 @@ impl Node {
{ {
match e { match e {
NodeError::MtuExceeded { mtu, .. } => { NodeError::MtuExceeded { mtu, .. } => {
self.stats_mut().forwarding.record_drop_mtu_exceeded(payload.len()); self.stats_mut()
.forwarding
.record_drop_mtu_exceeded(payload.len());
self.send_mtu_exceeded_error(&datagram, mtu).await; self.send_mtu_exceeded_error(&datagram, mtu).await;
} }
_ => { _ => {
self.stats_mut().forwarding.record_drop_send_error(payload.len()); self.stats_mut()
.forwarding
.record_drop_send_error(payload.len());
debug!( debug!(
next_hop = %next_hop_addr, next_hop = %next_hop_addr,
dest = %datagram.dest_addr, dest = %datagram.dest_addr,
@@ -146,54 +167,38 @@ impl Node {
.unwrap_or(0); .unwrap_or(0);
match prefix.phase { match prefix.phase {
FSP_PHASE_MSG1 => { FSP_PHASE_MSG1 => match SessionSetup::decode(inner) {
match SessionSetup::decode(inner) { Ok(setup) => {
Ok(setup) => { self.coord_cache_mut()
self.coord_cache_mut().insert( .insert(datagram.src_addr, setup.src_coords, now_ms);
datagram.src_addr, self.coord_cache_mut()
setup.src_coords, .insert(datagram.dest_addr, setup.dest_coords, now_ms);
now_ms, debug!(
); src = %datagram.src_addr,
self.coord_cache_mut().insert( dest = %datagram.dest_addr,
datagram.dest_addr, "Cached coords from SessionSetup"
setup.dest_coords, );
now_ms,
);
debug!(
src = %datagram.src_addr,
dest = %datagram.dest_addr,
"Cached coords from SessionSetup"
);
}
Err(e) => {
debug!(error = %e, "Failed to decode SessionSetup for cache warming");
}
} }
} Err(e) => {
FSP_PHASE_MSG2 => { debug!(error = %e, "Failed to decode SessionSetup for cache warming");
match SessionAck::decode(inner) {
Ok(ack) => {
self.coord_cache_mut().insert(
datagram.src_addr,
ack.src_coords,
now_ms,
);
self.coord_cache_mut().insert(
datagram.dest_addr,
ack.dest_coords,
now_ms,
);
debug!(
src = %datagram.src_addr,
dest = %datagram.dest_addr,
"Cached coords from SessionAck"
);
}
Err(e) => {
debug!(error = %e, "Failed to decode SessionAck for cache warming");
}
} }
} },
FSP_PHASE_MSG2 => match SessionAck::decode(inner) {
Ok(ack) => {
self.coord_cache_mut()
.insert(datagram.src_addr, ack.src_coords, now_ms);
self.coord_cache_mut()
.insert(datagram.dest_addr, ack.dest_coords, now_ms);
debug!(
src = %datagram.src_addr,
dest = %datagram.dest_addr,
"Cached coords from SessionAck"
);
}
Err(e) => {
debug!(error = %e, "Failed to decode SessionAck for cache warming");
}
},
FSP_PHASE_ESTABLISHED if prefix.has_coords() => { FSP_PHASE_ESTABLISHED if prefix.has_coords() => {
// CP flag set: coords in cleartext between header and ciphertext. // CP flag set: coords in cleartext between header and ciphertext.
// Parse coords from the cleartext section after the 12-byte header. // Parse coords from the cleartext section after the 12-byte header.
@@ -203,18 +208,12 @@ impl Node {
match parse_encrypted_coords(coord_data) { match parse_encrypted_coords(coord_data) {
Ok((src_coords, dest_coords, _bytes_consumed)) => { Ok((src_coords, dest_coords, _bytes_consumed)) => {
if let Some(coords) = src_coords { if let Some(coords) = src_coords {
self.coord_cache_mut().insert( self.coord_cache_mut()
datagram.src_addr, .insert(datagram.src_addr, coords, now_ms);
coords,
now_ms,
);
} }
if let Some(coords) = dest_coords { if let Some(coords) = dest_coords {
self.coord_cache_mut().insert( self.coord_cache_mut()
datagram.dest_addr, .insert(datagram.dest_addr, coords, now_ms);
coords,
now_ms,
);
} }
debug!( debug!(
src = %datagram.src_addr, src = %datagram.src_addr,
@@ -243,7 +242,10 @@ impl Node {
/// No cascading errors. /// No cascading errors.
async fn send_routing_error(&mut self, original: &SessionDatagram) { async fn send_routing_error(&mut self, original: &SessionDatagram) {
// Rate limit: one error signal per destination per 100ms // Rate limit: one error signal per destination per 100ms
if !self.routing_error_rate_limiter.should_send(&original.dest_addr) { if !self
.routing_error_rate_limiter
.should_send(&original.dest_addr)
{
return; return;
} }
@@ -303,23 +305,18 @@ impl Node {
/// Called when `send_encrypted_link_message()` fails with /// Called when `send_encrypted_link_message()` fails with
/// `NodeError::MtuExceeded` during forwarding. The signal tells the /// `NodeError::MtuExceeded` during forwarding. The signal tells the
/// source the bottleneck MTU so it can immediately reduce its path MTU. /// source the bottleneck MTU so it can immediately reduce its path MTU.
async fn send_mtu_exceeded_error( async fn send_mtu_exceeded_error(&mut self, original: &SessionDatagram, bottleneck_mtu: u16) {
&mut self,
original: &SessionDatagram,
bottleneck_mtu: u16,
) {
// Rate limit: reuse routing_error_rate_limiter keyed on dest_addr // Rate limit: reuse routing_error_rate_limiter keyed on dest_addr
if !self.routing_error_rate_limiter.should_send(&original.dest_addr) { if !self
.routing_error_rate_limiter
.should_send(&original.dest_addr)
{
return; return;
} }
let my_addr = *self.node_addr(); let my_addr = *self.node_addr();
let error_payload = MtuExceeded::new( let error_payload = MtuExceeded::new(original.dest_addr, my_addr, bottleneck_mtu).encode();
original.dest_addr,
my_addr,
bottleneck_mtu,
).encode();
let error_dg = SessionDatagram::new(my_addr, original.src_addr, error_payload) let error_dg = SessionDatagram::new(my_addr, original.src_addr, error_payload)
.with_ttl(self.config.node.session.default_ttl); .with_ttl(self.config.node.session.default_ttl);
@@ -403,7 +400,10 @@ impl Node {
} }
for tid in new_drop_events { for tid in new_drop_events {
self.stats_mut().congestion.record_kernel_drop_event(); self.stats_mut().congestion.record_kernel_drop_event();
warn!(transport_id = tid.as_u32(), "Kernel recv drops first observed on transport"); warn!(
transport_id = tid.as_u32(),
"Kernel recv drops first observed on transport"
);
} }
} }
} }

View File

@@ -1,12 +1,10 @@
//! Handshake handlers and connection promotion. //! Handshake handlers and connection promotion.
use crate::node::{Node, NodeError};
use crate::peer::{
cross_connection_winner, ActivePeer, PeerConnection, PromotionResult,
};
use crate::transport::{Link, LinkDirection, LinkId, ReceivedPacket};
use crate::node::wire::{build_msg2, Msg1Header, Msg2Header};
use crate::PeerIdentity; use crate::PeerIdentity;
use crate::node::wire::{Msg1Header, Msg2Header, build_msg2};
use crate::node::{Node, NodeError};
use crate::peer::{ActivePeer, PeerConnection, PromotionResult, cross_connection_winner};
use crate::transport::{Link, LinkDirection, LinkId, ReceivedPacket};
use std::time::Duration; use std::time::Duration;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
@@ -63,8 +61,7 @@ impl Node {
{ {
if link.direction() == LinkDirection::Inbound { if link.direction() == LinkDirection::Inbound {
// Check if this link belongs to an already-promoted active peer // Check if this link belongs to an already-promoted active peer
let is_active_peer = self.peers.values() let is_active_peer = self.peers.values().any(|p| p.link_id() == existing_link_id);
.any(|p| p.link_id() == existing_link_id);
if is_active_peer { if is_active_peer {
// Possible restart — fall through to decrypt and check epoch // Possible restart — fall through to decrypt and check epoch
@@ -100,8 +97,7 @@ impl Node {
// peer, this may be a rekey msg1 (same epoch) or a // peer, this may be a rekey msg1 (same epoch) or a
// restart (different epoch). Set possible_restart to enable // restart (different epoch). Set possible_restart to enable
// the epoch/rekey check below. // the epoch/rekey check below.
let is_active_peer = self.peers.values() let is_active_peer = self.peers.values().any(|p| p.link_id() == existing_link_id);
.any(|p| p.link_id() == existing_link_id);
if is_active_peer { if is_active_peer {
possible_restart = true; possible_restart = true;
} else { } else {
@@ -126,7 +122,12 @@ impl Node {
let our_keypair = self.identity.keypair(); let our_keypair = self.identity.keypair();
let noise_msg1 = &packet.data[header.noise_msg1_offset..]; let noise_msg1 = &packet.data[header.noise_msg1_offset..];
let msg2_response = match conn.receive_handshake_init(our_keypair, self.startup_epoch, noise_msg1, packet.timestamp_ms) { let msg2_response = match conn.receive_handshake_init(
our_keypair,
self.startup_epoch,
noise_msg1,
packet.timestamp_ms,
) {
Ok(m) => m, Ok(m) => m,
Err(e) => { Err(e) => {
self.msg1_rate_limiter.complete_handshake(); self.msg1_rate_limiter.complete_handshake();
@@ -162,9 +163,7 @@ impl Node {
// If we fell through from the addr_to_link check above with // If we fell through from the addr_to_link check above with
// possible_restart=true, we now have the decrypted epoch from msg1. // possible_restart=true, we now have the decrypted epoch from msg1.
// Compare it against the stored epoch for this peer. // Compare it against the stored epoch for this peer.
if possible_restart if possible_restart && let Some(existing_peer) = self.peers.get(&peer_node_addr) {
&& let Some(existing_peer) = self.peers.get(&peer_node_addr)
{
let new_epoch = conn.remote_epoch(); let new_epoch = conn.remote_epoch();
let existing_epoch = existing_peer.remote_epoch(); let existing_epoch = existing_peer.remote_epoch();
@@ -192,10 +191,8 @@ impl Node {
// During simultaneous connection, both sides promote // During simultaneous connection, both sides promote
// within the same tick and the peer's msg1 arrives // within the same tick and the peer's msg1 arrives
// immediately — a genuine rekey can't fire that fast. // immediately — a genuine rekey can't fire that fast.
let session_age_secs = existing_peer let session_age_secs =
.session_established_at() existing_peer.session_established_at().elapsed().as_secs();
.elapsed()
.as_secs();
if self.config.node.rekey.enabled if self.config.node.rekey.enabled
&& existing_peer.has_session() && existing_peer.has_session()
&& existing_peer.is_healthy() && existing_peer.is_healthy()
@@ -272,7 +269,8 @@ impl Node {
}; };
// Send msg2 response using the new handshake // Send msg2 response using the new handshake
let wire_msg2 = build_msg2(our_new_index, header.sender_idx, &msg2_response); let wire_msg2 =
build_msg2(our_new_index, header.sender_idx, &msg2_response);
if let Some(transport) = self.transports.get(&packet.transport_id) { if let Some(transport) = self.transports.get(&packet.transport_id) {
match transport.send(&packet.remote_addr, &wire_msg2).await { match transport.send(&packet.remote_addr, &wire_msg2).await {
Ok(_) => { Ok(_) => {
@@ -401,7 +399,8 @@ impl Node {
// Clean up on failure // Clean up on failure
self.connections.remove(&link_id); self.connections.remove(&link_id);
self.links.remove(&link_id); self.links.remove(&link_id);
self.addr_to_link.remove(&(packet.transport_id, packet.remote_addr)); self.addr_to_link
.remove(&(packet.transport_id, packet.remote_addr));
let _ = self.index_allocator.free(our_index); let _ = self.index_allocator.free(our_index);
self.msg1_rate_limiter.complete_handshake(); self.msg1_rate_limiter.complete_handshake();
return; return;
@@ -434,7 +433,10 @@ impl Node {
self.bloom_state.mark_update_needed(node_addr); self.bloom_state.mark_update_needed(node_addr);
self.reset_discovery_backoff(); self.reset_discovery_backoff();
} }
PromotionResult::CrossConnectionWon { loser_link_id, node_addr } => { PromotionResult::CrossConnectionWon {
loser_link_id,
node_addr,
} => {
// Store msg2 on peer for resend on duplicate msg1 // Store msg2 on peer for resend on duplicate msg1
if let Some(peer) = self.peers.get_mut(&node_addr) { if let Some(peer) = self.peers.get_mut(&node_addr) {
peer.set_handshake_msg2(wire_msg2.clone()); peer.set_handshake_msg2(wire_msg2.clone());
@@ -552,9 +554,7 @@ impl Node {
// Find peer with rekey in progress for this index // Find peer with rekey in progress for this index
let peer_addr = self.peers.iter().find_map(|(addr, peer)| { let peer_addr = self.peers.iter().find_map(|(addr, peer)| {
if peer.rekey_in_progress() if peer.rekey_in_progress() && peer.rekey_our_index() == Some(header.receiver_idx) {
&& peer.rekey_our_index() == Some(header.receiver_idx)
{
Some(*addr) Some(*addr)
} else { } else {
None None
@@ -568,15 +568,12 @@ impl Node {
if let Some(peer) = self.peers.get_mut(&peer_node_addr) { if let Some(peer) = self.peers.get_mut(&peer_node_addr) {
match peer.complete_rekey_msg2(noise_msg2) { match peer.complete_rekey_msg2(noise_msg2) {
Ok(session) => { Ok(session) => {
let our_index = peer.rekey_our_index() let our_index = peer.rekey_our_index().unwrap_or(header.receiver_idx);
.unwrap_or(header.receiver_idx);
peer.set_pending_session(session, our_index, header.sender_idx); peer.set_pending_session(session, our_index, header.sender_idx);
if let Some(transport_id) = peer.transport_id() { if let Some(transport_id) = peer.transport_id() {
self.peers_by_index.insert( self.peers_by_index
(transport_id, our_index.as_u32()), .insert((transport_id, our_index.as_u32()), peer_node_addr);
peer_node_addr,
);
} }
debug!( debug!(
@@ -679,15 +676,17 @@ impl Node {
let outbound_our_index = conn.our_index(); let outbound_our_index = conn.our_index();
let outbound_session = conn.take_session(); let outbound_session = conn.take_session();
let (outbound_session, outbound_our_index) = let (outbound_session, outbound_our_index) = match (
match (outbound_session, outbound_our_index) { outbound_session,
(Some(s), Some(idx)) => (s, idx), outbound_our_index,
_ => { ) {
warn!(peer = %self.peer_display_name(&peer_node_addr), "Incomplete outbound connection"); (Some(s), Some(idx)) => (s, idx),
self.pending_outbound.remove(&key); _ => {
return; warn!(peer = %self.peer_display_name(&peer_node_addr), "Incomplete outbound connection");
} self.pending_outbound.remove(&key);
}; return;
}
};
if let Some(peer) = self.peers.get_mut(&peer_node_addr) { if let Some(peer) = self.peers.get_mut(&peer_node_addr) {
let suppressed = peer.replay_suppressed_count(); let suppressed = peer.replay_suppressed_count();
@@ -700,13 +699,12 @@ impl Node {
// Update peers_by_index: remove old inbound index, add outbound // Update peers_by_index: remove old inbound index, add outbound
let transport_id = peer.transport_id().unwrap(); let transport_id = peer.transport_id().unwrap();
if let Some(old_idx) = old_our_index { if let Some(old_idx) = old_our_index {
self.peers_by_index.remove(&(transport_id, old_idx.as_u32())); self.peers_by_index
.remove(&(transport_id, old_idx.as_u32()));
let _ = self.index_allocator.free(old_idx); let _ = self.index_allocator.free(old_idx);
} }
self.peers_by_index.insert( self.peers_by_index
(transport_id, outbound_our_index.as_u32()), .insert((transport_id, outbound_our_index.as_u32()), peer_node_addr);
peer_node_addr,
);
if suppressed > 0 { if suppressed > 0 {
debug!( debug!(
@@ -791,7 +789,10 @@ impl Node {
self.bloom_state.mark_update_needed(node_addr); self.bloom_state.mark_update_needed(node_addr);
self.reset_discovery_backoff(); self.reset_discovery_backoff();
} }
PromotionResult::CrossConnectionWon { loser_link_id, node_addr } => { PromotionResult::CrossConnectionWon {
loser_link_id,
node_addr,
} => {
// Close the losing TCP connection (no-op for connectionless) // Close the losing TCP connection (no-op for connectionless)
if let Some(loser_link) = self.links.get(&loser_link_id) { if let Some(loser_link) = self.links.get(&loser_link_id) {
let loser_tid = loser_link.transport_id(); let loser_tid = loser_link.transport_id();
@@ -803,10 +804,8 @@ impl Node {
// Clean up the losing connection's link // Clean up the losing connection's link
self.remove_link(&loser_link_id); self.remove_link(&loser_link_id);
// Ensure addr_to_link points to the winning link // Ensure addr_to_link points to the winning link
self.addr_to_link.insert( self.addr_to_link
(packet.transport_id, packet.remote_addr.clone()), .insert((packet.transport_id, packet.remote_addr.clone()), link_id);
link_id,
);
info!( info!(
peer = %self.peer_display_name(&node_addr), peer = %self.peer_display_name(&node_addr),
loser_link_id = %loser_link_id, loser_link_id = %loser_link_id,
@@ -873,30 +872,31 @@ impl Node {
.take_session() .take_session()
.ok_or(NodeError::NoSession(link_id))?; .ok_or(NodeError::NoSession(link_id))?;
let our_index = connection.our_index().ok_or_else(|| { let our_index = connection
NodeError::PromotionFailed { .our_index()
.ok_or_else(|| NodeError::PromotionFailed {
link_id, link_id,
reason: "missing our_index".into(), reason: "missing our_index".into(),
} })?;
})?; let their_index = connection
let their_index = connection.their_index().ok_or_else(|| { .their_index()
NodeError::PromotionFailed { .ok_or_else(|| NodeError::PromotionFailed {
link_id, link_id,
reason: "missing their_index".into(), reason: "missing their_index".into(),
} })?;
})?; let transport_id = connection
let transport_id = connection.transport_id().ok_or_else(|| { .transport_id()
NodeError::PromotionFailed { .ok_or_else(|| NodeError::PromotionFailed {
link_id, link_id,
reason: "missing transport_id".into(), reason: "missing transport_id".into(),
} })?;
})?; let current_addr = connection
let current_addr = connection.source_addr().ok_or_else(|| { .source_addr()
NodeError::PromotionFailed { .ok_or_else(|| NodeError::PromotionFailed {
link_id, link_id,
reason: "missing source_addr".into(), reason: "missing source_addr".into(),
} })?
})?.clone(); .clone();
let link_stats = connection.link_stats().clone(); let link_stats = connection.link_stats().clone();
let remote_epoch = connection.remote_epoch(); let remote_epoch = connection.remote_epoch();
@@ -908,11 +908,8 @@ impl Node {
let existing_link_id = existing_peer.link_id(); let existing_link_id = existing_peer.link_id();
// Determine which connection wins // Determine which connection wins
let this_wins = cross_connection_winner( let this_wins =
self.identity.node_addr(), cross_connection_winner(self.identity.node_addr(), &peer_node_addr, is_outbound);
&peer_node_addr,
is_outbound,
);
if this_wins { if this_wins {
// This connection wins, replace the existing peer // This connection wins, replace the existing peer
@@ -923,8 +920,7 @@ impl Node {
if let (Some(old_tid), Some(old_idx)) = if let (Some(old_tid), Some(old_idx)) =
(old_peer.transport_id(), old_peer.our_index()) (old_peer.transport_id(), old_peer.our_index())
{ {
self.peers_by_index self.peers_by_index.remove(&(old_tid, old_idx.as_u32()));
.remove(&(old_tid, old_idx.as_u32()));
let _ = self.index_allocator.free(old_idx); let _ = self.index_allocator.free(old_idx);
} }
@@ -942,7 +938,9 @@ impl Node {
&self.config.node.mmp, &self.config.node.mmp,
remote_epoch, remote_epoch,
); );
new_peer.set_tree_announce_min_interval_ms(self.config.node.tree.announce_min_interval_ms); new_peer.set_tree_announce_min_interval_ms(
self.config.node.tree.announce_min_interval_ms,
);
self.peers.insert(peer_node_addr, new_peer); self.peers.insert(peer_node_addr, new_peer);
self.peers_by_index self.peers_by_index
@@ -1008,13 +1006,17 @@ impl Node {
// Normal promotion // Normal promotion
if self.max_peers > 0 && self.peers.len() >= self.max_peers { if self.max_peers > 0 && self.peers.len() >= self.max_peers {
let _ = self.index_allocator.free(our_index); let _ = self.index_allocator.free(our_index);
return Err(NodeError::MaxPeersExceeded { max: self.max_peers }); return Err(NodeError::MaxPeersExceeded {
max: self.max_peers,
});
} }
// Preserve tree announce rate-limit state from old peer (if reconnecting). // Preserve tree announce rate-limit state from old peer (if reconnecting).
// Without this, reconnection resets the rate limit window to zero, // Without this, reconnection resets the rate limit window to zero,
// allowing an immediate announce that can feed an announce loop. // allowing an immediate announce that can feed an announce loop.
let old_announce_ts = self.peers.get(&peer_node_addr) let old_announce_ts = self
.peers
.get(&peer_node_addr)
.map(|p| p.last_tree_announce_sent_ms()); .map(|p| p.last_tree_announce_sent_ms());
let mut new_peer = ActivePeer::with_session( let mut new_peer = ActivePeer::with_session(
@@ -1031,7 +1033,8 @@ impl Node {
&self.config.node.mmp, &self.config.node.mmp,
remote_epoch, remote_epoch,
); );
new_peer.set_tree_announce_min_interval_ms(self.config.node.tree.announce_min_interval_ms); new_peer
.set_tree_announce_min_interval_ms(self.config.node.tree.announce_min_interval_ms);
if let Some(ts) = old_announce_ts { if let Some(ts) = old_announce_ts {
new_peer.set_last_tree_announce_sent_ms(ts); new_peer.set_last_tree_announce_sent_ms(ts);
} }

View File

@@ -4,6 +4,7 @@
//! periodic report generation on the tick timer, and emits periodic //! periodic report generation on the tick timer, and emits periodic
//! and teardown metric logs. //! and teardown metric logs.
use crate::NodeAddr;
use crate::mmp::MmpMode; use crate::mmp::MmpMode;
use crate::mmp::MmpSessionState; use crate::mmp::MmpSessionState;
use crate::mmp::report::{ReceiverReport, SenderReport}; use crate::mmp::report::{ReceiverReport, SenderReport};
@@ -12,7 +13,6 @@ use crate::protocol::{
LinkMessageType, PathMtuNotification, SessionMessageType, SessionReceiverReport, LinkMessageType, PathMtuNotification, SessionMessageType, SessionReceiverReport,
SessionSenderReport, SessionSenderReport,
}; };
use crate::NodeAddr;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use tracing::{debug, info, trace, warn}; use tracing::{debug, info, trace, warn};
@@ -72,7 +72,11 @@ impl Node {
/// ///
/// The peer is telling us about what they received from us. We feed /// The peer is telling us about what they received from us. We feed
/// this to our metrics to compute RTT, loss rate, and trend indicators. /// this to our metrics to compute RTT, loss rate, and trend indicators.
pub(in crate::node) async fn handle_receiver_report(&mut self, from: &NodeAddr, payload: &[u8]) { pub(in crate::node) async fn handle_receiver_report(
&mut self,
from: &NodeAddr,
payload: &[u8],
) {
let rr = match ReceiverReport::decode(payload) { let rr = match ReceiverReport::decode(payload) {
Ok(rr) => rr, Ok(rr) => rr,
Err(e) => { Err(e) => {
@@ -101,7 +105,9 @@ impl Node {
// Process the report: computes RTT from timestamp echo, updates // Process the report: computes RTT from timestamp echo, updates
// loss rate, goodput rate, jitter trend, and ETX. // loss rate, goodput rate, jitter trend, and ETX.
let now = Instant::now(); let now = Instant::now();
let first_rtt = mmp.metrics.process_receiver_report(&rr, our_timestamp_ms, now); let first_rtt = mmp
.metrics
.process_receiver_report(&rr, our_timestamp_ms, now);
// Feed SRTT back to sender/receiver report interval tuning // Feed SRTT back to sender/receiver report interval tuning
if let Some(srtt_ms) = mmp.metrics.srtt_ms() { if let Some(srtt_ms) = mmp.metrics.srtt_ms() {
@@ -114,7 +120,8 @@ impl Node {
// (what fraction of peer's frames we received), using per-interval deltas. // (what fraction of peer's frames we received), using per-interval deltas.
let our_recv_packets = mmp.receiver.cumulative_packets_recv(); let our_recv_packets = mmp.receiver.cumulative_packets_recv();
let peer_highest = mmp.receiver.highest_counter(); let peer_highest = mmp.receiver.highest_counter();
mmp.metrics.update_reverse_delivery(our_recv_packets, peer_highest); mmp.metrics
.update_reverse_delivery(our_recv_packets, peer_highest);
trace!( trace!(
from = %peer_name, from = %peer_name,
@@ -128,7 +135,9 @@ impl Node {
// Trigger re-evaluation so the node doesn't wait for the next // Trigger re-evaluation so the node doesn't wait for the next
// periodic tick or TreeAnnounce. // periodic tick or TreeAnnounce.
if first_rtt { if first_rtt {
let peer_costs: std::collections::HashMap<crate::NodeAddr, f64> = self.peers.iter() let peer_costs: std::collections::HashMap<crate::NodeAddr, f64> = self
.peers
.iter()
.filter(|(_, p)| p.has_srtt()) .filter(|(_, p)| p.has_srtt())
.map(|(a, p)| (*a, p.link_cost())) .map(|(a, p)| (*a, p.link_cost()))
.collect(); .collect();
@@ -179,7 +188,9 @@ impl Node {
for (node_addr, peer) in self.peers.iter_mut() { for (node_addr, peer) in self.peers.iter_mut() {
// Compute display name before taking mutable MMP borrow // Compute display name before taking mutable MMP borrow
let peer_name = self.peer_aliases.get(node_addr) let peer_name = self
.peer_aliases
.get(node_addr)
.cloned() .cloned()
.unwrap_or_else(|| peer.identity().short_npub()); .unwrap_or_else(|| peer.identity().short_npub());
@@ -261,7 +272,7 @@ impl Node {
let rtt_str = match m.srtt_ms() { let rtt_str = match m.srtt_ms() {
Some(rtt) => format!("{:.1}ms", rtt), Some(rtt) => format!("{:.1}ms", rtt),
None => "n/a".to_string() None => "n/a".to_string(),
}; };
let loss_str = format!("{:.1}%", m.loss_rate() * 100.0); let loss_str = format!("{:.1}%", m.loss_rate() * 100.0);
@@ -294,7 +305,9 @@ impl Node {
for (dest_addr, entry) in self.sessions.iter_mut() { for (dest_addr, entry) in self.sessions.iter_mut() {
// Compute display name before taking mutable MMP borrow // Compute display name before taking mutable MMP borrow
let session_name = self.peer_aliases.get(dest_addr) let session_name = self
.peer_aliases
.get(dest_addr)
.cloned() .cloned()
.unwrap_or_else(|| { .unwrap_or_else(|| {
let (xonly, _) = entry.remote_pubkey().x_only_public_key(); let (xonly, _) = entry.remote_pubkey().x_only_public_key();
@@ -362,7 +375,9 @@ impl Node {
} }
Err(e) => { Err(e) => {
// Peek at current failure count for log suppression // Peek at current failure count for log suppression
let failures = self.sessions.get(&dest_addr) let failures = self
.sessions
.get(&dest_addr)
.and_then(|entry| entry.mmp()) .and_then(|entry| entry.mmp())
.map(|mmp| mmp.sender.consecutive_send_failures()) .map(|mmp| mmp.sender.consecutive_send_failures())
.unwrap_or(0); .unwrap_or(0);
@@ -541,7 +556,10 @@ impl Node {
if let Some(peer) = self.peers.get_mut(&addr) { if let Some(peer) = self.peers.get_mut(&addr) {
peer.mark_heartbeat_sent(now); peer.mark_heartbeat_sent(now);
} }
if let Err(e) = self.send_encrypted_link_message(&addr, &heartbeat_msg).await { if let Err(e) = self
.send_encrypted_link_message(&addr, &heartbeat_msg)
.await
{
trace!(peer = %self.peer_display_name(&addr), error = %e, "Failed to send heartbeat"); trace!(peer = %self.peer_display_name(&addr), error = %e, "Failed to send heartbeat");
} }
} }

View File

@@ -5,11 +5,11 @@
//! 2. Drain window expiry (clean up previous session after cutover) //! 2. Drain window expiry (clean up previous session after cutover)
//! 3. Initiator-side cutover (first send after handshake completion) //! 3. Initiator-side cutover (first send after handshake completion)
use crate::NodeAddr;
use crate::node::Node; use crate::node::Node;
use crate::node::wire::build_msg1; use crate::node::wire::build_msg1;
use crate::noise::HandshakeState; use crate::noise::HandshakeState;
use crate::protocol::{SessionDatagram, SessionSetup}; use crate::protocol::{SessionDatagram, SessionSetup};
use crate::NodeAddr;
use tracing::{debug, info, trace, warn}; use tracing::{debug, info, trace, warn};
/// Keep previous session alive for this long after cutover. /// Keep previous session alive for this long after cutover.
@@ -69,7 +69,8 @@ impl Node {
} }
let elapsed = peer.session_established_at().elapsed().as_secs(); let elapsed = peer.session_established_at().elapsed().as_secs();
let counter = peer.noise_session() let counter = peer
.noise_session()
.map(|s| s.current_send_counter()) .map(|s| s.current_send_counter())
.unwrap_or(0); .unwrap_or(0);
@@ -88,9 +89,10 @@ impl Node {
debug_assert!( debug_assert!(
peer.transport_id().is_some() peer.transport_id().is_some()
&& peer.our_index().is_some() && peer.our_index().is_some()
&& self.peers_by_index.contains_key( && self.peers_by_index.contains_key(&(
&(peer.transport_id().unwrap(), peer.our_index().unwrap().as_u32()) peer.transport_id().unwrap(),
), peer.our_index().unwrap().as_u32()
)),
"peers_by_index should contain pre-registered new index after cutover" "peers_by_index should contain pre-registered new index after cutover"
); );
info!( info!(
@@ -106,7 +108,8 @@ impl Node {
&& let Some(old_our_index) = peer.complete_drain() && let Some(old_our_index) = peer.complete_drain()
{ {
if let Some(transport_id) = peer.transport_id() { if let Some(transport_id) = peer.transport_id() {
self.peers_by_index.remove(&(transport_id, old_our_index.as_u32())); self.peers_by_index
.remove(&(transport_id, old_our_index.as_u32()));
} }
let _ = self.index_allocator.free(old_our_index); let _ = self.index_allocator.free(old_our_index);
trace!( trace!(
@@ -208,7 +211,8 @@ impl Node {
} }
// Register in pending_outbound for msg2 dispatch (maps to existing link) // Register in pending_outbound for msg2 dispatch (maps to existing link)
self.pending_outbound.insert((transport_id, our_index.as_u32()), link_id); self.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
} }
/// Resend pending rekey msg1s and abandon timed-out rekeys. /// Resend pending rekey msg1s and abandon timed-out rekeys.

View File

@@ -1,10 +1,12 @@
//! RX event loop and packet dispatch. //! RX event loop and packet dispatch.
use crate::control::{commands, ControlSocket};
use crate::control::queries; use crate::control::queries;
use crate::control::{ControlSocket, commands};
use crate::node::wire::{
COMMON_PREFIX_SIZE, CommonPrefix, FMP_VERSION, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2,
};
use crate::node::{Node, NodeError}; use crate::node::{Node, NodeError};
use crate::transport::ReceivedPacket; use crate::transport::ReceivedPacket;
use crate::node::wire::{CommonPrefix, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2, FMP_VERSION, COMMON_PREFIX_SIZE};
use std::time::Duration; use std::time::Duration;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
@@ -29,8 +31,7 @@ impl Node {
/// This method takes ownership of the packet_rx channel and runs /// This method takes ownership of the packet_rx channel and runs
/// until the channel is closed (typically when stop() is called). /// until the channel is closed (typically when stop() is called).
pub async fn run_rx_loop(&mut self) -> Result<(), NodeError> { pub async fn run_rx_loop(&mut self) -> Result<(), NodeError> {
let mut packet_rx = self.packet_rx.take() let mut packet_rx = self.packet_rx.take().ok_or(NodeError::NotStarted)?;
.ok_or(NodeError::NotStarted)?;
// Take the TUN outbound receiver, or create a dummy channel that never // Take the TUN outbound receiver, or create a dummy channel that never
// produces messages (when TUN is disabled). Holding the sender prevents // produces messages (when TUN is disabled). Holding the sender prevents
@@ -53,12 +54,12 @@ impl Node {
} }
}; };
let mut tick = tokio::time::interval(Duration::from_secs(self.config.node.tick_interval_secs)); let mut tick =
tokio::time::interval(Duration::from_secs(self.config.node.tick_interval_secs));
// Set up control socket channel // Set up control socket channel
let (control_tx, mut control_rx) = tokio::sync::mpsc::channel::< let (control_tx, mut control_rx) =
crate::control::ControlMessage, tokio::sync::mpsc::channel::<crate::control::ControlMessage>(32);
>(32);
if self.config.node.control.enabled { if self.config.node.control.enabled {
let config = self.config.node.control.clone(); let config = self.config.node.control.clone();

View File

@@ -5,25 +5,27 @@
//! SessionSetup (Noise XK msg1), SessionAck (msg2), SessionMsg3 (msg3), //! SessionSetup (Noise XK msg1), SessionAck (msg2), SessionMsg3 (msg3),
//! encrypted data, and error signals (CoordsRequired, PathBroken). //! encrypted data, and error signals (CoordsRequired, PathBroken).
use crate::node::session::{EndToEndState, SessionEntry}; use crate::NodeAddr;
use crate::node::session_wire::{
build_fsp_header, fsp_prepend_inner_header, fsp_strip_inner_header,
parse_encrypted_coords, FspCommonPrefix, FspEncryptedHeader, FSP_COMMON_PREFIX_SIZE,
FSP_FLAG_CP, FSP_FLAG_K, FSP_HEADER_SIZE, FSP_PHASE_ESTABLISHED, FSP_PHASE_MSG1,
FSP_PHASE_MSG2, FSP_PHASE_MSG3, FSP_PORT_HEADER_SIZE, FSP_PORT_IPV6_SHIM,
};
use crate::protocol::{coords_wire_size, encode_coords};
use crate::upper::icmp::FIPS_OVERHEAD;
use crate::node::{Node, NodeError};
use crate::noise::{HandshakeState, XK_HANDSHAKE_MSG1_SIZE, XK_HANDSHAKE_MSG2_SIZE, XK_HANDSHAKE_MSG3_SIZE};
use crate::mmp::report::ReceiverReport; use crate::mmp::report::ReceiverReport;
use crate::mmp::{MAX_SESSION_REPORT_INTERVAL_MS, MIN_SESSION_REPORT_INTERVAL_MS}; use crate::mmp::{MAX_SESSION_REPORT_INTERVAL_MS, MIN_SESSION_REPORT_INTERVAL_MS};
use crate::node::session::{EndToEndState, SessionEntry};
use crate::node::session_wire::{
FSP_COMMON_PREFIX_SIZE, FSP_FLAG_CP, FSP_FLAG_K, FSP_HEADER_SIZE, FSP_PHASE_ESTABLISHED,
FSP_PHASE_MSG1, FSP_PHASE_MSG2, FSP_PHASE_MSG3, FSP_PORT_HEADER_SIZE, FSP_PORT_IPV6_SHIM,
FspCommonPrefix, FspEncryptedHeader, build_fsp_header, fsp_prepend_inner_header,
fsp_strip_inner_header, parse_encrypted_coords,
};
use crate::node::{Node, NodeError};
use crate::noise::{
HandshakeState, XK_HANDSHAKE_MSG1_SIZE, XK_HANDSHAKE_MSG2_SIZE, XK_HANDSHAKE_MSG3_SIZE,
};
use crate::protocol::{ use crate::protocol::{
CoordsRequired, FspInnerFlags, MtuExceeded, PathBroken, PathMtuNotification, SessionAck, CoordsRequired, FspInnerFlags, MtuExceeded, PathBroken, PathMtuNotification, SessionAck,
SessionDatagram, SessionMessageType, SessionMsg3, SessionReceiverReport, SessionSenderReport, SessionDatagram, SessionMessageType, SessionMsg3, SessionReceiverReport, SessionSenderReport,
SessionSetup, SessionSetup,
}; };
use crate::NodeAddr; use crate::protocol::{coords_wire_size, encode_coords};
use crate::upper::icmp::FIPS_OVERHEAD;
use secp256k1::PublicKey; use secp256k1::PublicKey;
use tracing::{debug, info, trace}; use tracing::{debug, info, trace};
@@ -48,7 +50,10 @@ impl Node {
let prefix = match FspCommonPrefix::parse(payload) { let prefix = match FspCommonPrefix::parse(payload) {
Some(p) => p, Some(p) => p,
None => { None => {
debug!(len = payload.len(), "Session payload too short for FSP prefix"); debug!(
len = payload.len(),
"Session payload too short for FSP prefix"
);
return; return;
} }
}; };
@@ -89,7 +94,8 @@ impl Node {
} }
} }
FSP_PHASE_ESTABLISHED => { FSP_PHASE_ESTABLISHED => {
self.handle_encrypted_session_msg(src_addr, payload, path_mtu, ce_flag).await; self.handle_encrypted_session_msg(src_addr, payload, path_mtu, ce_flag)
.await;
} }
_ => { _ => {
debug!(phase = prefix.phase, "Unknown FSP phase"); debug!(phase = prefix.phase, "Unknown FSP phase");
@@ -106,12 +112,21 @@ impl Node {
/// 4. AEAD decrypt with AAD = header_bytes /// 4. AEAD decrypt with AAD = header_bytes
/// 5. Strip FSP inner header → timestamp, msg_type, inner_flags /// 5. Strip FSP inner header → timestamp, msg_type, inner_flags
/// 6. Dispatch by msg_type /// 6. Dispatch by msg_type
async fn handle_encrypted_session_msg(&mut self, src_addr: &NodeAddr, payload: &[u8], path_mtu: u16, ce_flag: bool) { async fn handle_encrypted_session_msg(
&mut self,
src_addr: &NodeAddr,
payload: &[u8],
path_mtu: u16,
ce_flag: bool,
) {
// Parse the 12-byte encrypted header (includes the 4-byte prefix) // Parse the 12-byte encrypted header (includes the 4-byte prefix)
let header = match FspEncryptedHeader::parse(payload) { let header = match FspEncryptedHeader::parse(payload) {
Some(h) => h, Some(h) => h,
None => { None => {
debug!(len = payload.len(), "Encrypted session message too short for FSP header"); debug!(
len = payload.len(),
"Encrypted session message too short for FSP header"
);
return; return;
} }
}; };
@@ -166,8 +181,8 @@ impl Node {
let received_k_bit = header.flags & FSP_FLAG_K != 0; let received_k_bit = header.flags & FSP_FLAG_K != 0;
{ {
let entry = self.sessions.get(src_addr).unwrap(); let entry = self.sessions.get(src_addr).unwrap();
let k_bit_flipped = received_k_bit != entry.current_k_bit() let k_bit_flipped =
&& entry.pending_new_session().is_some(); received_k_bit != entry.current_k_bit() && entry.pending_new_session().is_some();
if k_bit_flipped { if k_bit_flipped {
let display_name = self.peer_display_name(src_addr); let display_name = self.peer_display_name(src_addr);
@@ -234,7 +249,8 @@ impl Node {
self.sessions.insert(*src_addr, entry); self.sessions.insert(*src_addr, entry);
// Strip FSP inner header (6 bytes) // Strip FSP inner header (6 bytes)
let (timestamp, msg_type, inner_flags_byte, rest) = match fsp_strip_inner_header(&plaintext) { let (timestamp, msg_type, inner_flags_byte, rest) = match fsp_strip_inner_header(&plaintext)
{
Some(parts) => parts, Some(parts) => parts,
None => { None => {
debug!(src = %self.peer_display_name(src_addr), "Decrypted payload too short for FSP inner header"); debug!(src = %self.peer_display_name(src_addr), "Decrypted payload too short for FSP inner header");
@@ -247,16 +263,15 @@ impl Node {
&& let Some(mmp) = entry.mmp_mut() && let Some(mmp) = entry.mmp_mut()
{ {
let now = std::time::Instant::now(); let now = std::time::Instant::now();
mmp.receiver.record_recv( mmp.receiver
header.counter, timestamp, plaintext.len(), ce_flag, now, .record_recv(header.counter, timestamp, plaintext.len(), ce_flag, now);
);
// Spin bit: advance state machine for correct TX reflection. // Spin bit: advance state machine for correct TX reflection.
// RTT samples not fed into SRTT — timestamp-echo provides // RTT samples not fed into SRTT — timestamp-echo provides
// accurate RTT; spin bit includes variable inter-frame delays. // accurate RTT; spin bit includes variable inter-frame delays.
let inner_flags = FspInnerFlags::from_byte(inner_flags_byte); let inner_flags = FspInnerFlags::from_byte(inner_flags_byte);
let _spin_rtt = mmp.spin_bit.rx_observe( let _spin_rtt = mmp
inner_flags.spin_bit, header.counter, now, .spin_bit
); .rx_observe(inner_flags.spin_bit, header.counter, now);
} }
// Feed path_mtu from datagram envelope to MMP path MTU tracking. // Feed path_mtu from datagram envelope to MMP path MTU tracking.
@@ -283,9 +298,15 @@ impl Node {
FSP_PORT_IPV6_SHIM => { FSP_PORT_IPV6_SHIM => {
use crate::FipsAddress; use crate::FipsAddress;
let src_ipv6 = FipsAddress::from_node_addr(src_addr).to_ipv6().octets(); let src_ipv6 = FipsAddress::from_node_addr(src_addr).to_ipv6().octets();
let dst_ipv6 = FipsAddress::from_node_addr(self.node_addr()).to_ipv6().octets(); let dst_ipv6 = FipsAddress::from_node_addr(self.node_addr())
.to_ipv6()
.octets();
match crate::upper::ipv6_shim::decompress_ipv6(service_payload, src_ipv6, dst_ipv6) { match crate::upper::ipv6_shim::decompress_ipv6(
service_payload,
src_ipv6,
dst_ipv6,
) {
Some(mut packet) => { Some(mut packet) => {
if ce_flag { if ce_flag {
mark_ipv6_ecn_ce(&mut packet); mark_ipv6_ecn_ce(&mut packet);
@@ -533,7 +554,13 @@ impl Node {
let placeholder_pubkey = self.identity.keypair().public_key(); let placeholder_pubkey = self.identity.keypair().public_key();
let now_ms = Self::now_ms(); let now_ms = Self::now_ms();
let resend_interval = self.config.node.rate_limit.handshake_resend_interval_ms; let resend_interval = self.config.node.rate_limit.handshake_resend_interval_ms;
let mut entry = SessionEntry::new(*src_addr, placeholder_pubkey, EndToEndState::AwaitingMsg3(handshake), now_ms, false); let mut entry = SessionEntry::new(
*src_addr,
placeholder_pubkey,
EndToEndState::AwaitingMsg3(handshake),
now_ms,
false,
);
entry.set_handshake_payload(ack_payload, now_ms + resend_interval); entry.set_handshake_payload(ack_payload, now_ms + resend_interval);
self.sessions.insert(*src_addr, entry); self.sessions.insert(*src_addr, entry);
@@ -809,7 +836,13 @@ impl Node {
let now_ms = Self::now_ms(); let now_ms = Self::now_ms();
// Replace the placeholder pubkey with the real one // Replace the placeholder pubkey with the real one
let mut new_entry = SessionEntry::new(*src_addr, remote_pubkey, EndToEndState::Established(session), now_ms, false); let mut new_entry = SessionEntry::new(
*src_addr,
remote_pubkey,
EndToEndState::Established(session),
now_ms,
false,
);
new_entry.set_coords_warmup_remaining(self.config.node.session.coords_warmup_packets); new_entry.set_coords_warmup_remaining(self.config.node.session.coords_warmup_packets);
new_entry.mark_established(now_ms); new_entry.mark_established(now_ms);
new_entry.init_mmp(&self.config.node.session_mmp); new_entry.init_mmp(&self.config.node.session_mmp);
@@ -878,7 +911,8 @@ impl Node {
}; };
let now = std::time::Instant::now(); let now = std::time::Instant::now();
mmp.metrics.process_receiver_report(&rr, our_timestamp_ms, now); mmp.metrics
.process_receiver_report(&rr, our_timestamp_ms, now);
// Feed SRTT back to sender/receiver report interval tuning (session-layer bounds) // Feed SRTT back to sender/receiver report interval tuning (session-layer bounds)
if let Some(srtt_ms) = mmp.metrics.srtt_ms() { if let Some(srtt_ms) = mmp.metrics.srtt_ms() {
@@ -900,7 +934,8 @@ impl Node {
// Update reverse delivery ratio from our own receiver state, using per-interval deltas. // Update reverse delivery ratio from our own receiver state, using per-interval deltas.
let our_recv_packets = mmp.receiver.cumulative_packets_recv(); let our_recv_packets = mmp.receiver.cumulative_packets_recv();
let peer_highest = mmp.receiver.highest_counter(); let peer_highest = mmp.receiver.highest_counter();
mmp.metrics.update_reverse_delivery(our_recv_packets, peer_highest); mmp.metrics
.update_reverse_delivery(our_recv_packets, peer_highest);
trace!( trace!(
src = %peer_name, src = %peer_name,
@@ -975,7 +1010,10 @@ impl Node {
); );
// Send standalone CoordsWarmup immediately (rate-limited) // Send standalone CoordsWarmup immediately (rate-limited)
if self.coords_response_rate_limiter.should_send(&msg.dest_addr) { if self
.coords_response_rate_limiter
.should_send(&msg.dest_addr)
{
if let Some(entry) = self.sessions.get(&msg.dest_addr) if let Some(entry) = self.sessions.get(&msg.dest_addr)
&& entry.is_established() && entry.is_established()
&& let Err(e) = self.send_coords_warmup(&msg.dest_addr).await && let Err(e) = self.send_coords_warmup(&msg.dest_addr).await
@@ -1033,7 +1071,10 @@ impl Node {
); );
// Send standalone CoordsWarmup immediately (rate-limited) // Send standalone CoordsWarmup immediately (rate-limited)
if self.coords_response_rate_limiter.should_send(&msg.dest_addr) { if self
.coords_response_rate_limiter
.should_send(&msg.dest_addr)
{
if let Some(entry) = self.sessions.get(&msg.dest_addr) if let Some(entry) = self.sessions.get(&msg.dest_addr)
&& entry.is_established() && entry.is_established()
&& let Err(e) = self.send_coords_warmup(&msg.dest_addr).await && let Err(e) = self.send_coords_warmup(&msg.dest_addr).await
@@ -1139,16 +1180,17 @@ impl Node {
let our_keypair = self.identity.keypair(); let our_keypair = self.identity.keypair();
let mut handshake = HandshakeState::new_xk_initiator(our_keypair, dest_pubkey); let mut handshake = HandshakeState::new_xk_initiator(our_keypair, dest_pubkey);
handshake.set_local_epoch(self.startup_epoch); handshake.set_local_epoch(self.startup_epoch);
let msg1 = handshake.write_xk_message_1().map_err(|e| NodeError::SendFailed { let msg1 = handshake
node_addr: dest_addr, .write_xk_message_1()
reason: format!("Noise XK msg1 generation failed: {}", e), .map_err(|e| NodeError::SendFailed {
})?; node_addr: dest_addr,
reason: format!("Noise XK msg1 generation failed: {}", e),
})?;
// Build SessionSetup with coordinates // Build SessionSetup with coordinates
let our_coords = self.tree_state.my_coords().clone(); let our_coords = self.tree_state.my_coords().clone();
let dest_coords = self.get_dest_coords(&dest_addr); let dest_coords = self.get_dest_coords(&dest_addr);
let setup = SessionSetup::new(our_coords, dest_coords) let setup = SessionSetup::new(our_coords, dest_coords).with_handshake(msg1);
.with_handshake(msg1);
let setup_payload = setup.encode(); let setup_payload = setup.encode();
// Wrap in SessionDatagram // Wrap in SessionDatagram
@@ -1165,7 +1207,13 @@ impl Node {
// Store session entry with handshake payload for potential resend // Store session entry with handshake payload for potential resend
let now_ms = Self::now_ms(); let now_ms = Self::now_ms();
let resend_interval = self.config.node.rate_limit.handshake_resend_interval_ms; let resend_interval = self.config.node.rate_limit.handshake_resend_interval_ms;
let mut entry = SessionEntry::new(dest_addr, dest_pubkey, EndToEndState::Initiating(handshake), now_ms, true); let mut entry = SessionEntry::new(
dest_addr,
dest_pubkey,
EndToEndState::Initiating(handshake),
now_ms,
true,
);
entry.set_handshake_payload(setup_payload, now_ms + resend_interval); entry.set_handshake_payload(setup_payload, now_ms + resend_interval);
self.sessions.insert(dest_addr, entry); self.sessions.insert(dest_addr, entry);
@@ -1192,10 +1240,13 @@ impl Node {
let now_ms = Self::now_ms(); let now_ms = Self::now_ms();
// First borrow: read session metadata (NLL releases before coord decision) // First borrow: read session metadata (NLL releases before coord decision)
let entry = self.sessions.get(dest_addr).ok_or_else(|| NodeError::SendFailed { let entry = self
node_addr: *dest_addr, .sessions
reason: "no session".into(), .get(dest_addr)
})?; .ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let wants_coords = entry.coords_warmup_remaining() > 0; let wants_coords = entry.coords_warmup_remaining() > 0;
let timestamp = entry.session_timestamp(now_ms); let timestamp = entry.session_timestamp(now_ms);
let spin_bit = entry.mmp().is_some_and(|m| m.spin_bit.tx_bit()); let spin_bit = entry.mmp().is_some_and(|m| m.spin_bit.tx_bit());
@@ -1215,7 +1266,8 @@ impl Node {
// Build inner plaintext (doesn't depend on counter) // Build inner plaintext (doesn't depend on counter)
let msg_type = SessionMessageType::DataPacket.to_byte(); // 0x10 let msg_type = SessionMessageType::DataPacket.to_byte(); // 0x10
let inner_flags = FspInnerFlags { spin_bit }.to_byte(); let inner_flags = FspInnerFlags { spin_bit }.to_byte();
let inner_plaintext = fsp_prepend_inner_header(timestamp, msg_type, inner_flags, &port_payload); let inner_plaintext =
fsp_prepend_inner_header(timestamp, msg_type, inner_flags, &port_payload);
// Determine whether coords fit within transport MTU. // Determine whether coords fit within transport MTU.
// If not, send standalone CoordsWarmup before the data packet. // If not, send standalone CoordsWarmup before the data packet.
@@ -1223,7 +1275,8 @@ impl Node {
let src = self.tree_state.my_coords().clone(); let src = self.tree_state.my_coords().clone();
let dst = self.get_dest_coords(dest_addr); let dst = self.get_dest_coords(dest_addr);
let coords_size = coords_wire_size(&src) + coords_wire_size(&dst); let coords_size = coords_wire_size(&src) + coords_wire_size(&dst);
let total_wire = FIPS_OVERHEAD as usize + FSP_PORT_HEADER_SIZE + coords_size + payload.len(); let total_wire =
FIPS_OVERHEAD as usize + FSP_PORT_HEADER_SIZE + coords_size + payload.len();
if total_wire <= self.transport_mtu() as usize { if total_wire <= self.transport_mtu() as usize {
(true, Some(src), Some(dst)) (true, Some(src), Some(dst))
} else { } else {
@@ -1239,9 +1292,7 @@ impl Node {
}; };
// Decrement warmup counter if we sent coords (piggybacked or standalone) // Decrement warmup counter if we sent coords (piggybacked or standalone)
if wants_coords if wants_coords && let Some(entry) = self.sessions.get_mut(dest_addr) {
&& let Some(entry) = self.sessions.get_mut(dest_addr)
{
entry.set_coords_warmup_remaining(entry.coords_warmup_remaining() - 1); entry.set_coords_warmup_remaining(entry.coords_warmup_remaining() - 1);
} }
@@ -1254,10 +1305,13 @@ impl Node {
} }
// Borrow session for counter + encryption (after potential standalone send) // Borrow session for counter + encryption (after potential standalone send)
let entry = self.sessions.get_mut(dest_addr).ok_or_else(|| NodeError::SendFailed { let entry = self
node_addr: *dest_addr, .sessions
reason: "no session".into(), .get_mut(dest_addr)
})?; .ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let session = match entry.state_mut() { let session = match entry.state_mut() {
EndToEndState::Established(s) => s, EndToEndState::Established(s) => s,
_ => { _ => {
@@ -1274,12 +1328,12 @@ impl Node {
let header = build_fsp_header(counter, flags, payload_len); let header = build_fsp_header(counter, flags, payload_len);
// Encrypt with AAD binding to the FSP header // Encrypt with AAD binding to the FSP header
let ciphertext = session.encrypt_with_aad(&inner_plaintext, &header).map_err(|e| { let ciphertext = session
NodeError::SendFailed { .encrypt_with_aad(&inner_plaintext, &header)
.map_err(|e| NodeError::SendFailed {
node_addr: *dest_addr, node_addr: *dest_addr,
reason: format!("session encrypt failed: {}", e), reason: format!("session encrypt failed: {}", e),
} })?;
})?;
// Assemble: header(12) + [coords] + ciphertext // Assemble: header(12) + [coords] + ciphertext
let mut fsp_payload = Vec::with_capacity(FSP_HEADER_SIZE + ciphertext.len() + 200); let mut fsp_payload = Vec::with_capacity(FSP_HEADER_SIZE + ciphertext.len() + 200);
@@ -1317,13 +1371,19 @@ impl Node {
dest_addr: &NodeAddr, dest_addr: &NodeAddr,
ipv6_packet: &[u8], ipv6_packet: &[u8],
) -> Result<(), NodeError> { ) -> Result<(), NodeError> {
let compressed = crate::upper::ipv6_shim::compress_ipv6(ipv6_packet) let compressed = crate::upper::ipv6_shim::compress_ipv6(ipv6_packet).ok_or_else(|| {
.ok_or_else(|| NodeError::SendFailed { NodeError::SendFailed {
node_addr: *dest_addr, node_addr: *dest_addr,
reason: "IPv6 header compression failed".into(), reason: "IPv6 header compression failed".into(),
})?; }
self.send_session_data(dest_addr, FSP_PORT_IPV6_SHIM, FSP_PORT_IPV6_SHIM, &compressed) })?;
.await self.send_session_data(
dest_addr,
FSP_PORT_IPV6_SHIM,
FSP_PORT_IPV6_SHIM,
&compressed,
)
.await
} }
/// Send a non-data session message (reports, notifications) over an established session. /// Send a non-data session message (reports, notifications) over an established session.
@@ -1342,10 +1402,13 @@ impl Node {
let now_ms = Self::now_ms(); let now_ms = Self::now_ms();
// Read spin bit and session timestamp from entry // Read spin bit and session timestamp from entry
let entry = self.sessions.get(dest_addr).ok_or_else(|| NodeError::SendFailed { let entry = self
node_addr: *dest_addr, .sessions
reason: "no session".into(), .get(dest_addr)
})?; .ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let timestamp = entry.session_timestamp(now_ms); let timestamp = entry.session_timestamp(now_ms);
let spin_bit = entry.mmp().is_some_and(|m| m.spin_bit.tx_bit()); let spin_bit = entry.mmp().is_some_and(|m| m.spin_bit.tx_bit());
@@ -1353,10 +1416,13 @@ impl Node {
let inner_flags = FspInnerFlags { spin_bit }.to_byte(); let inner_flags = FspInnerFlags { spin_bit }.to_byte();
// Get mutable access for encryption // Get mutable access for encryption
let entry = self.sessions.get_mut(dest_addr).ok_or_else(|| NodeError::SendFailed { let entry = self
node_addr: *dest_addr, .sessions
reason: "no session".into(), .get_mut(dest_addr)
})?; .ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
// Read K-bit before mutable borrow of session state // Read K-bit before mutable borrow of session state
let k_flags = if entry.current_k_bit() { FSP_FLAG_K } else { 0 }; let k_flags = if entry.current_k_bit() { FSP_FLAG_K } else { 0 };
@@ -1381,12 +1447,12 @@ impl Node {
let header = build_fsp_header(counter, k_flags, payload_len); let header = build_fsp_header(counter, k_flags, payload_len);
// Encrypt with AAD // Encrypt with AAD
let ciphertext = session.encrypt_with_aad(&inner_plaintext, &header).map_err(|e| { let ciphertext = session
NodeError::SendFailed { .encrypt_with_aad(&inner_plaintext, &header)
.map_err(|e| NodeError::SendFailed {
node_addr: *dest_addr, node_addr: *dest_addr,
reason: format!("session encrypt failed: {}", e), reason: format!("session encrypt failed: {}", e),
} })?;
})?;
// Assemble: header(12) + ciphertext (no coords) // Assemble: header(12) + ciphertext (no coords)
let mut fsp_payload = Vec::with_capacity(FSP_HEADER_SIZE + ciphertext.len()); let mut fsp_payload = Vec::with_capacity(FSP_HEADER_SIZE + ciphertext.len());
@@ -1416,28 +1482,31 @@ impl Node {
/// coordinates via `try_warm_coord_cache()` (same as CP-flagged data /// coordinates via `try_warm_coord_cache()` (same as CP-flagged data
/// packets). The encrypted inner payload is the 6-byte inner header /// packets). The encrypted inner payload is the 6-byte inner header
/// with no application data. /// with no application data.
async fn send_coords_warmup( async fn send_coords_warmup(&mut self, dest_addr: &NodeAddr) -> Result<(), NodeError> {
&mut self,
dest_addr: &NodeAddr,
) -> Result<(), NodeError> {
let now_ms = Self::now_ms(); let now_ms = Self::now_ms();
let my_coords = self.tree_state.my_coords().clone(); let my_coords = self.tree_state.my_coords().clone();
let dest_coords = self.get_dest_coords(dest_addr); let dest_coords = self.get_dest_coords(dest_addr);
// Read session metadata // Read session metadata
let entry = self.sessions.get(dest_addr).ok_or_else(|| NodeError::SendFailed { let entry = self
node_addr: *dest_addr, .sessions
reason: "no session".into(), .get(dest_addr)
})?; .ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let timestamp = entry.session_timestamp(now_ms); let timestamp = entry.session_timestamp(now_ms);
let spin_bit = entry.mmp().is_some_and(|m| m.spin_bit.tx_bit()); let spin_bit = entry.mmp().is_some_and(|m| m.spin_bit.tx_bit());
// Get mutable access for encryption // Get mutable access for encryption
let entry = self.sessions.get_mut(dest_addr).ok_or_else(|| NodeError::SendFailed { let entry = self
node_addr: *dest_addr, .sessions
reason: "no session".into(), .get_mut(dest_addr)
})?; .ok_or_else(|| NodeError::SendFailed {
node_addr: *dest_addr,
reason: "no session".into(),
})?;
let session = match entry.state_mut() { let session = match entry.state_mut() {
EndToEndState::Established(s) => s, EndToEndState::Established(s) => s,
_ => { _ => {
@@ -1460,12 +1529,12 @@ impl Node {
let header = build_fsp_header(counter, FSP_FLAG_CP, payload_len); let header = build_fsp_header(counter, FSP_FLAG_CP, payload_len);
// Encrypt with AAD // Encrypt with AAD
let ciphertext = session.encrypt_with_aad(&inner_plaintext, &header).map_err(|e| { let ciphertext = session
NodeError::SendFailed { .encrypt_with_aad(&inner_plaintext, &header)
.map_err(|e| NodeError::SendFailed {
node_addr: *dest_addr, node_addr: *dest_addr,
reason: format!("session encrypt failed: {}", e), reason: format!("session encrypt failed: {}", e),
} })?;
})?;
// Assemble: header(12) + coords + ciphertext // Assemble: header(12) + coords + ciphertext
let coords_size = coords_wire_size(&my_coords) + coords_wire_size(&dest_coords); let coords_size = coords_wire_size(&my_coords) + coords_wire_size(&dest_coords);
@@ -1532,7 +1601,8 @@ impl Node {
} }
let encoded = datagram.encode(); let encoded = datagram.encode();
self.send_encrypted_link_message(&next_hop_addr, &encoded).await?; self.send_encrypted_link_message(&next_hop_addr, &encoded)
.await?;
self.stats_mut().forwarding.record_originated(encoded.len()); self.stats_mut().forwarding.record_originated(encoded.len());
Ok(()) Ok(())
} }
@@ -1639,19 +1709,20 @@ impl Node {
/// Send ICMPv6 Destination Unreachable back through TUN. /// Send ICMPv6 Destination Unreachable back through TUN.
pub(in crate::node) fn send_icmpv6_dest_unreachable(&self, original_packet: &[u8]) { pub(in crate::node) fn send_icmpv6_dest_unreachable(&self, original_packet: &[u8]) {
use crate::upper::icmp::{build_dest_unreachable, should_send_icmp_error, DestUnreachableCode};
use crate::FipsAddress; use crate::FipsAddress;
use crate::upper::icmp::{
DestUnreachableCode, build_dest_unreachable, should_send_icmp_error,
};
if !should_send_icmp_error(original_packet) { if !should_send_icmp_error(original_packet) {
return; return;
} }
let our_ipv6 = FipsAddress::from_node_addr(self.node_addr()).to_ipv6(); let our_ipv6 = FipsAddress::from_node_addr(self.node_addr()).to_ipv6();
if let Some(response) = build_dest_unreachable( if let Some(response) =
original_packet, build_dest_unreachable(original_packet, DestUnreachableCode::NoRoute, our_ipv6)
DestUnreachableCode::NoRoute, && let Some(tun_tx) = &self.tun_tx
our_ipv6, {
) && let Some(tun_tx) = &self.tun_tx {
let _ = tun_tx.send(response); let _ = tun_tx.send(response);
} }
} }
@@ -1708,10 +1779,7 @@ impl Node {
return; return;
} }
let queue = self let queue = self.pending_tun_packets.entry(dest_addr).or_default();
.pending_tun_packets
.entry(dest_addr)
.or_default();
if queue.len() >= self.config.node.session.pending_packets_per_dest { if queue.len() >= self.config.node.session.pending_packets_per_dest {
queue.pop_front(); // Drop oldest queue.pop_front(); // Drop oldest
} }

View File

@@ -22,7 +22,9 @@ impl Node {
.unwrap_or(0); .unwrap_or(0);
let timeout_ms = self.config.node.rate_limit.handshake_timeout_secs * 1000; let timeout_ms = self.config.node.rate_limit.handshake_timeout_secs * 1000;
let stale: Vec<LinkId> = self.connections.iter() let stale: Vec<LinkId> = self
.connections
.iter()
.filter(|(_, conn)| conn.is_timed_out(now_ms, timeout_ms) || conn.is_failed()) .filter(|(_, conn)| conn.is_timed_out(now_ms, timeout_ms) || conn.is_failed())
.map(|(link_id, _)| *link_id) .map(|(link_id, _)| *link_id)
.collect(); .collect();
@@ -96,7 +98,9 @@ impl Node {
// Collect resend candidates: outbound, in SentMsg1, with stored msg1, // Collect resend candidates: outbound, in SentMsg1, with stored msg1,
// under max resends, and past the scheduled time. // under max resends, and past the scheduled time.
let candidates: Vec<(LinkId, Vec<u8>)> = self.connections.iter() let candidates: Vec<(LinkId, Vec<u8>)> = self
.connections
.iter()
.filter(|(_, conn)| { .filter(|(_, conn)| {
conn.is_outbound() conn.is_outbound()
&& conn.handshake_state() == HandshakeState::SentMsg1 && conn.handshake_state() == HandshakeState::SentMsg1
@@ -136,9 +140,7 @@ impl Node {
false false
}; };
if sent if sent && let Some(conn) = self.connections.get_mut(&link_id) {
&& let Some(conn) = self.connections.get_mut(&link_id)
{
let count = conn.resend_count() + 1; let count = conn.resend_count() + 1;
let next = now_ms + (interval_ms as f64 * backoff.powi(count as i32)) as u64; let next = now_ms + (interval_ms as f64 * backoff.powi(count as i32)) as u64;
conn.record_resend(next); conn.record_resend(next);
@@ -169,10 +171,11 @@ impl Node {
let ttl = self.config.node.session.default_ttl; let ttl = self.config.node.session.default_ttl;
// First pass: find timed-out sessions to remove // First pass: find timed-out sessions to remove
let timed_out: Vec<crate::NodeAddr> = self.sessions.iter() let timed_out: Vec<crate::NodeAddr> = self
.sessions
.iter()
.filter(|(_, entry)| { .filter(|(_, entry)| {
!entry.is_established() !entry.is_established() && now_ms.saturating_sub(entry.last_activity()) > timeout_ms
&& now_ms.saturating_sub(entry.last_activity()) > timeout_ms
}) })
.map(|(addr, _)| *addr) .map(|(addr, _)| *addr)
.collect(); .collect();
@@ -186,7 +189,9 @@ impl Node {
// Second pass: collect resend candidates // Second pass: collect resend candidates
let my_addr = *self.node_addr(); let my_addr = *self.node_addr();
let candidates: Vec<(crate::NodeAddr, Vec<u8>)> = self.sessions.iter() let candidates: Vec<(crate::NodeAddr, Vec<u8>)> = self
.sessions
.iter()
.filter(|(_, entry)| { .filter(|(_, entry)| {
!entry.is_established() !entry.is_established()
&& entry.handshake_payload().is_some() && entry.handshake_payload().is_some()
@@ -200,8 +205,7 @@ impl Node {
for (dest_addr, payload) in candidates { for (dest_addr, payload) in candidates {
use crate::protocol::SessionDatagram; use crate::protocol::SessionDatagram;
let mut datagram = SessionDatagram::new(my_addr, dest_addr, payload) let mut datagram = SessionDatagram::new(my_addr, dest_addr, payload).with_ttl(ttl);
.with_ttl(ttl);
let sent = match self.send_session_datagram(&mut datagram).await { let sent = match self.send_session_datagram(&mut datagram).await {
Ok(_) => true, Ok(_) => true,
Err(e) => { Err(e) => {
@@ -214,9 +218,7 @@ impl Node {
} }
}; };
if sent if sent && let Some(entry) = self.sessions.get_mut(&dest_addr) {
&& let Some(entry) = self.sessions.get_mut(&dest_addr)
{
let count = entry.resend_count() + 1; let count = entry.resend_count() + 1;
let next = now_ms + (interval_ms as f64 * backoff.powi(count as i32)) as u64; let next = now_ms + (interval_ms as f64 * backoff.powi(count as i32)) as u64;
entry.record_resend(next); entry.record_resend(next);
@@ -239,10 +241,11 @@ impl Node {
return; // disabled return; // disabled
} }
let idle: Vec<_> = self.sessions.iter() let idle: Vec<_> = self
.sessions
.iter()
.filter(|(_, entry)| { .filter(|(_, entry)| {
entry.is_established() entry.is_established() && now_ms.saturating_sub(entry.last_activity()) > timeout_ms
&& now_ms.saturating_sub(entry.last_activity()) > timeout_ms
}) })
.map(|(addr, _)| *addr) .map(|(addr, _)| *addr)
.collect(); .collect();

View File

@@ -1,11 +1,11 @@
//! Node lifecycle management: start, stop, and peer connection initiation. //! Node lifecycle management: start, stop, and peer connection initiation.
use super::{Node, NodeError, NodeState}; use super::{Node, NodeError, NodeState};
use crate::node::wire::build_msg1;
use crate::peer::PeerConnection; use crate::peer::PeerConnection;
use crate::protocol::{Disconnect, DisconnectReason}; use crate::protocol::{Disconnect, DisconnectReason};
use crate::transport::{packet_channel, Link, LinkDirection, LinkId, TransportAddr, TransportId}; use crate::transport::{Link, LinkDirection, LinkId, TransportAddr, TransportId, packet_channel};
use crate::upper::tun::{run_tun_reader, shutdown_tun_interface, TunDevice, TunState}; use crate::upper::tun::{TunDevice, TunState, run_tun_reader, shutdown_tun_interface};
use crate::node::wire::build_msg1;
use crate::{NodeAddr, PeerIdentity}; use crate::{NodeAddr, PeerIdentity};
use std::thread; use std::thread;
use std::time::Duration; use std::time::Duration;
@@ -50,7 +50,10 @@ impl Node {
return; return;
} }
info!(count = peer_configs.len(), "Initiating static peer connections"); info!(
count = peer_configs.len(),
"Initiating static peer connections"
);
for peer_config in peer_configs { for peer_config in peer_configs {
if let Err(e) = self.initiate_peer_connection(&peer_config).await { if let Err(e) = self.initiate_peer_connection(&peer_config).await {
@@ -67,14 +70,16 @@ impl Node {
/// Initiate a connection to a single peer. /// Initiate a connection to a single peer.
/// ///
/// Creates a link, starts the Noise handshake, and sends the first message. /// Creates a link, starts the Noise handshake, and sends the first message.
pub(super) async fn initiate_peer_connection(&mut self, peer_config: &crate::config::PeerConfig) -> Result<(), NodeError> { pub(super) async fn initiate_peer_connection(
&mut self,
peer_config: &crate::config::PeerConfig,
) -> Result<(), NodeError> {
// Parse the peer's npub to get their identity // Parse the peer's npub to get their identity
let peer_identity = PeerIdentity::from_npub(&peer_config.npub).map_err(|e| { let peer_identity =
NodeError::InvalidPeerNpub { PeerIdentity::from_npub(&peer_config.npub).map_err(|e| NodeError::InvalidPeerNpub {
npub: peer_config.npub.clone(), npub: peer_config.npub.clone(),
reason: e.to_string(), reason: e.to_string(),
} })?;
})?;
let peer_node_addr = *peer_identity.node_addr(); let peer_node_addr = *peer_identity.node_addr();
@@ -134,7 +139,10 @@ impl Node {
(tid, TransportAddr::from_string(&addr.addr)) (tid, TransportAddr::from_string(&addr.addr))
}; };
match self.initiate_connection(transport_id, remote_addr, peer_identity).await { match self
.initiate_connection(transport_id, remote_addr, peer_identity)
.await
{
Ok(()) => return Ok(()), Ok(()) => return Ok(()),
Err(e) => { Err(e) => {
debug!( debug!(
@@ -173,7 +181,9 @@ impl Node {
) -> Result<(), NodeError> { ) -> Result<(), NodeError> {
let peer_node_addr = *peer_identity.node_addr(); let peer_node_addr = *peer_identity.node_addr();
let is_connection_oriented = self.transports.get(&transport_id) let is_connection_oriented = self
.transports
.get(&transport_id)
.map(|t| t.transport_type().connection_oriented) .map(|t| t.transport_type().connection_oriented)
.unwrap_or(false); .unwrap_or(false);
@@ -234,7 +244,8 @@ impl Node {
Ok(()) Ok(())
} else { } else {
// Connectionless: proceed with immediate handshake // Connectionless: proceed with immediate handshake
self.start_handshake(link_id, transport_id, remote_addr, peer_identity).await self.start_handshake(link_id, transport_id, remote_addr, peer_identity)
.await
} }
} }
@@ -271,16 +282,17 @@ impl Node {
// Start the Noise handshake and get message 1 // Start the Noise handshake and get message 1
let our_keypair = self.identity.keypair(); let our_keypair = self.identity.keypair();
let noise_msg1 = match connection.start_handshake(our_keypair, self.startup_epoch, current_time_ms) { let noise_msg1 =
Ok(msg) => msg, match connection.start_handshake(our_keypair, self.startup_epoch, current_time_ms) {
Err(e) => { Ok(msg) => msg,
// Clean up the index and link Err(e) => {
let _ = self.index_allocator.free(our_index); // Clean up the index and link
self.links.remove(&link_id); let _ = self.index_allocator.free(our_index);
self.addr_to_link.remove(&(transport_id, remote_addr)); self.links.remove(&link_id);
return Err(NodeError::HandshakeFailed(e.to_string())); self.addr_to_link.remove(&(transport_id, remote_addr));
} return Err(NodeError::HandshakeFailed(e.to_string()));
}; }
};
// Set index and transport info on the connection // Set index and transport info on the connection
connection.set_our_index(our_index); connection.set_our_index(our_index);
@@ -304,7 +316,8 @@ impl Node {
connection.set_handshake_msg1(wire_msg1.clone(), current_time_ms + resend_interval); connection.set_handshake_msg1(wire_msg1.clone(), current_time_ms + resend_interval);
// Track in pending_outbound for msg2 dispatch // Track in pending_outbound for msg2 dispatch
self.pending_outbound.insert((transport_id, our_index.as_u32()), link_id); self.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
self.connections.insert(link_id, connection); self.connections.insert(link_id, connection);
// Send the wire format handshake message // Send the wire format handshake message
@@ -395,7 +408,10 @@ impl Node {
remote_addr = %remote_addr, remote_addr = %remote_addr,
"Auto-connecting to discovered peer" "Auto-connecting to discovered peer"
); );
if let Err(e) = self.initiate_connection(transport_id, remote_addr, identity).await { if let Err(e) = self
.initiate_connection(transport_id, remote_addr, identity)
.await
{
warn!(error = %e, "Failed to auto-connect to discovered peer"); warn!(error = %e, "Failed to auto-connect to discovered peer");
} }
} }
@@ -457,12 +473,15 @@ impl Node {
); );
// Start the handshake now that the transport is connected // Start the handshake now that the transport is connected
if let Err(e) = self.start_handshake( if let Err(e) = self
pending.link_id, .start_handshake(
pending.transport_id, pending.link_id,
pending.remote_addr.clone(), pending.transport_id,
pending.peer_identity, pending.remote_addr.clone(),
).await { pending.peer_identity,
)
.await
{
warn!( warn!(
link_id = %pending.link_id, link_id = %pending.link_id,
error = %e, error = %e,
@@ -559,7 +578,7 @@ impl Node {
// Calculate max MSS for TCP clamping // Calculate max MSS for TCP clamping
let effective_mtu = self.effective_ipv6_mtu(); let effective_mtu = self.effective_ipv6_mtu();
let max_mss = effective_mtu.saturating_sub(40).saturating_sub(20); // IPv6 + TCP headers let max_mss = effective_mtu.saturating_sub(40).saturating_sub(20); // IPv6 + TCP headers
info!("effective MTU: {} bytes", effective_mtu); info!("effective MTU: {} bytes", effective_mtu);
info!(" max TCP MSS: {} bytes", max_mss); info!(" max TCP MSS: {} bytes", max_mss);
@@ -581,7 +600,14 @@ impl Node {
// Spawn reader thread // Spawn reader thread
let transport_mtu = self.transport_mtu(); let transport_mtu = self.transport_mtu();
let reader_handle = thread::spawn(move || { let reader_handle = thread::spawn(move || {
run_tun_reader(device, mtu, our_addr, reader_tun_tx, outbound_tx, transport_mtu); run_tun_reader(
device,
mtu,
our_addr,
reader_tun_tx,
outbound_tx,
transport_mtu,
);
}); });
self.tun_state = TunState::Active; self.tun_state = TunState::Active;
@@ -606,11 +632,19 @@ impl Node {
let dns_channel_size = self.config.node.buffers.dns_channel; let dns_channel_size = self.config.node.buffers.dns_channel;
let (identity_tx, identity_rx) = tokio::sync::mpsc::channel(dns_channel_size); let (identity_tx, identity_rx) = tokio::sync::mpsc::channel(dns_channel_size);
let dns_ttl = self.config.dns.ttl(); let dns_ttl = self.config.dns.ttl();
let base_hosts = crate::upper::hosts::HostMap::from_peer_configs(self.config.peers()); let base_hosts =
let hosts_path = std::path::PathBuf::from(crate::upper::hosts::DEFAULT_HOSTS_PATH); crate::upper::hosts::HostMap::from_peer_configs(self.config.peers());
let reloader = crate::upper::hosts::HostMapReloader::new(base_hosts, hosts_path); let hosts_path =
std::path::PathBuf::from(crate::upper::hosts::DEFAULT_HOSTS_PATH);
let reloader =
crate::upper::hosts::HostMapReloader::new(base_hosts, hosts_path);
info!(bind = %bind, hosts = reloader.hosts().len(), "DNS responder started for .fips domain (auto-reload enabled)"); info!(bind = %bind, hosts = reloader.hosts().len(), "DNS responder started for .fips domain (auto-reload enabled)");
let handle = tokio::spawn(crate::upper::dns::run_dns_responder(socket, identity_tx, dns_ttl, reloader)); let handle = tokio::spawn(crate::upper::dns::run_dns_responder(
socket,
identity_tx,
dns_ttl,
reloader,
));
self.dns_identity_rx = Some(identity_rx); self.dns_identity_rx = Some(identity_rx);
self.dns_task = Some(handle); self.dns_task = Some(handle);
} }
@@ -646,7 +680,8 @@ impl Node {
} }
// Send disconnect notifications to all active peers before closing transports // Send disconnect notifications to all active peers before closing transports
self.send_disconnect_to_all_peers(DisconnectReason::Shutdown).await; self.send_disconnect_to_all_peers(DisconnectReason::Shutdown)
.await;
// Shutdown transports (they're packet producers) // Shutdown transports (they're packet producers)
let transport_ids: Vec<_> = self.transports.keys().cloned().collect(); let transport_ids: Vec<_> = self.transports.keys().cloned().collect();
@@ -710,7 +745,9 @@ impl Node {
let plaintext = disconnect.encode(); let plaintext = disconnect.encode();
// Collect node_addrs to avoid borrow conflict with send helper // Collect node_addrs to avoid borrow conflict with send helper
let peer_addrs: Vec<NodeAddr> = self.peers.iter() let peer_addrs: Vec<NodeAddr> = self
.peers
.iter()
.filter(|(_, peer)| peer.can_send() && peer.has_session()) .filter(|(_, peer)| peer.can_send() && peer.has_session())
.map(|(addr, _)| *addr) .map(|(addr, _)| *addr)
.collect(); .collect();
@@ -725,7 +762,10 @@ impl Node {
let mut sent = 0usize; let mut sent = 0usize;
for node_addr in &peer_addrs { for node_addr in &peer_addrs {
match self.send_encrypted_link_message(node_addr, &plaintext).await { match self
.send_encrypted_link_message(node_addr, &plaintext)
.await
{
Ok(()) => sent += 1, Ok(()) => sent += 1,
Err(e) => { Err(e) => {
debug!( debug!(
@@ -790,8 +830,8 @@ impl Node {
/// ///
/// Removes the peer and suppresses auto-reconnect. /// Removes the peer and suppresses auto-reconnect.
pub(crate) fn api_disconnect(&mut self, npub: &str) -> Result<serde_json::Value, String> { pub(crate) fn api_disconnect(&mut self, npub: &str) -> Result<serde_json::Value, String> {
let peer_identity = PeerIdentity::from_npub(npub) let peer_identity =
.map_err(|e| format!("invalid npub '{npub}': {e}"))?; PeerIdentity::from_npub(npub).map_err(|e| format!("invalid npub '{npub}': {e}"))?;
let node_addr = *peer_identity.node_addr(); let node_addr = *peer_identity.node_addr();
if !self.peers.contains_key(&node_addr) { if !self.peers.contains_key(&node_addr) {

View File

@@ -5,41 +5,44 @@
//! Bloom filters, coordinate caches, transports, links, and peers. //! Bloom filters, coordinate caches, transports, links, and peers.
mod bloom; mod bloom;
mod discovery_rate_limit;
mod handlers; mod handlers;
mod lifecycle; mod lifecycle;
mod retry;
mod discovery_rate_limit;
mod rate_limit; mod rate_limit;
mod retry;
mod routing_error_rate_limit; mod routing_error_rate_limit;
pub(crate) mod session; pub(crate) mod session;
pub(crate) mod session_wire; pub(crate) mod session_wire;
pub(crate) mod wire;
pub(crate) mod stats; pub(crate) mod stats;
mod tree;
#[cfg(test)] #[cfg(test)]
mod tests; mod tests;
mod tree;
pub(crate) mod wire;
use crate::bloom::BloomState;
use crate::cache::CoordCache;
use crate::utils::index::IndexAllocator;
use crate::node::session::SessionEntry;
use crate::peer::{ActivePeer, PeerConnection};
use self::discovery_rate_limit::{DiscoveryBackoff, DiscoveryForwardRateLimiter}; use self::discovery_rate_limit::{DiscoveryBackoff, DiscoveryForwardRateLimiter};
use self::rate_limit::HandshakeRateLimiter; use self::rate_limit::HandshakeRateLimiter;
use self::routing_error_rate_limit::RoutingErrorRateLimiter; use self::routing_error_rate_limit::RoutingErrorRateLimiter;
use self::wire::{
FLAG_CE, FLAG_KEY_EPOCH, FLAG_SP, build_encrypted, build_established_header,
prepend_inner_header,
};
use crate::bloom::BloomState;
use crate::cache::CoordCache;
use crate::node::session::SessionEntry;
use crate::peer::{ActivePeer, PeerConnection};
#[cfg(target_os = "linux")]
use crate::transport::ethernet::EthernetTransport;
use crate::transport::tcp::TcpTransport;
use crate::transport::tor::TorTransport;
use crate::transport::udp::UdpTransport;
use crate::transport::{ use crate::transport::{
Link, LinkId, PacketRx, PacketTx, TransportAddr, TransportError, TransportHandle, TransportId, Link, LinkId, PacketRx, PacketTx, TransportAddr, TransportError, TransportHandle, TransportId,
}; };
use crate::transport::udp::UdpTransport;
use crate::transport::tcp::TcpTransport;
use crate::transport::tor::TorTransport;
#[cfg(target_os = "linux")]
use crate::transport::ethernet::EthernetTransport;
use crate::tree::TreeState; use crate::tree::TreeState;
use crate::upper::hosts::HostMap; use crate::upper::hosts::HostMap;
use crate::upper::icmp_rate_limit::IcmpRateLimiter; use crate::upper::icmp_rate_limit::IcmpRateLimiter;
use crate::upper::tun::{TunError, TunOutboundRx, TunState, TunTx}; use crate::upper::tun::{TunError, TunOutboundRx, TunState, TunTx};
use self::wire::{build_encrypted, build_established_header, prepend_inner_header, FLAG_CE, FLAG_KEY_EPOCH, FLAG_SP}; use crate::utils::index::IndexAllocator;
use crate::{Config, ConfigError, Identity, IdentityError, NodeAddr, PeerIdentity}; use crate::{Config, ConfigError, Identity, IdentityError, NodeAddr, PeerIdentity};
use rand::Rng; use rand::Rng;
use std::collections::{HashMap, VecDeque}; use std::collections::{HashMap, VecDeque};
@@ -106,7 +109,11 @@ pub enum NodeError {
SendFailed { node_addr: NodeAddr, reason: String }, SendFailed { node_addr: NodeAddr, reason: String },
#[error("mtu exceeded forwarding to {node_addr}: packet {packet_size} > mtu {mtu}")] #[error("mtu exceeded forwarding to {node_addr}: packet {packet_size} > mtu {mtu}")]
MtuExceeded { node_addr: NodeAddr, packet_size: usize, mtu: u16 }, MtuExceeded {
node_addr: NodeAddr,
packet_size: usize,
mtu: u16,
},
#[error("config error: {0}")] #[error("config error: {0}")]
Config(#[from] ConfigError), Config(#[from] ConfigError),
@@ -541,10 +548,7 @@ impl Node {
coords_response_rate_limiter: RoutingErrorRateLimiter::with_interval( coords_response_rate_limiter: RoutingErrorRateLimiter::with_interval(
std::time::Duration::from_millis(coords_response_interval_ms), std::time::Duration::from_millis(coords_response_interval_ms),
), ),
discovery_backoff: DiscoveryBackoff::with_params( discovery_backoff: DiscoveryBackoff::with_params(backoff_base_secs, backoff_max_secs),
backoff_base_secs,
backoff_max_secs,
),
discovery_forward_limiter: DiscoveryForwardRateLimiter::with_interval( discovery_forward_limiter: DiscoveryForwardRateLimiter::with_interval(
std::time::Duration::from_secs(forward_min_interval_secs), std::time::Duration::from_secs(forward_min_interval_secs),
), ),
@@ -708,7 +712,8 @@ impl Node {
let xonly = self.identity.pubkey(); let xonly = self.identity.pubkey();
for (name, eth_config) in eth_instances { for (name, eth_config) in eth_instances {
let transport_id = self.allocate_transport_id(); let transport_id = self.allocate_transport_id();
let mut eth = EthernetTransport::new(transport_id, name, eth_config, packet_tx.clone()); let mut eth =
EthernetTransport::new(transport_id, name, eth_config, packet_tx.clone());
eth.set_local_pubkey(xonly); eth.set_local_pubkey(xonly);
transports.push(TransportHandle::Ethernet(eth)); transports.push(TransportHandle::Ethernet(eth));
} }
@@ -985,9 +990,10 @@ impl Node {
let now = std::time::Instant::now(); let now = std::time::Instant::now();
let should_log = match self.last_mesh_size_log { let should_log = match self.last_mesh_size_log {
None => true, None => true,
Some(last) => now.duration_since(last) >= std::time::Duration::from_secs( Some(last) => {
self.config.node.mmp.log_interval_secs, now.duration_since(last)
), >= std::time::Duration::from_secs(self.config.node.mmp.log_interval_secs)
}
}; };
if should_log { if should_log {
tracing::info!( tracing::info!(
@@ -1036,7 +1042,6 @@ impl Node {
self.tun_name.as_deref() self.tun_name.as_deref()
} }
// === Resource Limits === // === Resource Limits ===
/// Set the maximum number of connections (handshake phase). /// Set the maximum number of connections (handshake phase).
@@ -1117,14 +1122,17 @@ impl Node {
/// Add a link. /// Add a link.
pub fn add_link(&mut self, link: Link) -> Result<(), NodeError> { pub fn add_link(&mut self, link: Link) -> Result<(), NodeError> {
if self.max_links > 0 && self.links.len() >= self.max_links { if self.max_links > 0 && self.links.len() >= self.max_links {
return Err(NodeError::MaxLinksExceeded { max: self.max_links }); return Err(NodeError::MaxLinksExceeded {
max: self.max_links,
});
} }
let link_id = link.link_id(); let link_id = link.link_id();
let transport_id = link.transport_id(); let transport_id = link.transport_id();
let remote_addr = link.remote_addr().clone(); let remote_addr = link.remote_addr().clone();
self.links.insert(link_id, link); self.links.insert(link_id, link);
self.addr_to_link.insert((transport_id, remote_addr), link_id); self.addr_to_link
.insert((transport_id, remote_addr), link_id);
Ok(()) Ok(())
} }
@@ -1139,8 +1147,14 @@ impl Node {
} }
/// Find link ID by transport address. /// Find link ID by transport address.
pub fn find_link_by_addr(&self, transport_id: TransportId, addr: &TransportAddr) -> Option<LinkId> { pub fn find_link_by_addr(
self.addr_to_link.get(&(transport_id, addr.clone())).copied() &self,
transport_id: TransportId,
addr: &TransportAddr,
) -> Option<LinkId> {
self.addr_to_link
.get(&(transport_id, addr.clone()))
.copied()
} }
/// Remove a link. /// Remove a link.
@@ -1286,11 +1300,14 @@ impl Node {
pub(crate) fn register_identity(&mut self, node_addr: NodeAddr, pubkey: secp256k1::PublicKey) { pub(crate) fn register_identity(&mut self, node_addr: NodeAddr, pubkey: secp256k1::PublicKey) {
let mut prefix = [0u8; 15]; let mut prefix = [0u8; 15];
prefix.copy_from_slice(&node_addr.as_bytes()[0..15]); prefix.copy_from_slice(&node_addr.as_bytes()[0..15]);
self.identity_cache.insert(prefix, (node_addr, pubkey, Self::now_ms())); self.identity_cache
.insert(prefix, (node_addr, pubkey, Self::now_ms()));
// LRU eviction // LRU eviction
let max = self.config.node.cache.identity_size; let max = self.config.node.cache.identity_size;
if self.identity_cache.len() > max if self.identity_cache.len() > max
&& let Some(oldest_key) = self.identity_cache.iter() && let Some(oldest_key) = self
.identity_cache
.iter()
.min_by_key(|(_, (_, _, ts))| *ts) .min_by_key(|(_, (_, _, ts))| *ts)
.map(|(k, _)| *k) .map(|(k, _)| *k)
{ {
@@ -1299,7 +1316,10 @@ impl Node {
} }
/// Look up a destination by FipsAddress prefix (bytes 1-15 of the IPv6 address). /// Look up a destination by FipsAddress prefix (bytes 1-15 of the IPv6 address).
pub(crate) fn lookup_by_fips_prefix(&mut self, prefix: &[u8; 15]) -> Option<(NodeAddr, secp256k1::PublicKey)> { pub(crate) fn lookup_by_fips_prefix(
&mut self,
prefix: &[u8; 15],
) -> Option<(NodeAddr, secp256k1::PublicKey)> {
if let Some(entry) = self.identity_cache.get_mut(prefix) { if let Some(entry) = self.identity_cache.get_mut(prefix) {
entry.2 = Self::now_ms(); // LRU touch entry.2 = Self::now_ms(); // LRU touch
Some((entry.0, entry.1)) Some((entry.0, entry.1))
@@ -1338,9 +1358,7 @@ impl Node {
/// has declared us as their parent (making them our child). /// has declared us as their parent (making them our child).
pub(crate) fn is_tree_peer(&self, peer_addr: &NodeAddr) -> bool { pub(crate) fn is_tree_peer(&self, peer_addr: &NodeAddr) -> bool {
// Peer is our parent // Peer is our parent
if !self.tree_state.is_root() if !self.tree_state.is_root() && self.tree_state.my_declaration().parent_id() == peer_addr {
&& self.tree_state.my_declaration().parent_id() == peer_addr
{
return true; return true;
} }
// Peer is our child (their declaration names us as parent) // Peer is our child (their declaration names us as parent)
@@ -1389,7 +1407,10 @@ impl Node {
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis() as u64) .map(|d| d.as_millis() as u64)
.unwrap_or(0); .unwrap_or(0);
let dest_coords = self.coord_cache.get_and_touch(dest_node_addr, now_ms)?.clone(); let dest_coords = self
.coord_cache
.get_and_touch(dest_node_addr, now_ms)?
.clone();
// 3. Bloom filter candidates — requires dest_coords for loop-free selection. // 3. Bloom filter candidates — requires dest_coords for loop-free selection.
// If no candidate is strictly closer, fall through to tree routing. // If no candidate is strictly closer, fall through to tree routing.
@@ -1491,7 +1512,8 @@ impl Node {
node_addr: &NodeAddr, node_addr: &NodeAddr,
plaintext: &[u8], plaintext: &[u8],
) -> Result<(), NodeError> { ) -> Result<(), NodeError> {
self.send_encrypted_link_message_with_ce(node_addr, plaintext, false).await self.send_encrypted_link_message_with_ce(node_addr, plaintext, false)
.await
} }
/// Like `send_encrypted_link_message` but allows setting the FMP CE flag. /// Like `send_encrypted_link_message` but allows setting the FMP CE flag.
@@ -1503,7 +1525,9 @@ impl Node {
plaintext: &[u8], plaintext: &[u8],
ce_flag: bool, ce_flag: bool,
) -> Result<(), NodeError> { ) -> Result<(), NodeError> {
let peer = self.peers.get_mut(node_addr) let peer = self
.peers
.get_mut(node_addr)
.ok_or(NodeError::PeerNotFound(*node_addr))?; .ok_or(NodeError::PeerNotFound(*node_addr))?;
let their_index = peer.their_index().ok_or_else(|| NodeError::SendFailed { let their_index = peer.their_index().ok_or_else(|| NodeError::SendFailed {
@@ -1514,18 +1538,19 @@ impl Node {
node_addr: *node_addr, node_addr: *node_addr,
reason: "no transport_id".into(), reason: "no transport_id".into(),
})?; })?;
let remote_addr = peer.current_addr().cloned().ok_or_else(|| NodeError::SendFailed { let remote_addr = peer
node_addr: *node_addr, .current_addr()
reason: "no current_addr".into(), .cloned()
})?; .ok_or_else(|| NodeError::SendFailed {
node_addr: *node_addr,
reason: "no current_addr".into(),
})?;
// Prepend 4-byte session-relative timestamp (inner header) // Prepend 4-byte session-relative timestamp (inner header)
let timestamp_ms = peer.session_elapsed_ms(); let timestamp_ms = peer.session_elapsed_ms();
// MMP: read spin bit value before entering session borrow // MMP: read spin bit value before entering session borrow
let sp_flag = peer.mmp() let sp_flag = peer.mmp().map(|mmp| mmp.spin_bit.tx_bit()).unwrap_or(false);
.map(|mmp| mmp.spin_bit.tx_bit())
.unwrap_or(false);
let mut flags = if sp_flag { FLAG_SP } else { 0 }; let mut flags = if sp_flag { FLAG_SP } else { 0 };
if ce_flag { if ce_flag {
flags |= FLAG_CE; flags |= FLAG_CE;
@@ -1534,10 +1559,12 @@ impl Node {
flags |= FLAG_KEY_EPOCH; flags |= FLAG_KEY_EPOCH;
} }
let session = peer.noise_session_mut().ok_or_else(|| NodeError::SendFailed { let session = peer
node_addr: *node_addr, .noise_session_mut()
reason: "no noise session".into(), .ok_or_else(|| NodeError::SendFailed {
})?; node_addr: *node_addr,
reason: "no noise session".into(),
})?;
// Inner plaintext: [timestamp:4 LE][msg_type][payload...] // Inner plaintext: [timestamp:4 LE][msg_type][payload...]
let inner_plaintext = prepend_inner_header(timestamp_ms, plaintext); let inner_plaintext = prepend_inner_header(timestamp_ms, plaintext);
@@ -1548,18 +1575,24 @@ impl Node {
let header = build_established_header(their_index, counter, flags, payload_len); let header = build_established_header(their_index, counter, flags, payload_len);
// Encrypt with AAD binding to the outer header // Encrypt with AAD binding to the outer header
let ciphertext = session.encrypt_with_aad(&inner_plaintext, &header).map_err(|e| NodeError::SendFailed { let ciphertext = session
node_addr: *node_addr, .encrypt_with_aad(&inner_plaintext, &header)
reason: format!("encryption failed: {}", e), .map_err(|e| NodeError::SendFailed {
})?; node_addr: *node_addr,
reason: format!("encryption failed: {}", e),
})?;
let wire_packet = build_encrypted(&header, &ciphertext); let wire_packet = build_encrypted(&header, &ciphertext);
// Re-borrow peer for stats update after sending // Re-borrow peer for stats update after sending
let transport = self.transports.get(&transport_id) let transport = self
.transports
.get(&transport_id)
.ok_or(NodeError::TransportNotFound(transport_id))?; .ok_or(NodeError::TransportNotFound(transport_id))?;
let bytes_sent = transport.send(&remote_addr, &wire_packet).await let bytes_sent = transport
.send(&remote_addr, &wire_packet)
.await
.map_err(|e| match e { .map_err(|e| match e {
TransportError::MtuExceeded { packet_size, mtu } => NodeError::MtuExceeded { TransportError::MtuExceeded { packet_size, mtu } => NodeError::MtuExceeded {
node_addr: *node_addr, node_addr: *node_addr,

View File

@@ -237,7 +237,6 @@ impl HandshakeRateLimiter {
} }
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;

View File

@@ -5,9 +5,9 @@
//! (not PeerConnection) because each retry creates a fresh connection. //! (not PeerConnection) because each retry creates a fresh connection.
use super::Node; use super::Node;
use crate::PeerIdentity;
use crate::config::PeerConfig; use crate::config::PeerConfig;
use crate::identity::NodeAddr; use crate::identity::NodeAddr;
use crate::PeerIdentity;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
// MAX_BACKOFF_MS is now derived from config: node.retry.max_backoff_secs * 1000 // MAX_BACKOFF_MS is now derived from config: node.retry.max_backoff_secs * 1000
@@ -44,7 +44,9 @@ impl RetryState {
/// capped at `MAX_BACKOFF_MS`. /// capped at `MAX_BACKOFF_MS`.
pub fn backoff_ms(&self, base_interval_ms: u64, max_backoff_ms: u64) -> u64 { pub fn backoff_ms(&self, base_interval_ms: u64, max_backoff_ms: u64) -> u64 {
let multiplier = 1u64.checked_shl(self.retry_count).unwrap_or(u64::MAX); let multiplier = 1u64.checked_shl(self.retry_count).unwrap_or(u64::MAX);
base_interval_ms.saturating_mul(multiplier).min(max_backoff_ms) base_interval_ms
.saturating_mul(multiplier)
.min(max_backoff_ms)
} }
} }
@@ -55,11 +57,7 @@ impl Node {
/// have not been exhausted (unless `reconnect` is true, which retries /// have not been exhausted (unless `reconnect` is true, which retries
/// indefinitely). Does nothing if the peer is already connected or has /// indefinitely). Does nothing if the peer is already connected or has
/// a connection in progress. /// a connection in progress.
pub(super) fn schedule_retry( pub(super) fn schedule_retry(&mut self, node_addr: NodeAddr, now_ms: u64) {
&mut self,
node_addr: NodeAddr,
now_ms: u64,
) {
let retry_cfg = &self.config.node.retry; let retry_cfg = &self.config.node.retry;
let max_retries = retry_cfg.max_retries; let max_retries = retry_cfg.max_retries;
if max_retries == 0 { if max_retries == 0 {
@@ -240,8 +238,7 @@ impl Node {
// succeeds, promote_connection() clears retry_pending. If // succeeds, promote_connection() clears retry_pending. If
// it times out, check_timeouts() calls schedule_retry() // it times out, check_timeouts() calls schedule_retry()
// which bumps the counter and applies proper backoff. // which bumps the counter and applies proper backoff.
let hs_timeout_ms = let hs_timeout_ms = self.config.node.rate_limit.handshake_timeout_secs * 1000;
self.config.node.rate_limit.handshake_timeout_secs * 1000;
if let Some(state) = self.retry_pending.get_mut(&node_addr) { if let Some(state) = self.retry_pending.get_mut(&node_addr) {
state.retry_after_ms = now_ms + hs_timeout_ms; state.retry_after_ms = now_ms + hs_timeout_ms;
} }
@@ -317,7 +314,10 @@ mod tests {
retry_after_ms: 0, retry_after_ms: 0,
reconnect: false, reconnect: false,
}; };
assert_eq!(state.backoff_ms(5000, TEST_MAX_BACKOFF_MS), TEST_MAX_BACKOFF_MS); assert_eq!(
state.backoff_ms(5000, TEST_MAX_BACKOFF_MS),
TEST_MAX_BACKOFF_MS
);
} }
#[test] #[test]

View File

@@ -70,7 +70,6 @@ impl RoutingErrorRateLimiter {
pub fn len(&self) -> usize { pub fn len(&self) -> usize {
self.last_sent.len() self.last_sent.len()
} }
} }
impl Default for RoutingErrorRateLimiter { impl Default for RoutingErrorRateLimiter {

View File

@@ -7,10 +7,10 @@
use std::time::Instant; use std::time::Instant;
use crate::NodeAddr;
use crate::config::SessionMmpConfig; use crate::config::SessionMmpConfig;
use crate::mmp::MmpSessionState; use crate::mmp::MmpSessionState;
use crate::noise::{HandshakeState, NoiseSession}; use crate::noise::{HandshakeState, NoiseSession};
use crate::NodeAddr;
use secp256k1::PublicKey; use secp256k1::PublicKey;
/// State machine for an end-to-end session. /// State machine for an end-to-end session.
@@ -159,12 +159,16 @@ impl SessionEntry {
/// Get the current session state. /// Get the current session state.
#[cfg(test)] #[cfg(test)]
pub(crate) fn state(&self) -> &EndToEndState { pub(crate) fn state(&self) -> &EndToEndState {
self.state.as_ref().expect("session state taken but not restored") self.state
.as_ref()
.expect("session state taken but not restored")
} }
/// Get mutable access to the session state. /// Get mutable access to the session state.
pub(crate) fn state_mut(&mut self) -> &mut EndToEndState { pub(crate) fn state_mut(&mut self) -> &mut EndToEndState {
self.state.as_mut().expect("session state taken but not restored") self.state
.as_mut()
.expect("session state taken but not restored")
} }
/// Replace the session state. /// Replace the session state.
@@ -278,7 +282,12 @@ impl SessionEntry {
/// Get traffic counters: (packets_sent, packets_recv, bytes_sent, bytes_recv). /// Get traffic counters: (packets_sent, packets_recv, bytes_sent, bytes_recv).
pub(crate) fn traffic_counters(&self) -> (u64, u64, u64, u64) { pub(crate) fn traffic_counters(&self) -> (u64, u64, u64, u64) {
(self.packets_sent, self.packets_recv, self.bytes_sent, self.bytes_recv) (
self.packets_sent,
self.packets_recv,
self.bytes_sent,
self.bytes_recv,
)
} }
// === Handshake Resend === // === Handshake Resend ===

View File

@@ -208,8 +208,7 @@ impl FspEncryptedHeader {
let payload_len = u16::from_le_bytes([data[2], data[3]]); let payload_len = u16::from_le_bytes([data[2], data[3]]);
let counter = u64::from_le_bytes([ let counter = u64::from_le_bytes([
data[4], data[5], data[6], data[7], data[4], data[5], data[6], data[7], data[8], data[9], data[10], data[11],
data[8], data[9], data[10], data[11],
]); ]);
let mut header_bytes = [0u8; FSP_HEADER_SIZE]; let mut header_bytes = [0u8; FSP_HEADER_SIZE];
@@ -242,11 +241,7 @@ impl FspEncryptedHeader {
/// Build the 12-byte cleartext header for an encrypted FSP message. /// Build the 12-byte cleartext header for an encrypted FSP message.
/// ///
/// Returns the header bytes for use as AEAD AAD. /// Returns the header bytes for use as AEAD AAD.
pub fn build_fsp_header( pub fn build_fsp_header(counter: u64, flags: u8, payload_len: u16) -> [u8; FSP_HEADER_SIZE] {
counter: u64,
flags: u8,
payload_len: u16,
) -> [u8; FSP_HEADER_SIZE] {
let mut header = [0u8; FSP_HEADER_SIZE]; let mut header = [0u8; FSP_HEADER_SIZE];
header[0] = FspCommonPrefix::ver_phase_byte(FSP_VERSION, FSP_PHASE_ESTABLISHED); header[0] = FspCommonPrefix::ver_phase_byte(FSP_VERSION, FSP_PHASE_ESTABLISHED);
header[1] = flags; header[1] = flags;
@@ -323,12 +318,15 @@ pub fn fsp_strip_inner_header(plaintext: &[u8]) -> Option<(u32, u8, u8, &[u8])>
if plaintext.len() < FSP_INNER_HEADER_SIZE { if plaintext.len() < FSP_INNER_HEADER_SIZE {
return None; return None;
} }
let timestamp = u32::from_le_bytes([ let timestamp = u32::from_le_bytes([plaintext[0], plaintext[1], plaintext[2], plaintext[3]]);
plaintext[0], plaintext[1], plaintext[2], plaintext[3],
]);
let msg_type = plaintext[4]; let msg_type = plaintext[4];
let inner_flags = plaintext[5]; let inner_flags = plaintext[5];
Some((timestamp, msg_type, inner_flags, &plaintext[FSP_INNER_HEADER_SIZE..])) Some((
timestamp,
msg_type,
inner_flags,
&plaintext[FSP_INNER_HEADER_SIZE..],
))
} }
// ============================================================================ // ============================================================================
@@ -465,8 +463,8 @@ mod tests {
assert_eq!(u16::from_le_bytes([header[2], header[3]]), 200); assert_eq!(u16::from_le_bytes([header[2], header[3]]), 200);
assert_eq!( assert_eq!(
u64::from_le_bytes([ u64::from_le_bytes([
header[4], header[5], header[6], header[7], header[4], header[5], header[6], header[7], header[8], header[9], header[10],
header[8], header[9], header[10], header[11], header[11],
]), ]),
1000 1000
); );

View File

@@ -16,10 +16,7 @@ fn get_tree_edges(nodes: &[TestNode]) -> Vec<(usize, usize)> {
let ts = tn.node.tree_state(); let ts = tn.node.tree_state();
if !ts.is_root() { if !ts.is_root() {
let parent_addr = ts.my_declaration().parent_id(); let parent_addr = ts.my_declaration().parent_id();
if let Some(j) = nodes if let Some(j) = nodes.iter().position(|n| n.node.node_addr() == parent_addr) {
.iter()
.position(|n| n.node.node_addr() == parent_addr)
{
edges.push((i, j)); edges.push((i, j));
} }
} }
@@ -174,8 +171,7 @@ async fn test_bloom_filter_star() {
/// entries, and so on. Both endpoints should see all other nodes. /// entries, and so on. Both endpoints should see all other nodes.
#[tokio::test] #[tokio::test]
async fn test_bloom_filter_chain_propagation() { async fn test_bloom_filter_chain_propagation() {
let edges: Vec<(usize, usize)> = let edges: Vec<(usize, usize)> = vec![(0, 1), (1, 2), (2, 3), (3, 4), (4, 5), (5, 6), (6, 7)];
vec![(0, 1), (1, 2), (2, 3), (3, 4), (4, 5), (5, 6), (6, 7)];
let mut nodes = run_tree_test(8, &edges, false).await; let mut nodes = run_tree_test(8, &edges, false).await;
verify_tree_convergence(&nodes); verify_tree_convergence(&nodes);
verify_filter_exchange(&nodes, &edges); verify_filter_exchange(&nodes, &edges);
@@ -315,8 +311,7 @@ fn collect_subtree(
#[tokio::test] #[tokio::test]
async fn test_bloom_filter_split_horizon() { async fn test_bloom_filter_split_horizon() {
// Pure tree: 7 nodes, 6 edges // Pure tree: 7 nodes, 6 edges
let edges: Vec<(usize, usize)> = let edges: Vec<(usize, usize)> = vec![(0, 1), (0, 2), (1, 3), (1, 4), (2, 5), (5, 6)];
vec![(0, 1), (0, 2), (1, 3), (1, 4), (2, 5), (5, 6)];
let mut nodes = run_tree_test(7, &edges, false).await; let mut nodes = run_tree_test(7, &edges, false).await;
verify_tree_convergence(&nodes); verify_tree_convergence(&nodes);
verify_filter_exchange(&nodes, &edges); verify_filter_exchange(&nodes, &edges);
@@ -340,9 +335,7 @@ async fn test_bloom_filter_split_horizon() {
// - parent's filter to child contains the complement only // - parent's filter to child contains the complement only
for &(child_idx, parent_idx) in &tree_edges { for &(child_idx, parent_idx) in &tree_edges {
let child_subtree = collect_subtree(child_idx, Some(parent_idx), &tree_adj); let child_subtree = collect_subtree(child_idx, Some(parent_idx), &tree_adj);
let complement: Vec<usize> = (0..n) let complement: Vec<usize> = (0..n).filter(|i| !child_subtree.contains(i)).collect();
.filter(|i| !child_subtree.contains(i))
.collect();
// --- Upward filter: child → parent --- // --- Upward filter: child → parent ---
// This is stored as parent's inbound filter from child // This is stored as parent's inbound filter from child
@@ -358,7 +351,9 @@ async fn test_bloom_filter_split_horizon() {
assert!( assert!(
filter_up.contains(&addrs[idx]), filter_up.contains(&addrs[idx]),
"Upward filter (n{}→n{}): should contain subtree member n{} but doesn't", "Upward filter (n{}→n{}): should contain subtree member n{} but doesn't",
child_idx, parent_idx, idx child_idx,
parent_idx,
idx
); );
} }
@@ -367,7 +362,9 @@ async fn test_bloom_filter_split_horizon() {
assert!( assert!(
!filter_up.contains(&addrs[idx]), !filter_up.contains(&addrs[idx]),
"Upward filter (n{}→n{}): should NOT contain complement member n{} but does", "Upward filter (n{}→n{}): should NOT contain complement member n{} but does",
child_idx, parent_idx, idx child_idx,
parent_idx,
idx
); );
} }
@@ -376,7 +373,10 @@ async fn test_bloom_filter_split_horizon() {
assert!( assert!(
(up_est - child_subtree.len() as f64).abs() < 1.5, (up_est - child_subtree.len() as f64).abs() < 1.5,
"Upward filter (n{}→n{}): expected ~{} entries, got {:.1}", "Upward filter (n{}→n{}): expected ~{} entries, got {:.1}",
child_idx, parent_idx, child_subtree.len(), up_est child_idx,
parent_idx,
child_subtree.len(),
up_est
); );
// --- Downward filter: parent → child --- // --- Downward filter: parent → child ---
@@ -393,7 +393,9 @@ async fn test_bloom_filter_split_horizon() {
assert!( assert!(
filter_down.contains(&addrs[idx]), filter_down.contains(&addrs[idx]),
"Downward filter (n{}→n{}): should contain complement member n{} but doesn't", "Downward filter (n{}→n{}): should contain complement member n{} but doesn't",
parent_idx, child_idx, idx parent_idx,
child_idx,
idx
); );
} }
@@ -405,7 +407,9 @@ async fn test_bloom_filter_split_horizon() {
assert!( assert!(
!filter_down.contains(&addrs[idx]), !filter_down.contains(&addrs[idx]),
"Downward filter (n{}→n{}): should NOT contain subtree member n{} but does", "Downward filter (n{}→n{}): should NOT contain subtree member n{} but does",
parent_idx, child_idx, idx parent_idx,
child_idx,
idx
); );
} }
@@ -414,7 +418,10 @@ async fn test_bloom_filter_split_horizon() {
assert!( assert!(
(down_est - complement.len() as f64).abs() < 1.5, (down_est - complement.len() as f64).abs() < 1.5,
"Downward filter (n{}→n{}): expected ~{} entries, got {:.1}", "Downward filter (n{}→n{}): expected ~{} entries, got {:.1}",
parent_idx, child_idx, complement.len(), down_est parent_idx,
child_idx,
complement.len(),
down_est
); );
// Together, subtree + complement = all nodes // Together, subtree + complement = all nodes

View File

@@ -258,10 +258,8 @@ async fn test_disconnect_clears_session() {
{ {
let our_identity = nodes[1].node.identity(); let our_identity = nodes[1].node.identity();
let mut initiator = HandshakeState::new_initiator( let mut initiator =
our_identity.keypair(), HandshakeState::new_initiator(our_identity.keypair(), remote_identity.pubkey_full());
remote_identity.pubkey_full(),
);
let mut responder = HandshakeState::new_responder(remote_identity.keypair()); let mut responder = HandshakeState::new_responder(remote_identity.keypair());
let mut init_epoch = [0u8; 8]; let mut init_epoch = [0u8; 8];
rand::Rng::fill_bytes(&mut rand::rng(), &mut init_epoch); rand::Rng::fill_bytes(&mut rand::rng(), &mut init_epoch);
@@ -285,8 +283,16 @@ async fn test_disconnect_clears_session() {
nodes[1].node.sessions.insert(node0_addr, entry); nodes[1].node.sessions.insert(node0_addr, entry);
} }
assert_eq!(nodes[1].node.session_count(), 1, "Session should exist before disconnect"); assert_eq!(
assert_eq!(nodes[1].node.peer_count(), 1, "Peer should exist before disconnect"); nodes[1].node.session_count(),
1,
"Session should exist before disconnect"
);
assert_eq!(
nodes[1].node.peer_count(),
1,
"Peer should exist before disconnect"
);
// Node 0 sends Disconnect to node 1. // Node 0 sends Disconnect to node 1.
let disconnect = crate::protocol::Disconnect::new(DisconnectReason::Shutdown); let disconnect = crate::protocol::Disconnect::new(DisconnectReason::Shutdown);
@@ -301,7 +307,8 @@ async fn test_disconnect_clears_session() {
// Peer must be gone. // Peer must be gone.
assert_eq!( assert_eq!(
nodes[1].node.peer_count(), 0, nodes[1].node.peer_count(),
0,
"Peer should be removed after disconnect" "Peer should be removed after disconnect"
); );
@@ -309,7 +316,8 @@ async fn test_disconnect_clears_session() {
// Before the fix, session_count() would still be 1 here because // Before the fix, session_count() would still be 1 here because
// remove_active_peer didn't remove self.sessions[node0_addr]. // remove_active_peer didn't remove self.sessions[node0_addr].
assert_eq!( assert_eq!(
nodes[1].node.session_count(), 0, nodes[1].node.session_count(),
0,
"Session must be cleaned up when peer is removed (regression: issue #5)" "Session must be cleaned up when peer is removed (regression: issue #5)"
); );

View File

@@ -149,10 +149,8 @@ async fn test_response_transit_needs_recent_request() {
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.unwrap() .unwrap()
.as_millis() as u64; .as_millis() as u64;
node.recent_requests.insert( node.recent_requests
444, .insert(444, RecentRequest::new(make_node_addr(0xDD), now_ms));
RecentRequest::new(make_node_addr(0xDD), now_ms),
);
// Handle response — should try to reverse-path forward to 0xDD // Handle response — should try to reverse-path forward to 0xDD
// (will fail silently since 0xDD is not an actual peer) // (will fail silently since 0xDD is not an actual peer)
@@ -282,11 +280,7 @@ async fn test_response_coord_substitution_detected() {
let target = *target_identity.node_addr(); let target = *target_identity.node_addr();
let root = make_node_addr(0xF0); let root = make_node_addr(0xF0);
let real_coords = TreeCoordinate::from_addrs(vec![target, root]).unwrap(); let real_coords = TreeCoordinate::from_addrs(vec![target, root]).unwrap();
let fake_coords = TreeCoordinate::from_addrs(vec![ let fake_coords = TreeCoordinate::from_addrs(vec![target, make_node_addr(0xEE), root]).unwrap();
target,
make_node_addr(0xEE),
root,
]).unwrap();
// Register target in identity_cache // Register target in identity_cache
node.register_identity(target, target_identity.pubkey_full()); node.register_identity(target, target_identity.pubkey_full());
@@ -325,16 +319,12 @@ async fn test_recent_request_expiry() {
.as_millis() as u64; .as_millis() as u64;
// Insert an old request (11 seconds ago) // Insert an old request (11 seconds ago)
node.recent_requests.insert( node.recent_requests
123, .insert(123, RecentRequest::new(make_node_addr(1), now_ms - 11_000));
RecentRequest::new(make_node_addr(1), now_ms - 11_000),
);
// Insert a recent request // Insert a recent request
node.recent_requests.insert( node.recent_requests
456, .insert(456, RecentRequest::new(make_node_addr(2), now_ms));
RecentRequest::new(make_node_addr(2), now_ms),
);
assert_eq!(node.recent_requests.len(), 2); assert_eq!(node.recent_requests.len(), 2);
@@ -344,7 +334,8 @@ async fn test_recent_request_expiry() {
let coords = TreeCoordinate::from_addrs(vec![origin, make_node_addr(0)]).unwrap(); let coords = TreeCoordinate::from_addrs(vec![origin, make_node_addr(0)]).unwrap();
let request = LookupRequest::new(789, target, origin, coords, 3, 0); let request = LookupRequest::new(789, target, origin, coords, 3, 0);
let payload = &request.encode()[1..]; let payload = &request.encode()[1..];
node.handle_lookup_request(&make_node_addr(0xAA), payload).await; node.handle_lookup_request(&make_node_addr(0xAA), payload)
.await;
// Old entry (123) should be purged, recent entry (456) and new entry (789) kept // Old entry (123) should be purged, recent entry (456) and new entry (789) kept
assert!(!node.recent_requests.contains_key(&123)); assert!(!node.recent_requests.contains_key(&123));
@@ -381,7 +372,10 @@ async fn test_request_forwarding_two_node() {
// Process packets — node1 should receive the forwarded request // Process packets — node1 should receive the forwarded request
tokio::time::sleep(Duration::from_millis(50)).await; tokio::time::sleep(Duration::from_millis(50)).await;
let count = process_available_packets(&mut nodes).await; let count = process_available_packets(&mut nodes).await;
assert!(count > 0, "Expected forwarded LookupRequest to arrive at node 1"); assert!(
count > 0,
"Expected forwarded LookupRequest to arrive at node 1"
);
// Node1 should have recorded the request // Node1 should have recorded the request
assert!( assert!(
@@ -546,10 +540,7 @@ async fn test_discovery_100_nodes() {
} }
// Collect all node addresses and public keys for lookup targets // Collect all node addresses and public keys for lookup targets
let all_addrs: Vec<NodeAddr> = nodes let all_addrs: Vec<NodeAddr> = nodes.iter().map(|tn| *tn.node.node_addr()).collect();
.iter()
.map(|tn| *tn.node.node_addr())
.collect();
let all_pubkeys: Vec<secp256k1::PublicKey> = nodes let all_pubkeys: Vec<secp256k1::PublicKey> = nodes
.iter() .iter()
.map(|tn| tn.node.identity().pubkey_full()) .map(|tn| tn.node.identity().pubkey_full())
@@ -563,7 +554,8 @@ async fn test_discovery_100_nodes() {
if src == dst { if src == dst {
continue; continue;
} }
node.node.register_identity(all_addrs[dst], all_pubkeys[dst]); node.node
.register_identity(all_addrs[dst], all_pubkeys[dst]);
} }
} }
@@ -588,10 +580,7 @@ async fn test_discovery_100_nodes() {
let mut initiated = false; let mut initiated = false;
for &(s, dst) in &lookup_pairs { for &(s, dst) in &lookup_pairs {
if s == src { if s == src {
nodes[src] nodes[src].node.initiate_lookup(&all_addrs[dst], TTL).await;
.node
.initiate_lookup(&all_addrs[dst], TTL)
.await;
initiated = true; initiated = true;
} }
} }
@@ -628,7 +617,11 @@ async fn test_discovery_100_nodes() {
let mut failed_pairs: Vec<(usize, usize)> = Vec::new(); let mut failed_pairs: Vec<(usize, usize)> = Vec::new();
for &(src, dst) in &lookup_pairs { for &(src, dst) in &lookup_pairs {
if nodes[src].node.coord_cache().contains(&all_addrs[dst], now_ms) { if nodes[src]
.node
.coord_cache()
.contains(&all_addrs[dst], now_ms)
{
resolved += 1; resolved += 1;
} else { } else {
failed += 1; failed += 1;
@@ -638,9 +631,7 @@ async fn test_discovery_100_nodes() {
} }
} }
eprintln!( eprintln!("\n === Discovery 100-Node Test ===",);
"\n === Discovery 100-Node Test ===",
);
eprintln!( eprintln!(
" Lookups: {} | Resolved: {} | Failed: {} | Success rate: {:.1}%", " Lookups: {} | Resolved: {} | Failed: {} | Success rate: {:.1}%",
total_lookups, total_lookups,
@@ -651,8 +642,16 @@ async fn test_discovery_100_nodes() {
// Report coord_cache stats across all nodes // Report coord_cache stats across all nodes
let total_cached: usize = nodes.iter().map(|tn| tn.node.coord_cache().len()).sum(); let total_cached: usize = nodes.iter().map(|tn| tn.node.coord_cache().len()).sum();
let min_cached = nodes.iter().map(|tn| tn.node.coord_cache().len()).min().unwrap(); let min_cached = nodes
let max_cached = nodes.iter().map(|tn| tn.node.coord_cache().len()).max().unwrap(); .iter()
.map(|tn| tn.node.coord_cache().len())
.min()
.unwrap();
let max_cached = nodes
.iter()
.map(|tn| tn.node.coord_cache().len())
.max()
.unwrap();
eprintln!( eprintln!(
" Coord cache entries: total={} min={} max={} avg={:.1}", " Coord cache entries: total={} min={} max={} avg={:.1}",
total_cached, total_cached,
@@ -663,21 +662,32 @@ async fn test_discovery_100_nodes() {
// Detailed diagnostics for failures (to aid future debugging) // Detailed diagnostics for failures (to aid future debugging)
if !failed_pairs.is_empty() { if !failed_pairs.is_empty() {
eprintln!(" --- Failure Diagnostics ({} failures) ---", failed_pairs.len()); eprintln!(
" --- Failure Diagnostics ({} failures) ---",
failed_pairs.len()
);
for &(src, dst) in &failed_pairs { for &(src, dst) in &failed_pairs {
let src_coords = nodes[src].node.tree_state().my_coords().clone(); let src_coords = nodes[src].node.tree_state().my_coords().clone();
let dst_coords = nodes[dst].node.tree_state().my_coords().clone(); let dst_coords = nodes[dst].node.tree_state().my_coords().clone();
let tree_dist = src_coords.distance_to(&dst_coords); let tree_dist = src_coords.distance_to(&dst_coords);
let reverse_cached = nodes[dst].node.coord_cache().contains(&all_addrs[src], now_ms); let reverse_cached = nodes[dst]
.node
.coord_cache()
.contains(&all_addrs[src], now_ms);
let src_peers = nodes[src].node.peers.len(); let src_peers = nodes[src].node.peers.len();
let dst_peers = nodes[dst].node.peers.len(); let dst_peers = nodes[dst].node.peers.len();
eprintln!( eprintln!(
" node {} -> node {}: tree_dist={} src_depth={} dst_depth={} \ " node {} -> node {}: tree_dist={} src_depth={} dst_depth={} \
src_peers={} dst_peers={} reverse_cached={}", src_peers={} dst_peers={} reverse_cached={}",
src, dst, tree_dist, src,
src_coords.depth(), dst_coords.depth(), dst,
src_peers, dst_peers, reverse_cached tree_dist,
src_coords.depth(),
dst_coords.depth(),
src_peers,
dst_peers,
reverse_cached
); );
} }
} }
@@ -730,7 +740,9 @@ async fn test_response_path_mtu_two_node() {
// Check that path_mtu was stored in the cache entry // Check that path_mtu was stored in the cache entry
let entry = nodes[0].node.coord_cache().get_entry(&node1_addr).unwrap(); let entry = nodes[0].node.coord_cache().get_entry(&node1_addr).unwrap();
let path_mtu = entry.path_mtu().expect("path_mtu should be set from discovery"); let path_mtu = entry
.path_mtu()
.expect("path_mtu should be set from discovery");
// In a 2-node setup, no transit node applies the min() so path_mtu stays u16::MAX // In a 2-node setup, no transit node applies the min() so path_mtu stays u16::MAX
assert_eq!( assert_eq!(
path_mtu, path_mtu,
@@ -774,7 +786,9 @@ async fn test_response_path_mtu_three_node_chain() {
// Node1 is transit and applies min(u16::MAX, 1280) = 1280 // Node1 is transit and applies min(u16::MAX, 1280) = 1280
let entry = nodes[0].node.coord_cache().get_entry(&node2_addr).unwrap(); let entry = nodes[0].node.coord_cache().get_entry(&node2_addr).unwrap();
let path_mtu = entry.path_mtu().expect("path_mtu should be set from discovery"); let path_mtu = entry
.path_mtu()
.expect("path_mtu should be set from discovery");
assert_eq!( assert_eq!(
path_mtu, 1280, path_mtu, 1280,
"Three-node chain path_mtu should reflect transit node's transport MTU (1280)" "Three-node chain path_mtu should reflect transit node's transport MTU (1280)"
@@ -796,12 +810,8 @@ async fn test_cache_entry_path_mtu_stored() {
let coords = TreeCoordinate::from_addrs(vec![target, make_node_addr(0)]).unwrap(); let coords = TreeCoordinate::from_addrs(vec![target, make_node_addr(0)]).unwrap();
let now_ms = 1000u64; let now_ms = 1000u64;
node.coord_cache_mut().insert_with_path_mtu( node.coord_cache_mut()
target, .insert_with_path_mtu(target, coords, now_ms, 1280);
coords,
now_ms,
1280,
);
let entry = node.coord_cache().get_entry(&target).unwrap(); let entry = node.coord_cache().get_entry(&target).unwrap();
assert_eq!(entry.path_mtu(), Some(1280)); assert_eq!(entry.path_mtu(), Some(1280));

View File

@@ -6,8 +6,8 @@
use super::*; use super::*;
use crate::config::EthernetConfig; use crate::config::EthernetConfig;
use crate::transport::ethernet::EthernetTransport; use crate::transport::ethernet::EthernetTransport;
use crate::transport::{packet_channel, TransportAddr, TransportHandle, TransportId}; use crate::transport::{TransportAddr, TransportHandle, TransportId, packet_channel};
use spanning_tree::{cleanup_nodes, drain_all_packets, initiate_handshake, TestNode}; use spanning_tree::{TestNode, cleanup_nodes, drain_all_packets, initiate_handshake};
use std::process::Command; use std::process::Command;
use std::sync::atomic::{AtomicU32, Ordering}; use std::sync::atomic::{AtomicU32, Ordering};
@@ -40,7 +40,9 @@ impl VethPair {
// Create veth pair // Create veth pair
let status = Command::new("ip") let status = Command::new("ip")
.args(["link", "add", &name_a, "type", "veth", "peer", "name", &name_b]) .args([
"link", "add", &name_a, "type", "veth", "peer", "name", &name_b,
])
.status() .status()
.expect("failed to run 'ip link add'"); .expect("failed to run 'ip link add'");
assert!(status.success(), "failed to create veth pair"); assert!(status.success(), "failed to create veth pair");
@@ -91,7 +93,9 @@ async fn make_test_node_ethernet(interface: &str) -> TestNode {
let mut transport = EthernetTransport::new(transport_id, None, config, packet_tx); let mut transport = EthernetTransport::new(transport_id, None, config, packet_tx);
transport.start_async().await.unwrap(); transport.start_async().await.unwrap();
let mac = transport.local_mac().expect("transport should have MAC after start"); let mac = transport
.local_mac()
.expect("transport should have MAC after start");
let addr = TransportAddr::from_bytes(&mac); let addr = TransportAddr::from_bytes(&mac);
node.transports node.transports

View File

@@ -5,12 +5,11 @@
//! multi-hop forwarding through live node topologies. //! multi-hop forwarding through live node topologies.
use super::*; use super::*;
use crate::node::session_wire::{build_fsp_header, FSP_FLAG_CP}; use crate::node::session_wire::{FSP_FLAG_CP, build_fsp_header};
use crate::protocol::{SessionAck, SessionDatagram, SessionSetup, encode_coords}; use crate::protocol::{SessionAck, SessionDatagram, SessionSetup, encode_coords};
use crate::tree::TreeCoordinate; use crate::tree::TreeCoordinate;
use spanning_tree::{ use spanning_tree::{
cleanup_nodes, process_available_packets, run_tree_test, verify_tree_convergence, TestNode, cleanup_nodes, process_available_packets, run_tree_test, verify_tree_convergence,
TestNode,
}; };
// ============================================================================ // ============================================================================
@@ -35,11 +34,11 @@ async fn test_forwarding_hop_limit_exhausted() {
let from = make_node_addr(0xAA); let from = make_node_addr(0xAA);
let src = make_node_addr(0x01); let src = make_node_addr(0x01);
let dest = make_node_addr(0x02); let dest = make_node_addr(0x02);
let dg = SessionDatagram::new(src, dest, vec![0x10, 0x00, 0x00, 0x00]) let dg = SessionDatagram::new(src, dest, vec![0x10, 0x00, 0x00, 0x00]).with_ttl(0);
.with_ttl(0);
let encoded = dg.encode(); let encoded = dg.encode();
// Dispatch with payload after msg_type byte // Dispatch with payload after msg_type byte
node.handle_session_datagram(&from, &encoded[1..], false).await; node.handle_session_datagram(&from, &encoded[1..], false)
.await;
// No panic, no send (node has no peers) // No panic, no send (node has no peers)
} }
@@ -52,11 +51,11 @@ async fn test_forwarding_hop_limit_one_drops_at_transit() {
let from = make_node_addr(0xAA); let from = make_node_addr(0xAA);
let my_addr = *node.node_addr(); let my_addr = *node.node_addr();
let src = make_node_addr(0x01); let src = make_node_addr(0x01);
let dg = SessionDatagram::new(src, my_addr, vec![0x10, 0x00, 0x00, 0x00]) let dg = SessionDatagram::new(src, my_addr, vec![0x10, 0x00, 0x00, 0x00]).with_ttl(1);
.with_ttl(1);
let encoded = dg.encode(); let encoded = dg.encode();
// Should succeed — ttl=1 decrements to 0 but packet is still processed // Should succeed — ttl=1 decrements to 0 but packet is still processed
node.handle_session_datagram(&from, &encoded[1..], false).await; node.handle_session_datagram(&from, &encoded[1..], false)
.await;
} }
// --- Local delivery --- // --- Local delivery ---
@@ -69,7 +68,8 @@ async fn test_forwarding_local_delivery() {
let dg = SessionDatagram::new(from, my_addr, vec![0x10, 0x00, 0x00, 0x00]); let dg = SessionDatagram::new(from, my_addr, vec![0x10, 0x00, 0x00, 0x00]);
let encoded = dg.encode(); let encoded = dg.encode();
// Should detect local delivery and return without forwarding // Should detect local delivery and return without forwarding
node.handle_session_datagram(&from, &encoded[1..], false).await; node.handle_session_datagram(&from, &encoded[1..], false)
.await;
} }
// --- Direct peer forwarding --- // --- Direct peer forwarding ---
@@ -135,7 +135,8 @@ async fn test_coord_cache_warming_session_setup() {
// Handle the datagram (will be local delivery or no-route, but cache warming // Handle the datagram (will be local delivery or no-route, but cache warming
// happens before routing decision) // happens before routing decision)
node.handle_session_datagram(&from, &encoded[1..], false).await; node.handle_session_datagram(&from, &encoded[1..], false)
.await;
// After: both src and dest coords should be cached // After: both src and dest coords should be cached
let cached_src = node.coord_cache().get(&src_addr, now_ms); let cached_src = node.coord_cache().get(&src_addr, now_ms);
@@ -175,15 +176,22 @@ async fn test_coord_cache_warming_session_ack() {
assert!(node.coord_cache().get(&src_addr, now_ms).is_none()); assert!(node.coord_cache().get(&src_addr, now_ms).is_none());
assert!(node.coord_cache().get(&dest_addr, now_ms).is_none()); assert!(node.coord_cache().get(&dest_addr, now_ms).is_none());
node.handle_session_datagram(&from, &encoded[1..], false).await; node.handle_session_datagram(&from, &encoded[1..], false)
.await;
// SessionAck caches both src_coords and dest_coords // SessionAck caches both src_coords and dest_coords
let cached_src = node.coord_cache().get(&src_addr, now_ms); let cached_src = node.coord_cache().get(&src_addr, now_ms);
assert!(cached_src.is_some(), "src_addr coords not cached from SessionAck"); assert!(
cached_src.is_some(),
"src_addr coords not cached from SessionAck"
);
assert_eq!(cached_src.unwrap().root_id(), &root_addr); assert_eq!(cached_src.unwrap().root_id(), &root_addr);
let cached_dest = node.coord_cache().get(&dest_addr, now_ms); let cached_dest = node.coord_cache().get(&dest_addr, now_ms);
assert!(cached_dest.is_some(), "dest_addr coords not cached from SessionAck"); assert!(
cached_dest.is_some(),
"dest_addr coords not cached from SessionAck"
);
assert_eq!(cached_dest.unwrap().root_id(), &root_addr); assert_eq!(cached_dest.unwrap().root_id(), &root_addr);
} }
@@ -217,7 +225,8 @@ async fn test_coord_cache_warming_encrypted_msg_with_coords() {
assert!(node.coord_cache().get(&src_addr, now_ms).is_none()); assert!(node.coord_cache().get(&src_addr, now_ms).is_none());
assert!(node.coord_cache().get(&dest_addr, now_ms).is_none()); assert!(node.coord_cache().get(&dest_addr, now_ms).is_none());
node.handle_session_datagram(&from, &encoded[1..], false).await; node.handle_session_datagram(&from, &encoded[1..], false)
.await;
assert!( assert!(
node.coord_cache().get(&src_addr, now_ms).is_some(), node.coord_cache().get(&src_addr, now_ms).is_some(),
@@ -250,7 +259,8 @@ async fn test_coord_cache_warming_encrypted_msg_no_coords() {
.unwrap() .unwrap()
.as_millis() as u64; .as_millis() as u64;
node.handle_session_datagram(&from, &encoded[1..], false).await; node.handle_session_datagram(&from, &encoded[1..], false)
.await;
assert!( assert!(
node.coord_cache().get(&src_addr, now_ms).is_none(), node.coord_cache().get(&src_addr, now_ms).is_none(),
@@ -512,8 +522,16 @@ async fn test_forwarding_with_cache_warming_enables_routing() {
// Give each node coords for its direct peers only // Give each node coords for its direct peers only
let j_addr = *nodes[j].node.node_addr(); let j_addr = *nodes[j].node.node_addr();
if nodes[i].node.get_peer(&j_addr).is_some() { if nodes[i].node.get_peer(&j_addr).is_some() {
let coords = all_coords.iter().find(|(a, _)| a == &j_addr).unwrap().1.clone(); let coords = all_coords
nodes[i].node.coord_cache_mut().insert(j_addr, coords, now_ms); .iter()
.find(|(a, _)| a == &j_addr)
.unwrap()
.1
.clone();
nodes[i]
.node
.coord_cache_mut()
.insert(j_addr, coords, now_ms);
} }
} }
} }
@@ -572,8 +590,8 @@ async fn test_forwarding_with_cache_warming_enables_routing() {
// ECN Tests // ECN Tests
// ============================================================================ // ============================================================================
use crate::node::handlers::session::mark_ipv6_ecn_ce;
use crate::node::TransportDropState; use crate::node::TransportDropState;
use crate::node::handlers::session::mark_ipv6_ecn_ce;
use crate::transport::TransportId; use crate::transport::TransportId;
/// Build a minimal IPv6 header (40 bytes) with specified ECN bits. /// Build a minimal IPv6 header (40 bytes) with specified ECN bits.
@@ -721,10 +739,13 @@ fn test_detect_congestion_with_transport_drops() {
// Simulate transport kernel drops // Simulate transport kernel drops
let tid = TransportId::new(1); let tid = TransportId::new(1);
node.transport_drops.insert(tid, TransportDropState { node.transport_drops.insert(
prev_drops: 100, tid,
dropping: true, TransportDropState {
}); prev_drops: 100,
dropping: true,
},
);
// Now detect_congestion should return true (local transport congestion) // Now detect_congestion should return true (local transport congestion)
assert!(node.detect_congestion(&fake_addr)); assert!(node.detect_congestion(&fake_addr));
@@ -741,10 +762,13 @@ fn test_detect_congestion_disabled_ecn() {
// Even with transport drops, disabled ECN should return false // Even with transport drops, disabled ECN should return false
let tid = TransportId::new(1); let tid = TransportId::new(1);
node.transport_drops.insert(tid, TransportDropState { node.transport_drops.insert(
prev_drops: 50, tid,
dropping: true, TransportDropState {
}); prev_drops: 50,
dropping: true,
},
);
let fake_addr = NodeAddr::from_bytes([1; 16]); let fake_addr = NodeAddr::from_bytes([1; 16]);
assert!(!node.detect_congestion(&fake_addr)); assert!(!node.detect_congestion(&fake_addr));
@@ -756,10 +780,13 @@ fn test_sample_transport_congestion() {
// Insert a transport drop state with a baseline // Insert a transport drop state with a baseline
let tid = TransportId::new(1); let tid = TransportId::new(1);
node.transport_drops.insert(tid, TransportDropState { node.transport_drops.insert(
prev_drops: 0, tid,
dropping: false, TransportDropState {
}); prev_drops: 0,
dropping: false,
},
);
// No transports registered — sample_transport_congestion is a no-op // No transports registered — sample_transport_congestion is a no-op
// (transport_drops entry stays unchanged) // (transport_drops entry stays unchanged)

View File

@@ -5,9 +5,11 @@ use super::*;
#[tokio::test] #[tokio::test]
async fn test_two_node_handshake_udp() { async fn test_two_node_handshake_udp() {
use crate::config::UdpConfig; use crate::config::UdpConfig;
use crate::node::wire::{
build_encrypted, build_established_header, build_msg1, prepend_inner_header,
};
use crate::transport::udp::UdpTransport; use crate::transport::udp::UdpTransport;
use crate::node::wire::{build_encrypted, build_established_header, build_msg1, prepend_inner_header}; use tokio::time::{Duration, timeout};
use tokio::time::{timeout, Duration};
// === Setup: Two nodes with UDP transports on localhost === // === Setup: Two nodes with UDP transports on localhost ===
@@ -26,10 +28,8 @@ async fn test_two_node_handshake_udp() {
let (packet_tx_a, mut packet_rx_a) = packet_channel(64); let (packet_tx_a, mut packet_rx_a) = packet_channel(64);
let (packet_tx_b, mut packet_rx_b) = packet_channel(64); let (packet_tx_b, mut packet_rx_b) = packet_channel(64);
let mut transport_a = let mut transport_a = UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a); let mut transport_b = UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
let mut transport_b =
UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
transport_a.start_async().await.unwrap(); transport_a.start_async().await.unwrap();
transport_b.start_async().await.unwrap(); transport_b.start_async().await.unwrap();
@@ -49,23 +49,20 @@ async fn test_two_node_handshake_udp() {
// === Phase 1: Node A initiates handshake to Node B === // === Phase 1: Node A initiates handshake to Node B ===
// Create peer identity for B (must use full key for ECDH parity) // Create peer identity for B (must use full key for ECDH parity)
let peer_b_identity = let peer_b_identity = PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_node_addr = *peer_b_identity.node_addr(); let peer_b_node_addr = *peer_b_identity.node_addr();
let link_id_a = node_a.allocate_link_id(); let link_id_a = node_a.allocate_link_id();
let mut conn_a = PeerConnection::outbound( let mut conn_a = PeerConnection::outbound(link_id_a, peer_b_identity, 1000);
link_id_a,
peer_b_identity,
1000,
);
// Allocate session index for A's outbound // Allocate session index for A's outbound
let our_index_a = node_a.index_allocator.allocate().unwrap(); let our_index_a = node_a.index_allocator.allocate().unwrap();
// Start handshake (generates Noise IK msg1) // Start handshake (generates Noise IK msg1)
let our_keypair_a = node_a.identity.keypair(); let our_keypair_a = node_a.identity.keypair();
let noise_msg1 = conn_a.start_handshake(our_keypair_a, node_a.startup_epoch, 1000).unwrap(); let noise_msg1 = conn_a
.start_handshake(our_keypair_a, node_a.startup_epoch, 1000)
.unwrap();
conn_a.set_our_index(our_index_a); conn_a.set_our_index(our_index_a);
conn_a.set_transport_id(transport_id_a); conn_a.set_transport_id(transport_id_a);
conn_a.set_source_addr(remote_addr_b.clone()); conn_a.set_source_addr(remote_addr_b.clone());
@@ -82,10 +79,9 @@ async fn test_two_node_handshake_udp() {
); );
node_a.links.insert(link_id_a, link_a); node_a.links.insert(link_id_a, link_a);
node_a.connections.insert(link_id_a, conn_a); node_a.connections.insert(link_id_a, conn_a);
node_a.pending_outbound.insert( node_a
(transport_id_a, our_index_a.as_u32()), .pending_outbound
link_id_a, .insert((transport_id_a, our_index_a.as_u32()), link_id_a);
);
// Send msg1 from A to B over UDP // Send msg1 from A to B over UDP
let transport = node_a.transports.get(&transport_id_a).unwrap(); let transport = node_a.transports.get(&transport_id_a).unwrap();
@@ -104,11 +100,13 @@ async fn test_two_node_handshake_udp() {
node_b.handle_msg1(packet_b).await; node_b.handle_msg1(packet_b).await;
// Verify B promoted the inbound connection // Verify B promoted the inbound connection
let peer_a_node_addr = *PeerIdentity::from_pubkey_full( let peer_a_node_addr =
node_a.identity.pubkey_full(), *PeerIdentity::from_pubkey_full(node_a.identity.pubkey_full()).node_addr();
) assert_eq!(
.node_addr(); node_b.peer_count(),
assert_eq!(node_b.peer_count(), 1, "Node B should have 1 peer after msg1"); 1,
"Node B should have 1 peer after msg1"
);
let peer_a_on_b = node_b let peer_a_on_b = node_b
.get_peer(&peer_a_node_addr) .get_peer(&peer_a_node_addr)
.expect("Node B should have peer A"); .expect("Node B should have peer A");
@@ -134,7 +132,11 @@ async fn test_two_node_handshake_udp() {
node_a.handle_msg2(packet_a).await; node_a.handle_msg2(packet_a).await;
// Verify A promoted the outbound connection // Verify A promoted the outbound connection
assert_eq!(node_a.peer_count(), 1, "Node A should have 1 peer after msg2"); assert_eq!(
node_a.peer_count(),
1,
"Node A should have 1 peer after msg2"
);
let peer_b_on_a = node_a let peer_b_on_a = node_a
.get_peer(&peer_b_node_addr) .get_peer(&peer_b_node_addr)
.expect("Node A should have peer B"); .expect("Node A should have peer B");
@@ -241,8 +243,8 @@ async fn test_two_node_handshake_udp() {
#[tokio::test] #[tokio::test]
async fn test_run_rx_loop_handshake() { async fn test_run_rx_loop_handshake() {
use crate::config::UdpConfig; use crate::config::UdpConfig;
use crate::transport::udp::UdpTransport;
use crate::node::wire::build_msg1; use crate::node::wire::build_msg1;
use crate::transport::udp::UdpTransport;
use tokio::time::Duration; use tokio::time::Duration;
// === Setup: Two nodes with UDP transports on localhost === // === Setup: Two nodes with UDP transports on localhost ===
@@ -262,10 +264,8 @@ async fn test_run_rx_loop_handshake() {
let (packet_tx_a, packet_rx_a) = packet_channel(64); let (packet_tx_a, packet_rx_a) = packet_channel(64);
let (packet_tx_b, packet_rx_b) = packet_channel(64); let (packet_tx_b, packet_rx_b) = packet_channel(64);
let mut transport_a = let mut transport_a = UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a); let mut transport_b = UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
let mut transport_b =
UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
transport_a.start_async().await.unwrap(); transport_a.start_async().await.unwrap();
transport_b.start_async().await.unwrap(); transport_b.start_async().await.unwrap();
@@ -290,20 +290,17 @@ async fn test_run_rx_loop_handshake() {
// === Phase 1: Node A initiates handshake to Node B === // === Phase 1: Node A initiates handshake to Node B ===
let peer_b_identity = let peer_b_identity = PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_node_addr = *peer_b_identity.node_addr(); let peer_b_node_addr = *peer_b_identity.node_addr();
let link_id_a = node_a.allocate_link_id(); let link_id_a = node_a.allocate_link_id();
let mut conn_a = PeerConnection::outbound( let mut conn_a = PeerConnection::outbound(link_id_a, peer_b_identity, 1000);
link_id_a,
peer_b_identity,
1000,
);
let our_index_a = node_a.index_allocator.allocate().unwrap(); let our_index_a = node_a.index_allocator.allocate().unwrap();
let our_keypair_a = node_a.identity.keypair(); let our_keypair_a = node_a.identity.keypair();
let noise_msg1 = conn_a.start_handshake(our_keypair_a, node_a.startup_epoch, 1000).unwrap(); let noise_msg1 = conn_a
.start_handshake(our_keypair_a, node_a.startup_epoch, 1000)
.unwrap();
conn_a.set_our_index(our_index_a); conn_a.set_our_index(our_index_a);
conn_a.set_transport_id(transport_id_a); conn_a.set_transport_id(transport_id_a);
conn_a.set_source_addr(remote_addr_b.clone()); conn_a.set_source_addr(remote_addr_b.clone());
@@ -319,10 +316,9 @@ async fn test_run_rx_loop_handshake() {
); );
node_a.links.insert(link_id_a, link_a); node_a.links.insert(link_id_a, link_a);
node_a.connections.insert(link_id_a, conn_a); node_a.connections.insert(link_id_a, conn_a);
node_a.pending_outbound.insert( node_a
(transport_id_a, our_index_a.as_u32()), .pending_outbound
link_id_a, .insert((transport_id_a, our_index_a.as_u32()), link_id_a);
);
// Send msg1 from A to B over real UDP // Send msg1 from A to B over real UDP
let transport = node_a.transports.get(&transport_id_a).unwrap(); let transport = node_a.transports.get(&transport_id_a).unwrap();
@@ -350,12 +346,14 @@ async fn test_run_rx_loop_handshake() {
} }
// Verify Node B promoted the inbound connection via rx loop dispatch // Verify Node B promoted the inbound connection via rx loop dispatch
let peer_a_node_addr = *PeerIdentity::from_pubkey_full( let peer_a_node_addr =
node_a.identity.pubkey_full(), *PeerIdentity::from_pubkey_full(node_a.identity.pubkey_full()).node_addr();
)
.node_addr();
assert_eq!(node_b.peer_count(), 1, "Node B should have 1 peer after rx loop processed msg1"); assert_eq!(
node_b.peer_count(),
1,
"Node B should have 1 peer after rx loop processed msg1"
);
let peer_a_on_b = node_b let peer_a_on_b = node_b
.get_peer(&peer_a_node_addr) .get_peer(&peer_a_node_addr)
.expect("Node B should have peer A"); .expect("Node B should have peer A");
@@ -390,7 +388,11 @@ async fn test_run_rx_loop_handshake() {
} }
// Verify Node A promoted the outbound connection via rx loop dispatch // Verify Node A promoted the outbound connection via rx loop dispatch
assert_eq!(node_a.peer_count(), 1, "Node A should have 1 peer after rx loop processed msg2"); assert_eq!(
node_a.peer_count(),
1,
"Node A should have 1 peer after rx loop processed msg2"
);
let peer_b_on_a = node_a let peer_b_on_a = node_a
.get_peer(&peer_b_node_addr) .get_peer(&peer_b_node_addr)
.expect("Node A should have peer B"); .expect("Node A should have peer B");
@@ -432,9 +434,9 @@ async fn test_run_rx_loop_handshake() {
#[tokio::test] #[tokio::test]
async fn test_cross_connection_both_initiate() { async fn test_cross_connection_both_initiate() {
use crate::config::UdpConfig; use crate::config::UdpConfig;
use crate::transport::udp::UdpTransport;
use crate::node::wire::build_msg1; use crate::node::wire::build_msg1;
use tokio::time::{timeout, Duration}; use crate::transport::udp::UdpTransport;
use tokio::time::{Duration, timeout};
// === Setup: Two nodes with UDP transports on localhost === // === Setup: Two nodes with UDP transports on localhost ===
@@ -453,10 +455,8 @@ async fn test_cross_connection_both_initiate() {
let (packet_tx_a, mut packet_rx_a) = packet_channel(64); let (packet_tx_a, mut packet_rx_a) = packet_channel(64);
let (packet_tx_b, mut packet_rx_b) = packet_channel(64); let (packet_tx_b, mut packet_rx_b) = packet_channel(64);
let mut transport_a = let mut transport_a = UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a);
UdpTransport::new(transport_id_a, None, udp_config.clone(), packet_tx_a); let mut transport_b = UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
let mut transport_b =
UdpTransport::new(transport_id_b, None, udp_config, packet_tx_b);
transport_a.start_async().await.unwrap(); transport_a.start_async().await.unwrap();
transport_b.start_async().await.unwrap(); transport_b.start_async().await.unwrap();
@@ -474,11 +474,9 @@ async fn test_cross_connection_both_initiate() {
.insert(transport_id_b, TransportHandle::Udp(transport_b)); .insert(transport_id_b, TransportHandle::Udp(transport_b));
// Peer identities (must use full key for ECDH parity) // Peer identities (must use full key for ECDH parity)
let peer_b_identity = let peer_b_identity = PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
PeerIdentity::from_pubkey_full(node_b.identity.pubkey_full());
let peer_b_node_addr = *peer_b_identity.node_addr(); let peer_b_node_addr = *peer_b_identity.node_addr();
let peer_a_identity = let peer_a_identity = PeerIdentity::from_pubkey_full(node_a.identity.pubkey_full());
PeerIdentity::from_pubkey_full(node_a.identity.pubkey_full());
let peer_a_node_addr = *peer_a_identity.node_addr(); let peer_a_node_addr = *peer_a_identity.node_addr();
// === Phase 1: Both nodes initiate handshakes (simulate auto_connect) === // === Phase 1: Both nodes initiate handshakes (simulate auto_connect) ===
@@ -488,7 +486,9 @@ async fn test_cross_connection_both_initiate() {
let mut conn_a = PeerConnection::outbound(link_id_a_out, peer_b_identity, 1000); let mut conn_a = PeerConnection::outbound(link_id_a_out, peer_b_identity, 1000);
let our_index_a = node_a.index_allocator.allocate().unwrap(); let our_index_a = node_a.index_allocator.allocate().unwrap();
let our_keypair_a = node_a.identity.keypair(); let our_keypair_a = node_a.identity.keypair();
let noise_msg1_a = conn_a.start_handshake(our_keypair_a, node_a.startup_epoch, 1000).unwrap(); let noise_msg1_a = conn_a
.start_handshake(our_keypair_a, node_a.startup_epoch, 1000)
.unwrap();
conn_a.set_our_index(our_index_a); conn_a.set_our_index(our_index_a);
conn_a.set_transport_id(transport_id_a); conn_a.set_transport_id(transport_id_a);
conn_a.set_source_addr(remote_addr_b.clone()); conn_a.set_source_addr(remote_addr_b.clone());
@@ -496,20 +496,29 @@ async fn test_cross_connection_both_initiate() {
let wire_msg1_a = build_msg1(our_index_a, &noise_msg1_a); let wire_msg1_a = build_msg1(our_index_a, &noise_msg1_a);
let link_a_out = Link::connectionless( let link_a_out = Link::connectionless(
link_id_a_out, transport_id_a, remote_addr_b.clone(), link_id_a_out,
LinkDirection::Outbound, Duration::from_millis(100), transport_id_a,
remote_addr_b.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
); );
node_a.links.insert(link_id_a_out, link_a_out); node_a.links.insert(link_id_a_out, link_a_out);
node_a.addr_to_link.insert((transport_id_a, remote_addr_b.clone()), link_id_a_out); node_a
.addr_to_link
.insert((transport_id_a, remote_addr_b.clone()), link_id_a_out);
node_a.connections.insert(link_id_a_out, conn_a); node_a.connections.insert(link_id_a_out, conn_a);
node_a.pending_outbound.insert((transport_id_a, our_index_a.as_u32()), link_id_a_out); node_a
.pending_outbound
.insert((transport_id_a, our_index_a.as_u32()), link_id_a_out);
// Node B initiates to Node A // Node B initiates to Node A
let link_id_b_out = node_b.allocate_link_id(); let link_id_b_out = node_b.allocate_link_id();
let mut conn_b = PeerConnection::outbound(link_id_b_out, peer_a_identity, 1000); let mut conn_b = PeerConnection::outbound(link_id_b_out, peer_a_identity, 1000);
let our_index_b = node_b.index_allocator.allocate().unwrap(); let our_index_b = node_b.index_allocator.allocate().unwrap();
let our_keypair_b = node_b.identity.keypair(); let our_keypair_b = node_b.identity.keypair();
let noise_msg1_b = conn_b.start_handshake(our_keypair_b, node_b.startup_epoch, 1000).unwrap(); let noise_msg1_b = conn_b
.start_handshake(our_keypair_b, node_b.startup_epoch, 1000)
.unwrap();
conn_b.set_our_index(our_index_b); conn_b.set_our_index(our_index_b);
conn_b.set_transport_id(transport_id_b); conn_b.set_transport_id(transport_id_b);
conn_b.set_source_addr(remote_addr_a.clone()); conn_b.set_source_addr(remote_addr_a.clone());
@@ -517,20 +526,33 @@ async fn test_cross_connection_both_initiate() {
let wire_msg1_b = build_msg1(our_index_b, &noise_msg1_b); let wire_msg1_b = build_msg1(our_index_b, &noise_msg1_b);
let link_b_out = Link::connectionless( let link_b_out = Link::connectionless(
link_id_b_out, transport_id_b, remote_addr_a.clone(), link_id_b_out,
LinkDirection::Outbound, Duration::from_millis(100), transport_id_b,
remote_addr_a.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
); );
node_b.links.insert(link_id_b_out, link_b_out); node_b.links.insert(link_id_b_out, link_b_out);
node_b.addr_to_link.insert((transport_id_b, remote_addr_a.clone()), link_id_b_out); node_b
.addr_to_link
.insert((transport_id_b, remote_addr_a.clone()), link_id_b_out);
node_b.connections.insert(link_id_b_out, conn_b); node_b.connections.insert(link_id_b_out, conn_b);
node_b.pending_outbound.insert((transport_id_b, our_index_b.as_u32()), link_id_b_out); node_b
.pending_outbound
.insert((transport_id_b, our_index_b.as_u32()), link_id_b_out);
// Both send msg1 over UDP // Both send msg1 over UDP
let transport = node_a.transports.get(&transport_id_a).unwrap(); let transport = node_a.transports.get(&transport_id_a).unwrap();
transport.send(&remote_addr_b, &wire_msg1_a).await.expect("A send msg1"); transport
.send(&remote_addr_b, &wire_msg1_a)
.await
.expect("A send msg1");
let transport = node_b.transports.get(&transport_id_b).unwrap(); let transport = node_b.transports.get(&transport_id_b).unwrap();
transport.send(&remote_addr_a, &wire_msg1_b).await.expect("B send msg1"); transport
.send(&remote_addr_a, &wire_msg1_b)
.await
.expect("B send msg1");
// === Phase 2: Both nodes receive the other's msg1 === // === Phase 2: Both nodes receive the other's msg1 ===
// Before the fix, addr_to_link would reject these because outbound links // Before the fix, addr_to_link would reject these because outbound links
@@ -538,21 +560,39 @@ async fn test_cross_connection_both_initiate() {
// B receives A's msg1 // B receives A's msg1
let packet_at_b = timeout(Duration::from_secs(1), packet_rx_b.recv()) let packet_at_b = timeout(Duration::from_secs(1), packet_rx_b.recv())
.await.expect("Timeout").expect("Channel closed"); .await
.expect("Timeout")
.expect("Channel closed");
node_b.handle_msg1(packet_at_b).await; node_b.handle_msg1(packet_at_b).await;
// B should have promoted the inbound connection // B should have promoted the inbound connection
assert_eq!(node_b.peer_count(), 1, "Node B should have 1 peer after processing A's msg1"); assert_eq!(
assert!(node_b.get_peer(&peer_a_node_addr).is_some(), "Node B should have peer A"); node_b.peer_count(),
1,
"Node B should have 1 peer after processing A's msg1"
);
assert!(
node_b.get_peer(&peer_a_node_addr).is_some(),
"Node B should have peer A"
);
// A receives B's msg1 // A receives B's msg1
let packet_at_a = timeout(Duration::from_secs(1), packet_rx_a.recv()) let packet_at_a = timeout(Duration::from_secs(1), packet_rx_a.recv())
.await.expect("Timeout").expect("Channel closed"); .await
.expect("Timeout")
.expect("Channel closed");
node_a.handle_msg1(packet_at_a).await; node_a.handle_msg1(packet_at_a).await;
// A should have promoted the inbound connection // A should have promoted the inbound connection
assert_eq!(node_a.peer_count(), 1, "Node A should have 1 peer after processing B's msg1"); assert_eq!(
assert!(node_a.get_peer(&peer_b_node_addr).is_some(), "Node A should have peer B"); node_a.peer_count(),
1,
"Node A should have 1 peer after processing B's msg1"
);
assert!(
node_a.get_peer(&peer_b_node_addr).is_some(),
"Node A should have peer B"
);
// === Phase 3: Both nodes receive msg2 responses === // === Phase 3: Both nodes receive msg2 responses ===
// The msg2 was sent during handle_msg1 processing. When handle_msg2 // The msg2 was sent during handle_msg1 processing. When handle_msg2
@@ -560,21 +600,37 @@ async fn test_cross_connection_both_initiate() {
// A receives B's msg2 (response to A's original msg1) // A receives B's msg2 (response to A's original msg1)
let msg2_at_a = timeout(Duration::from_secs(1), packet_rx_a.recv()) let msg2_at_a = timeout(Duration::from_secs(1), packet_rx_a.recv())
.await.expect("Timeout waiting for msg2 at A").expect("Channel closed"); .await
.expect("Timeout waiting for msg2 at A")
.expect("Channel closed");
node_a.handle_msg2(msg2_at_a).await; node_a.handle_msg2(msg2_at_a).await;
// B receives A's msg2 (response to B's original msg1) // B receives A's msg2 (response to B's original msg1)
let msg2_at_b = timeout(Duration::from_secs(1), packet_rx_b.recv()) let msg2_at_b = timeout(Duration::from_secs(1), packet_rx_b.recv())
.await.expect("Timeout waiting for msg2 at B").expect("Channel closed"); .await
.expect("Timeout waiting for msg2 at B")
.expect("Channel closed");
node_b.handle_msg2(msg2_at_b).await; node_b.handle_msg2(msg2_at_b).await;
// === Verification === // === Verification ===
// Both nodes should have exactly 1 peer each after cross-connection resolution // Both nodes should have exactly 1 peer each after cross-connection resolution
assert_eq!(node_a.peer_count(), 1, "Node A should have exactly 1 peer after cross-connection"); assert_eq!(
assert_eq!(node_b.peer_count(), 1, "Node B should have exactly 1 peer after cross-connection"); node_a.peer_count(),
1,
"Node A should have exactly 1 peer after cross-connection"
);
assert_eq!(
node_b.peer_count(),
1,
"Node B should have exactly 1 peer after cross-connection"
);
let peer_b_on_a = node_a.get_peer(&peer_b_node_addr).expect("A should have peer B"); let peer_b_on_a = node_a
let peer_a_on_b = node_b.get_peer(&peer_a_node_addr).expect("B should have peer A"); .get_peer(&peer_b_node_addr)
.expect("A should have peer B");
let peer_a_on_b = node_b
.get_peer(&peer_a_node_addr)
.expect("B should have peer A");
assert!(peer_b_on_a.has_session(), "Peer B on A should have session"); assert!(peer_b_on_a.has_session(), "Peer B on A should have session");
assert!(peer_a_on_b.has_session(), "Peer A on B should have session"); assert!(peer_a_on_b.has_session(), "Peer A on B should have session");
@@ -611,25 +667,35 @@ async fn test_stale_connection_cleanup() {
// Allocate session index and set transport info // Allocate session index and set transport info
let our_index = node.index_allocator.allocate().unwrap(); let our_index = node.index_allocator.allocate().unwrap();
let our_keypair = node.identity.keypair(); let our_keypair = node.identity.keypair();
let _noise_msg1 = conn.start_handshake(our_keypair, node.startup_epoch, past_time_ms).unwrap(); let _noise_msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, past_time_ms)
.unwrap();
conn.set_our_index(our_index); conn.set_our_index(our_index);
conn.set_transport_id(transport_id); conn.set_transport_id(transport_id);
conn.set_source_addr(remote_addr.clone()); conn.set_source_addr(remote_addr.clone());
// Set up all the state that initiate_peer_connection would create // Set up all the state that initiate_peer_connection would create
let link = Link::connectionless( let link = Link::connectionless(
link_id, transport_id, remote_addr.clone(), link_id,
LinkDirection::Outbound, Duration::from_millis(100), transport_id,
remote_addr.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
); );
node.links.insert(link_id, link); node.links.insert(link_id, link);
node.addr_to_link.insert((transport_id, remote_addr.clone()), link_id); node.addr_to_link
.insert((transport_id, remote_addr.clone()), link_id);
node.connections.insert(link_id, conn); node.connections.insert(link_id, conn);
node.pending_outbound.insert((transport_id, our_index.as_u32()), link_id); node.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
// Verify state before timeout check // Verify state before timeout check
assert_eq!(node.connection_count(), 1); assert_eq!(node.connection_count(), 1);
assert_eq!(node.link_count(), 1); assert_eq!(node.link_count(), 1);
assert!(node.pending_outbound.contains_key(&(transport_id, our_index.as_u32()))); assert!(
node.pending_outbound
.contains_key(&(transport_id, our_index.as_u32()))
);
assert_eq!(node.index_allocator.count(), 1); assert_eq!(node.index_allocator.count(), 1);
// Connection was created at time 1000ms. check_timeouts uses SystemTime::now(), // Connection was created at time 1000ms. check_timeouts uses SystemTime::now(),
@@ -637,13 +703,27 @@ async fn test_stale_connection_cleanup() {
node.check_timeouts(); node.check_timeouts();
// Verify everything was cleaned up // Verify everything was cleaned up
assert_eq!(node.connection_count(), 0, "Stale connection should be removed"); assert_eq!(
node.connection_count(),
0,
"Stale connection should be removed"
);
assert_eq!(node.link_count(), 0, "Stale link should be removed"); assert_eq!(node.link_count(), 0, "Stale link should be removed");
assert!(!node.pending_outbound.contains_key(&(transport_id, our_index.as_u32())), assert!(
"pending_outbound should be cleaned up"); !node
assert_eq!(node.index_allocator.count(), 0, "Session index should be freed"); .pending_outbound
assert!(!node.addr_to_link.contains_key(&(transport_id, remote_addr)), .contains_key(&(transport_id, our_index.as_u32())),
"addr_to_link should be cleaned up"); "pending_outbound should be cleaned up"
);
assert_eq!(
node.index_allocator.count(),
0,
"Session index should be freed"
);
assert!(
!node.addr_to_link.contains_key(&(transport_id, remote_addr)),
"addr_to_link should be cleaned up"
);
} }
/// Test that failed connections are cleaned up by check_timeouts(). /// Test that failed connections are cleaned up by check_timeouts().
@@ -665,29 +745,44 @@ async fn test_failed_connection_cleanup() {
let our_index = node.index_allocator.allocate().unwrap(); let our_index = node.index_allocator.allocate().unwrap();
let our_keypair = node.identity.keypair(); let our_keypair = node.identity.keypair();
let _noise_msg1 = conn.start_handshake(our_keypair, node.startup_epoch, now_ms).unwrap(); let _noise_msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, now_ms)
.unwrap();
conn.set_our_index(our_index); conn.set_our_index(our_index);
conn.set_transport_id(transport_id); conn.set_transport_id(transport_id);
conn.set_source_addr(remote_addr.clone()); conn.set_source_addr(remote_addr.clone());
conn.mark_failed(); // Simulate send failure conn.mark_failed(); // Simulate send failure
let link = Link::connectionless( let link = Link::connectionless(
link_id, transport_id, remote_addr.clone(), link_id,
LinkDirection::Outbound, Duration::from_millis(100), transport_id,
remote_addr.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
); );
node.links.insert(link_id, link); node.links.insert(link_id, link);
node.addr_to_link.insert((transport_id, remote_addr.clone()), link_id); node.addr_to_link
.insert((transport_id, remote_addr.clone()), link_id);
node.connections.insert(link_id, conn); node.connections.insert(link_id, conn);
node.pending_outbound.insert((transport_id, our_index.as_u32()), link_id); node.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
assert_eq!(node.connection_count(), 1); assert_eq!(node.connection_count(), 1);
// Failed connections should be cleaned up immediately regardless of age // Failed connections should be cleaned up immediately regardless of age
node.check_timeouts(); node.check_timeouts();
assert_eq!(node.connection_count(), 0, "Failed connection should be removed"); assert_eq!(
node.connection_count(),
0,
"Failed connection should be removed"
);
assert_eq!(node.link_count(), 0, "Failed link should be removed"); assert_eq!(node.link_count(), 0, "Failed link should be removed");
assert_eq!(node.index_allocator.count(), 0, "Session index should be freed"); assert_eq!(
node.index_allocator.count(),
0,
"Session index should be freed"
);
} }
/// Test that msg1 bytes are stored on connection for resend. /// Test that msg1 bytes are stored on connection for resend.
@@ -710,7 +805,9 @@ async fn test_msg1_stored_for_resend() {
let our_index = node.index_allocator.allocate().unwrap(); let our_index = node.index_allocator.allocate().unwrap();
let our_keypair = node.identity.keypair(); let our_keypair = node.identity.keypair();
let noise_msg1 = conn.start_handshake(our_keypair, node.startup_epoch, now_ms).unwrap(); let noise_msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, now_ms)
.unwrap();
conn.set_our_index(our_index); conn.set_our_index(our_index);
conn.set_transport_id(transport_id); conn.set_transport_id(transport_id);
conn.set_source_addr(remote_addr.clone()); conn.set_source_addr(remote_addr.clone());
@@ -741,7 +838,9 @@ async fn test_resend_scheduling() {
let our_index = node.index_allocator.allocate().unwrap(); let our_index = node.index_allocator.allocate().unwrap();
let our_keypair = node.identity.keypair(); let our_keypair = node.identity.keypair();
let noise_msg1 = conn.start_handshake(our_keypair, node.startup_epoch, now_ms).unwrap(); let noise_msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, now_ms)
.unwrap();
conn.set_our_index(our_index); conn.set_our_index(our_index);
conn.set_transport_id(transport_id); conn.set_transport_id(transport_id);
conn.set_source_addr(remote_addr.clone()); conn.set_source_addr(remote_addr.clone());
@@ -751,12 +850,17 @@ async fn test_resend_scheduling() {
conn.set_handshake_msg1(wire_msg1, now_ms + 1000); conn.set_handshake_msg1(wire_msg1, now_ms + 1000);
let link = Link::connectionless( let link = Link::connectionless(
link_id, transport_id, remote_addr.clone(), link_id,
LinkDirection::Outbound, Duration::from_millis(100), transport_id,
remote_addr.clone(),
LinkDirection::Outbound,
Duration::from_millis(100),
); );
node.links.insert(link_id, link); node.links.insert(link_id, link);
node.addr_to_link.insert((transport_id, remote_addr), link_id); node.addr_to_link
node.pending_outbound.insert((transport_id, our_index.as_u32()), link_id); .insert((transport_id, remote_addr), link_id);
node.pending_outbound
.insert((transport_id, our_index.as_u32()), link_id);
node.connections.insert(link_id, conn); node.connections.insert(link_id, conn);
// Before resend time: nothing should happen (no transport = can't send, // Before resend time: nothing should happen (no transport = can't send,
@@ -772,7 +876,11 @@ async fn test_resend_scheduling() {
// No transport registered, so send fails — count stays 0. // No transport registered, so send fails — count stays 0.
// That's the expected behavior (transport absence is a transient condition). // That's the expected behavior (transport absence is a transient condition).
let conn = node.connections.get(&link_id).unwrap(); let conn = node.connections.get(&link_id).unwrap();
assert_eq!(conn.resend_count(), 0, "No transport means no resend recorded"); assert_eq!(
conn.resend_count(),
0,
"No transport means no resend recorded"
);
} }
/// Test that msg2 is stored on PeerConnection for responder resend. /// Test that msg2 is stored on PeerConnection for responder resend.

View File

@@ -1,7 +1,7 @@
use super::*; use super::*;
use crate::utils::index::SessionIndex;
use crate::transport::{packet_channel, LinkDirection, TransportAddr};
use crate::PeerIdentity; use crate::PeerIdentity;
use crate::transport::{LinkDirection, TransportAddr, packet_channel};
use crate::utils::index::SessionIndex;
use std::time::Duration; use std::time::Duration;
mod bloom; mod bloom;
@@ -53,7 +53,9 @@ pub(super) fn make_completed_connection(
// Run initiator side of handshake // Run initiator side of handshake
let our_keypair = node.identity.keypair(); let our_keypair = node.identity.keypair();
let msg1 = conn.start_handshake(our_keypair, node.startup_epoch, current_time_ms).unwrap(); let msg1 = conn
.start_handshake(our_keypair, node.startup_epoch, current_time_ms)
.unwrap();
// Run responder side to generate msg2 // Run responder side to generate msg2
let mut resp_conn = PeerConnection::inbound(LinkId::new(999), current_time_ms); let mut resp_conn = PeerConnection::inbound(LinkId::new(999), current_time_ms);

View File

@@ -7,8 +7,8 @@ use super::*;
use crate::bloom::BloomFilter; use crate::bloom::BloomFilter;
use crate::tree::{ParentDeclaration, TreeCoordinate}; use crate::tree::{ParentDeclaration, TreeCoordinate};
use spanning_tree::{ use spanning_tree::{
cleanup_nodes, drain_all_packets, generate_random_edges, initiate_handshake, make_test_node, TestNode, cleanup_nodes, drain_all_packets, generate_random_edges, initiate_handshake,
run_tree_test, verify_tree_convergence, TestNode, make_test_node, run_tree_test, verify_tree_convergence,
}; };
use std::collections::HashSet; use std::collections::HashSet;
@@ -83,8 +83,7 @@ fn test_routing_bloom_filter_hit() {
// Destination not directly connected — placed under peer1 in the tree // Destination not directly connected — placed under peer1 in the tree
let dest = make_node_addr(99); let dest = make_node_addr(99);
let dest_coords = let dest_coords = TreeCoordinate::from_addrs(vec![dest, peer1_addr, my_addr]).unwrap();
TreeCoordinate::from_addrs(vec![dest, peer1_addr, my_addr]).unwrap();
let now_ms = std::time::SystemTime::now() let now_ms = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis() as u64) .map(|d| d.as_millis() as u64)
@@ -126,17 +125,14 @@ fn test_routing_bloom_filter_multiple_hits_tiebreak() {
// Set up tree: we are root, all peers are our children (equidistant) // Set up tree: we are root, all peers are our children (equidistant)
for &addr in &peer_addrs { for &addr in &peer_addrs {
let coords = TreeCoordinate::from_addrs(vec![addr, my_addr]).unwrap(); let coords = TreeCoordinate::from_addrs(vec![addr, my_addr]).unwrap();
node.tree_state_mut().update_peer( node.tree_state_mut()
ParentDeclaration::new(addr, my_addr, 1, 1000), .update_peer(ParentDeclaration::new(addr, my_addr, 1, 1000), coords);
coords,
);
} }
// Destination placed under the first peer (arbitrary — all peers are // Destination placed under the first peer (arbitrary — all peers are
// equidistant from dest since dest is 2 hops from root via any child) // equidistant from dest since dest is 2 hops from root via any child)
let dest = make_node_addr(99); let dest = make_node_addr(99);
let dest_coords = let dest_coords = TreeCoordinate::from_addrs(vec![dest, peer_addrs[0], my_addr]).unwrap();
TreeCoordinate::from_addrs(vec![dest, peer_addrs[0], my_addr]).unwrap();
let now_ms = std::time::SystemTime::now() let now_ms = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH) .duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis() as u64) .map(|d| d.as_millis() as u64)
@@ -187,8 +183,7 @@ fn test_routing_tree_fallback() {
// Destination: a node under our peer in the tree // Destination: a node under our peer in the tree
let dest = make_node_addr(99); let dest = make_node_addr(99);
let dest_coords = let dest_coords = TreeCoordinate::from_addrs(vec![dest, peer_addr, my_addr]).unwrap();
TreeCoordinate::from_addrs(vec![dest, peer_addr, my_addr]).unwrap();
// Put dest coords in the cache // Put dest coords in the cache
let now_ms = std::time::SystemTime::now() let now_ms = std::time::SystemTime::now()
@@ -239,8 +234,7 @@ fn test_routing_refreshes_coord_cache_ttl() {
// Set up tree coordinates // Set up tree coordinates
let dest = make_node_addr(99); let dest = make_node_addr(99);
let dest_coords = let dest_coords = TreeCoordinate::from_addrs(vec![dest, peer_addr, my_addr]).unwrap();
TreeCoordinate::from_addrs(vec![dest, peer_addr, my_addr]).unwrap();
node.tree_state_mut().update_peer( node.tree_state_mut().update_peer(
ParentDeclaration::new(peer_addr, my_addr, 1, 1000), ParentDeclaration::new(peer_addr, my_addr, 1, 1000),
TreeCoordinate::from_addrs(vec![peer_addr, my_addr]).unwrap(), TreeCoordinate::from_addrs(vec![peer_addr, my_addr]).unwrap(),
@@ -252,7 +246,8 @@ fn test_routing_refreshes_coord_cache_ttl() {
.map(|d| d.as_millis() as u64) .map(|d| d.as_millis() as u64)
.unwrap_or(0); .unwrap_or(0);
let short_ttl = 10_000; // 10 seconds let short_ttl = 10_000; // 10 seconds
node.coord_cache_mut().insert_with_ttl(dest, dest_coords, now_ms, short_ttl); node.coord_cache_mut()
.insert_with_ttl(dest, dest_coords, now_ms, short_ttl);
let original_expiry = node.coord_cache().get_entry(&dest).unwrap().expires_at(); let original_expiry = node.coord_cache().get_entry(&dest).unwrap().expires_at();
// find_next_hop should succeed and refresh TTL to now + default_ttl (300s) // find_next_hop should succeed and refresh TTL to now + default_ttl (300s)
@@ -263,7 +258,8 @@ fn test_routing_refreshes_coord_cache_ttl() {
assert!( assert!(
new_expiry > original_expiry, new_expiry > original_expiry,
"find_next_hop should refresh the coord_cache TTL: original={}, new={}", "find_next_hop should refresh the coord_cache TTL: original={}, new={}",
original_expiry, new_expiry, original_expiry,
new_expiry,
); );
} }
@@ -384,11 +380,7 @@ async fn test_routing_chain_topology() {
// Verify tree convergence // Verify tree convergence
let root = nodes.iter().map(|n| *n.node.node_addr()).min().unwrap(); let root = nodes.iter().map(|n| *n.node.node_addr()).min().unwrap();
for tn in &nodes { for tn in &nodes {
assert_eq!( assert_eq!(*tn.node.tree_state().root(), root, "Tree not converged");
*tn.node.tree_state().root(),
root,
"Tree not converged"
);
} }
// Populate coord caches: each node caches the far-end node's coords // Populate coord caches: each node caches the far-end node's coords
@@ -453,8 +445,13 @@ async fn test_routing_bloom_preferred_over_tree() {
// filter routing selects peer2 (strictly closer to dest than us). // filter routing selects peer2 (strictly closer to dest than us).
let dest = make_node_addr(99); let dest = make_node_addr(99);
let peer2_addr = *nodes[2].node.node_addr(); let peer2_addr = *nodes[2].node.node_addr();
let mut dest_path: Vec<NodeAddr> = let mut dest_path: Vec<NodeAddr> = nodes[2]
nodes[2].node.tree_state().my_coords().node_addrs().copied().collect(); .node
.tree_state()
.my_coords()
.node_addrs()
.copied()
.collect();
dest_path.insert(0, dest); dest_path.insert(0, dest);
let dest_coords = TreeCoordinate::from_addrs(dest_path).unwrap(); let dest_coords = TreeCoordinate::from_addrs(dest_path).unwrap();
let now_ms = std::time::SystemTime::now() let now_ms = std::time::SystemTime::now()
@@ -602,13 +599,20 @@ async fn test_routing_reachability_100_nodes() {
// Collect all (addr, coords) pairs first to avoid borrow issues // Collect all (addr, coords) pairs first to avoid borrow issues
let all_coords: Vec<(NodeAddr, TreeCoordinate)> = nodes let all_coords: Vec<(NodeAddr, TreeCoordinate)> = nodes
.iter() .iter()
.map(|tn| (*tn.node.node_addr(), tn.node.tree_state().my_coords().clone())) .map(|tn| {
(
*tn.node.node_addr(),
tn.node.tree_state().my_coords().clone(),
)
})
.collect(); .collect();
for node in &mut nodes { for node in &mut nodes {
for (addr, coords) in &all_coords { for (addr, coords) in &all_coords {
if addr != node.node.node_addr() { if addr != node.node.node_addr() {
node.node.coord_cache_mut().insert(*addr, coords.clone(), now_ms); node.node
.coord_cache_mut()
.insert(*addr, coords.clone(), now_ms);
} }
} }
} }
@@ -654,10 +658,7 @@ async fn test_routing_reachability_100_nodes() {
0.0 0.0
}; };
eprintln!( eprintln!("\n === Routing Reachability ({} nodes) ===", NUM_NODES);
"\n === Routing Reachability ({} nodes) ===",
NUM_NODES
);
eprintln!( eprintln!(
" Pairs tested: {} | Delivered: {} | Failed: {} | Loops: {}", " Pairs tested: {} | Delivered: {} | Failed: {} | Loops: {}",
total_pairs, total_pairs,
@@ -665,10 +666,7 @@ async fn test_routing_reachability_100_nodes() {
failures.len(), failures.len(),
loops.len() loops.len()
); );
eprintln!( eprintln!(" Hops: avg={:.1} max={}", avg_hops, max_hops);
" Hops: avg={:.1} max={}",
avg_hops, max_hops
);
if !failures.is_empty() { if !failures.is_empty() {
let show = failures.len().min(10); let show = failures.len().min(10);
@@ -736,13 +734,20 @@ async fn test_routing_stops_after_peer_removal() {
let all_coords: Vec<(NodeAddr, crate::tree::TreeCoordinate)> = nodes let all_coords: Vec<(NodeAddr, crate::tree::TreeCoordinate)> = nodes
.iter() .iter()
.map(|tn| (*tn.node.node_addr(), tn.node.tree_state().my_coords().clone())) .map(|tn| {
(
*tn.node.node_addr(),
tn.node.tree_state().my_coords().clone(),
)
})
.collect(); .collect();
for node in &mut nodes { for node in &mut nodes {
for (addr, coords) in &all_coords { for (addr, coords) in &all_coords {
if addr != node.node.node_addr() { if addr != node.node.node_addr() {
node.node.coord_cache_mut().insert(*addr, coords.clone(), now_ms); node.node
.coord_cache_mut()
.insert(*addr, coords.clone(), now_ms);
} }
} }
} }
@@ -792,9 +797,12 @@ async fn test_routing_stops_after_peer_removal() {
// matters is that delivery does NOT succeed. // matters is that delivery does NOT succeed.
match simulate_forwarding(&mut nodes, &addr_index, 0, 3) { match simulate_forwarding(&mut nodes, &addr_index, 0, 3) {
ForwardResult::NoRoute { .. } => {} // Expected: can't reach node 3 ForwardResult::NoRoute { .. } => {} // Expected: can't reach node 3
ForwardResult::Loop { .. } => {} // Also acceptable: stale coords cause loop detection ForwardResult::Loop { .. } => {} // Also acceptable: stale coords cause loop detection
ForwardResult::Delivered(hops) => { ForwardResult::Delivered(hops) => {
panic!("Should NOT deliver after partition, but got delivery in {} hops", hops); panic!(
"Should NOT deliver after partition, but got delivery in {} hops",
hops
);
} }
} }
@@ -905,7 +913,12 @@ async fn test_routing_source_only_coords_100_nodes() {
// Collect all coords for injection // Collect all coords for injection
let all_coords: Vec<(NodeAddr, crate::tree::TreeCoordinate)> = nodes let all_coords: Vec<(NodeAddr, crate::tree::TreeCoordinate)> = nodes
.iter() .iter()
.map(|tn| (*tn.node.node_addr(), tn.node.tree_state().my_coords().clone())) .map(|tn| {
(
*tn.node.node_addr(),
tn.node.tree_state().my_coords().clone(),
)
})
.collect(); .collect();
let addr_index = build_addr_index(&nodes); let addr_index = build_addr_index(&nodes);
@@ -946,7 +959,10 @@ async fn test_routing_source_only_coords_100_nodes() {
ForwardResult::Delivered(_) => source_only_delivered += 1, ForwardResult::Delivered(_) => source_only_delivered += 1,
ForwardResult::NoRoute { .. } => source_only_failed += 1, ForwardResult::NoRoute { .. } => source_only_failed += 1,
ForwardResult::Loop { .. } => { ForwardResult::Loop { .. } => {
panic!("Routing loop detected with source-only coords: {} -> {}", src, dst); panic!(
"Routing loop detected with source-only coords: {} -> {}",
src, dst
);
} }
} }
} }
@@ -977,7 +993,9 @@ async fn test_routing_source_only_coords_100_nodes() {
for node in &mut nodes { for node in &mut nodes {
for (addr, coords) in &all_coords { for (addr, coords) in &all_coords {
if addr != node.node.node_addr() { if addr != node.node.node_addr() {
node.node.coord_cache_mut().insert(*addr, coords.clone(), now_ms); node.node
.coord_cache_mut()
.insert(*addr, coords.clone(), now_ms);
} }
} }
} }
@@ -996,4 +1014,3 @@ async fn test_routing_source_only_coords_100_nodes() {
cleanup_nodes(&mut nodes).await; cleanup_nodes(&mut nodes).await;
} }

View File

@@ -3,8 +3,8 @@
use super::*; use super::*;
use crate::node::session::EndToEndState; use crate::node::session::EndToEndState;
use crate::node::tests::spanning_tree::{ use crate::node::tests::spanning_tree::{
cleanup_nodes, generate_random_edges, process_available_packets, run_tree_test, TestNode, cleanup_nodes, generate_random_edges, process_available_packets, run_tree_test,
run_tree_test_with_mtus, verify_tree_convergence, TestNode, run_tree_test_with_mtus, verify_tree_convergence,
}; };
use crate::protocol::{SessionAck, SessionDatagram}; use crate::protocol::{SessionAck, SessionDatagram};
@@ -50,10 +50,7 @@ fn test_session_entry_new_initiating() {
let identity_a = Identity::generate(); let identity_a = Identity::generate();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_initiator( let handshake = HandshakeState::new_initiator(identity_a.keypair(), identity_b.pubkey_full());
identity_a.keypair(),
identity_b.pubkey_full(),
);
let entry = crate::node::session::SessionEntry::new( let entry = crate::node::session::SessionEntry::new(
*identity_b.node_addr(), *identity_b.node_addr(),
@@ -77,10 +74,7 @@ fn test_session_entry_touch() {
let identity_a = Identity::generate(); let identity_a = Identity::generate();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_initiator( let handshake = HandshakeState::new_initiator(identity_a.keypair(), identity_b.pubkey_full());
identity_a.keypair(),
identity_b.pubkey_full(),
);
let mut entry = crate::node::session::SessionEntry::new( let mut entry = crate::node::session::SessionEntry::new(
*identity_b.node_addr(), *identity_b.node_addr(),
@@ -102,10 +96,8 @@ fn test_session_table_operations() {
let mut node = make_node(); let mut node = make_node();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_initiator( let handshake =
node.identity().keypair(), HandshakeState::new_initiator(node.identity().keypair(), identity_b.pubkey_full());
identity_b.pubkey_full(),
);
let dest_addr = *identity_b.node_addr(); let dest_addr = *identity_b.node_addr();
let entry = crate::node::session::SessionEntry::new( let entry = crate::node::session::SessionEntry::new(
@@ -151,12 +143,14 @@ async fn test_session_direct_peer_handshake() {
// Node 0 should have a session in Initiating state // Node 0 should have a session in Initiating state
assert_eq!(nodes[0].node.session_count(), 1); assert_eq!(nodes[0].node.session_count(), 1);
assert!(nodes[0] assert!(
.node nodes[0]
.get_session(&node1_addr) .node
.unwrap() .get_session(&node1_addr)
.state() .unwrap()
.is_initiating()); .state()
.is_initiating()
);
// Process packets: SessionSetup arrives at Node 1 // Process packets: SessionSetup arrives at Node 1
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
@@ -165,12 +159,14 @@ async fn test_session_direct_peer_handshake() {
// Node 1 should now have a session in AwaitingMsg3 state (XK: identity not yet known) // Node 1 should now have a session in AwaitingMsg3 state (XK: identity not yet known)
assert_eq!(nodes[1].node.session_count(), 1); assert_eq!(nodes[1].node.session_count(), 1);
assert!(nodes[1] assert!(
.node nodes[1]
.get_session(&node0_addr) .node
.unwrap() .get_session(&node0_addr)
.state() .unwrap()
.is_awaiting_msg3()); .state()
.is_awaiting_msg3()
);
// Process packets: SessionAck arrives at Node 0, Node 0 sends SessionMsg3 // Process packets: SessionAck arrives at Node 0, Node 0 sends SessionMsg3
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
@@ -178,12 +174,14 @@ async fn test_session_direct_peer_handshake() {
assert!(count > 0, "Expected SessionAck packet to arrive"); assert!(count > 0, "Expected SessionAck packet to arrive");
// Node 0 should now be Established (transitions after sending msg3) // Node 0 should now be Established (transitions after sending msg3)
assert!(nodes[0] assert!(
.node nodes[0]
.get_session(&node1_addr) .node
.unwrap() .get_session(&node1_addr)
.state() .unwrap()
.is_established()); .state()
.is_established()
);
// Process packets: SessionMsg3 arrives at Node 1 // Process packets: SessionMsg3 arrives at Node 1
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
@@ -191,12 +189,14 @@ async fn test_session_direct_peer_handshake() {
assert!(count > 0, "Expected SessionMsg3 packet to arrive"); assert!(count > 0, "Expected SessionMsg3 packet to arrive");
// Node 1 should now be Established (transitions after processing msg3) // Node 1 should now be Established (transitions after processing msg3)
assert!(nodes[1] assert!(
.node nodes[1]
.get_session(&node0_addr) .node
.unwrap() .get_session(&node0_addr)
.state() .unwrap()
.is_established()); .state()
.is_established()
);
cleanup_nodes(&mut nodes).await; cleanup_nodes(&mut nodes).await;
} }
@@ -226,18 +226,22 @@ async fn test_session_direct_peer_data_transfer() {
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await; // Msg3 → Node 1 process_available_packets(&mut nodes).await; // Msg3 → Node 1
assert!(nodes[0] assert!(
.node nodes[0]
.get_session(&node1_addr) .node
.unwrap() .get_session(&node1_addr)
.state() .unwrap()
.is_established()); .state()
assert!(nodes[1] .is_established()
.node );
.get_session(&node0_addr) assert!(
.unwrap() nodes[1]
.state() .node
.is_established()); .get_session(&node0_addr)
.unwrap()
.state()
.is_established()
);
// Send data from Node 0 to Node 1 // Send data from Node 0 to Node 1
let test_data = b"Hello, FIPS session!"; let test_data = b"Hello, FIPS session!";
@@ -291,12 +295,14 @@ async fn test_session_3node_forwarded_handshake() {
nodes[2].node.get_session(&node0_addr).is_some(), nodes[2].node.get_session(&node0_addr).is_some(),
"Node 2 should have a session entry for Node 0" "Node 2 should have a session entry for Node 0"
); );
assert!(nodes[2] assert!(
.node nodes[2]
.get_session(&node0_addr) .node
.unwrap() .get_session(&node0_addr)
.state() .unwrap()
.is_awaiting_msg3()); .state()
.is_awaiting_msg3()
);
// Process: SessionAck: 2→1 (forwarded by transit B) // Process: SessionAck: 2→1 (forwarded by transit B)
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
@@ -307,12 +313,14 @@ async fn test_session_3node_forwarded_handshake() {
process_available_packets(&mut nodes).await; process_available_packets(&mut nodes).await;
// Node 0 should now be Established (transitions after sending msg3) // Node 0 should now be Established (transitions after sending msg3)
assert!(nodes[0] assert!(
.node nodes[0]
.get_session(&node2_addr) .node
.unwrap() .get_session(&node2_addr)
.state() .unwrap()
.is_established()); .state()
.is_established()
);
// Process: SessionMsg3: 0→1 (forwarded by transit B) // Process: SessionMsg3: 0→1 (forwarded by transit B)
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
@@ -323,12 +331,14 @@ async fn test_session_3node_forwarded_handshake() {
process_available_packets(&mut nodes).await; process_available_packets(&mut nodes).await;
// Node 2 should now be Established (transitions after processing msg3) // Node 2 should now be Established (transitions after processing msg3)
assert!(nodes[2] assert!(
.node nodes[2]
.get_session(&node0_addr) .node
.unwrap() .get_session(&node0_addr)
.state() .unwrap()
.is_established()); .state()
.is_established()
);
// Transit node B should NOT have a session // Transit node B should NOT have a session
assert_eq!( assert_eq!(
@@ -389,12 +399,14 @@ async fn test_session_3node_forwarded_data() {
} }
// Node 2 should be Established (transitioned during XK handshake msg3) // Node 2 should be Established (transitioned during XK handshake msg3)
assert!(nodes[2] assert!(
.node nodes[2]
.get_session(&node0_addr) .node
.unwrap() .get_session(&node0_addr)
.state() .unwrap()
.is_established()); .state()
.is_established()
);
cleanup_nodes(&mut nodes).await; cleanup_nodes(&mut nodes).await;
} }
@@ -520,12 +532,7 @@ async fn test_session_100_nodes() {
// Collect identities: (node_addr, pubkey) for all nodes // Collect identities: (node_addr, pubkey) for all nodes
let all_info: Vec<(NodeAddr, secp256k1::PublicKey)> = nodes let all_info: Vec<(NodeAddr, secp256k1::PublicKey)> = nodes
.iter() .iter()
.map(|tn| { .map(|tn| (*tn.node.node_addr(), tn.node.identity().pubkey_full()))
(
*tn.node.node_addr(),
tn.node.identity().pubkey_full(),
)
})
.collect(); .collect();
// Each node picks one random target for its outbound session. // Each node picks one random target for its outbound session.
@@ -640,11 +647,7 @@ async fn test_session_100_nodes() {
// (Responder should already be Established after XK msg3) // (Responder should already be Established after XK msg3)
let rev_payload = format!("rev-{}", pair_idx).into_bytes(); let rev_payload = format!("rev-{}", pair_idx).into_bytes();
let rev_ipv6 = build_ipv6_packet(&dst_fips, &src_fips, &rev_payload); let rev_ipv6 = build_ipv6_packet(&dst_fips, &src_fips, &rev_payload);
match nodes[dst] match nodes[dst].node.send_ipv6_packet(&src_addr, &rev_ipv6).await {
.node
.send_ipv6_packet(&src_addr, &rev_ipv6)
.await
{
Ok(()) => send_reverse_ok += 1, Ok(()) => send_reverse_ok += 1,
Err(_) => send_reverse_err += 1, Err(_) => send_reverse_err += 1,
} }
@@ -723,10 +726,7 @@ async fn test_session_100_nodes() {
} }
} }
let session_counts: Vec<usize> = nodes let session_counts: Vec<usize> = nodes.iter().map(|tn| tn.node.session_count()).collect();
.iter()
.map(|tn| tn.node.session_count())
.collect();
let total_sessions: usize = session_counts.iter().sum(); let total_sessions: usize = session_counts.iter().sum();
let min_sessions = *session_counts.iter().min().unwrap(); let min_sessions = *session_counts.iter().min().unwrap();
let max_sessions = *session_counts.iter().max().unwrap(); let max_sessions = *session_counts.iter().max().unwrap();
@@ -770,10 +770,8 @@ async fn test_session_100_nodes() {
}; };
// Coord cache stats // Coord cache stats
let coord_cache_sizes: Vec<usize> = nodes let coord_cache_sizes: Vec<usize> =
.iter() nodes.iter().map(|tn| tn.node.coord_cache().len()).collect();
.map(|tn| tn.node.coord_cache().len())
.collect();
let total_coord_entries: usize = coord_cache_sizes.iter().sum(); let total_coord_entries: usize = coord_cache_sizes.iter().sum();
let min_coord = *coord_cache_sizes.iter().min().unwrap(); let min_coord = *coord_cache_sizes.iter().min().unwrap();
let max_coord = *coord_cache_sizes.iter().max().unwrap(); let max_coord = *coord_cache_sizes.iter().max().unwrap();
@@ -884,10 +882,7 @@ async fn test_session_100_nodes() {
// === Assertions === // === Assertions ===
assert_eq!( assert_eq!(send_forward_err, 0, "All forward sends should succeed");
send_forward_err, 0,
"All forward sends should succeed"
);
assert_eq!( assert_eq!(
send_reverse_err, 0, send_reverse_err, 0,
"All reverse sends should succeed (responder Established after XK msg3)" "All reverse sends should succeed (responder Established after XK msg3)"
@@ -915,7 +910,11 @@ async fn test_session_100_nodes() {
// ============================================================================ // ============================================================================
/// Build a minimal valid IPv6 packet with given source and destination addresses. /// Build a minimal valid IPv6 packet with given source and destination addresses.
fn build_ipv6_packet(src: &crate::FipsAddress, dst: &crate::FipsAddress, payload: &[u8]) -> Vec<u8> { fn build_ipv6_packet(
src: &crate::FipsAddress,
dst: &crate::FipsAddress,
payload: &[u8],
) -> Vec<u8> {
let payload_len = payload.len() as u16; let payload_len = payload.len() as u16;
let mut packet = vec![0u8; 40 + payload.len()]; let mut packet = vec![0u8; 40 + payload.len()];
// Version (6) + traffic class high nibble // Version (6) + traffic class high nibble
@@ -944,17 +943,14 @@ fn test_identity_cache_populated_on_promote() {
let transport_id = TransportId::new(1); let transport_id = TransportId::new(1);
let link_id = LinkId::new(1); let link_id = LinkId::new(1);
let (conn, peer_identity) = make_completed_connection( let (conn, peer_identity) = make_completed_connection(&mut node, link_id, transport_id, 1000);
&mut node,
link_id,
transport_id,
1000,
);
node.add_connection(conn).unwrap(); node.add_connection(conn).unwrap();
// Promote // Promote
let result = node.promote_connection(link_id, peer_identity, 2000).unwrap(); let result = node
.promote_connection(link_id, peer_identity, 2000)
.unwrap();
assert!(matches!(result, PromotionResult::Promoted(_))); assert!(matches!(result, PromotionResult::Promoted(_)));
// Identity cache should contain the peer // Identity cache should contain the peer
@@ -962,7 +958,10 @@ fn test_identity_cache_populated_on_promote() {
let mut prefix = [0u8; 15]; let mut prefix = [0u8; 15];
prefix.copy_from_slice(&peer_addr.as_bytes()[0..15]); prefix.copy_from_slice(&peer_addr.as_bytes()[0..15]);
let cached = node.lookup_by_fips_prefix(&prefix); let cached = node.lookup_by_fips_prefix(&prefix);
assert!(cached.is_some(), "Identity cache should contain promoted peer"); assert!(
cached.is_some(),
"Identity cache should contain promoted peer"
);
let (cached_addr, cached_pk) = cached.unwrap(); let (cached_addr, cached_pk) = cached.unwrap();
assert_eq!(cached_addr, peer_addr); assert_eq!(cached_addr, peer_addr);
assert_eq!(cached_pk, peer_identity.pubkey_full()); assert_eq!(cached_pk, peer_identity.pubkey_full());
@@ -986,7 +985,11 @@ async fn test_tun_outbound_established_session() {
let dst_fips = crate::FipsAddress::from_node_addr(&node1_addr); let dst_fips = crate::FipsAddress::from_node_addr(&node1_addr);
// Establish session (XK: 3 messages — Setup, Ack, Msg3) // Establish session (XK: 3 messages — Setup, Ack, Msg3)
nodes[0].node.initiate_session(node1_addr, node1_pubkey).await.unwrap(); nodes[0]
.node
.initiate_session(node1_addr, node1_pubkey)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await; // Setup → Node 1 process_available_packets(&mut nodes).await; // Setup → Node 1
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
@@ -994,7 +997,14 @@ async fn test_tun_outbound_established_session() {
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await; // Msg3 → Node 1 process_available_packets(&mut nodes).await; // Msg3 → Node 1
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established()); assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// Install TUN receiver on Node 1 // Install TUN receiver on Node 1
let (tun_tx, tun_rx) = std::sync::mpsc::channel(); let (tun_tx, tun_rx) = std::sync::mpsc::channel();
@@ -1013,7 +1023,10 @@ async fn test_tun_outbound_established_session() {
// Verify plaintext arrived at Node 1's TUN // Verify plaintext arrived at Node 1's TUN
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect(); let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(delivered.len(), 1, "Exactly one packet should be delivered"); assert_eq!(delivered.len(), 1, "Exactly one packet should be delivered");
assert_eq!(delivered[0], ipv6_packet, "Delivered packet should match original"); assert_eq!(
delivered[0], ipv6_packet,
"Delivered packet should match original"
);
cleanup_nodes(&mut nodes).await; cleanup_nodes(&mut nodes).await;
} }
@@ -1049,17 +1062,35 @@ async fn test_tun_outbound_triggers_session_initiation() {
// Session should now be initiating // Session should now be initiating
assert_eq!(nodes[0].node.session_count(), 1); assert_eq!(nodes[0].node.session_count(), 1);
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_initiating()); assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_initiating()
);
// Drain packets until session established and queued packet delivered // Drain packets until session established and queued packet delivered
drain_to_quiescence(&mut nodes).await; drain_to_quiescence(&mut nodes).await;
// Session should be established on Node 0 // Session should be established on Node 0
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established()); assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// Verify the queued packet was delivered to Node 1 // Verify the queued packet was delivered to Node 1
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect(); let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(delivered.len(), 1, "Queued packet should be delivered after handshake"); assert_eq!(
delivered.len(),
1,
"Queued packet should be delivered after handshake"
);
assert_eq!(delivered[0], ipv6_packet); assert_eq!(delivered[0], ipv6_packet);
cleanup_nodes(&mut nodes).await; cleanup_nodes(&mut nodes).await;
@@ -1087,12 +1118,19 @@ async fn test_tun_outbound_unknown_destination() {
// Should receive ICMPv6 Destination Unreachable back on TUN // Should receive ICMPv6 Destination Unreachable back on TUN
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect(); let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(delivered.len(), 1, "Should receive ICMPv6 Destination Unreachable"); assert_eq!(
delivered.len(),
1,
"Should receive ICMPv6 Destination Unreachable"
);
// Verify it's an ICMPv6 Destination Unreachable (type 1, code 0) // Verify it's an ICMPv6 Destination Unreachable (type 1, code 0)
// ICMPv6 header starts at byte 40, type at byte 40, code at byte 41 // ICMPv6 header starts at byte 40, type at byte 40, code at byte 41
assert!(delivered[0].len() >= 48, "ICMPv6 response too short"); assert!(delivered[0].len() >= 48, "ICMPv6 response too short");
assert_eq!(delivered[0][6], 58, "Next header should be ICMPv6 (58)"); assert_eq!(delivered[0][6], 58, "Next header should be ICMPv6 (58)");
assert_eq!(delivered[0][40], 1, "ICMPv6 type should be Destination Unreachable (1)"); assert_eq!(
delivered[0][40], 1,
"ICMPv6 type should be Destination Unreachable (1)"
);
assert_eq!(delivered[0][41], 0, "ICMPv6 code should be No Route (0)"); assert_eq!(delivered[0][41], 0, "ICMPv6 code should be No Route (0)");
cleanup_nodes(&mut nodes).await; cleanup_nodes(&mut nodes).await;
@@ -1131,7 +1169,14 @@ async fn test_tun_outbound_3node_forwarded() {
drain_to_quiescence(&mut nodes).await; drain_to_quiescence(&mut nodes).await;
// Session should be established // Session should be established
assert!(nodes[0].node.get_session(&node2_addr).unwrap().state().is_established()); assert!(
nodes[0]
.node
.get_session(&node2_addr)
.unwrap()
.state()
.is_established()
);
// Verify packet delivered to Node 2 // Verify packet delivered to Node 2
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect(); let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
@@ -1170,16 +1215,34 @@ async fn test_tun_outbound_pending_queue_flush() {
// First packet triggers session initiation, rest are queued // First packet triggers session initiation, rest are queued
assert_eq!(nodes[0].node.session_count(), 1); assert_eq!(nodes[0].node.session_count(), 1);
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_initiating()); assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_initiating()
);
// Drain until session established and queued packets flushed // Drain until session established and queued packets flushed
drain_to_quiescence(&mut nodes).await; drain_to_quiescence(&mut nodes).await;
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established()); assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// All 5 packets should have been delivered // All 5 packets should have been delivered
let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect(); let delivered: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(delivered.len(), 5, "All 5 queued packets should be delivered"); assert_eq!(
delivered.len(),
5,
"All 5 queued packets should be delivered"
);
for (i, pkt) in delivered.iter().enumerate() { for (i, pkt) in delivered.iter().enumerate() {
assert_eq!(*pkt, packets[i], "Packet {} should match", i); assert_eq!(*pkt, packets[i], "Packet {} should match", i);
} }
@@ -1198,10 +1261,8 @@ fn make_noise_session(
) -> crate::noise::NoiseSession { ) -> crate::noise::NoiseSession {
use crate::noise::HandshakeState; use crate::noise::HandshakeState;
let mut initiator = HandshakeState::new_initiator( let mut initiator =
our_identity.keypair(), HandshakeState::new_initiator(our_identity.keypair(), remote_identity.pubkey_full());
remote_identity.pubkey_full(),
);
let mut responder = HandshakeState::new_responder(remote_identity.keypair()); let mut responder = HandshakeState::new_responder(remote_identity.keypair());
// Set epochs for both sides (required for handshake message encryption) // Set epochs for both sides (required for handshake message encryption)
@@ -1270,7 +1331,11 @@ fn test_purge_idle_sessions_keeps_active() {
let now_ms = 92_000; let now_ms = 92_000;
node.purge_idle_sessions(now_ms); node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 1, "Active session should survive purge"); assert_eq!(
node.session_count(),
1,
"Active session should survive purge"
);
} }
#[test] #[test]
@@ -1281,10 +1346,7 @@ fn test_purge_idle_sessions_ignores_initiating() {
let remote = Identity::generate(); let remote = Identity::generate();
let remote_addr = *remote.node_addr(); let remote_addr = *remote.node_addr();
let handshake = HandshakeState::new_initiator( let handshake = HandshakeState::new_initiator(node.identity().keypair(), remote.pubkey_full());
node.identity().keypair(),
remote.pubkey_full(),
);
let entry = crate::node::session::SessionEntry::new( let entry = crate::node::session::SessionEntry::new(
remote_addr, remote_addr,
remote.pubkey_full(), remote.pubkey_full(),
@@ -1299,7 +1361,11 @@ fn test_purge_idle_sessions_ignores_initiating() {
let now_ms = 1000 + 200_000; let now_ms = 1000 + 200_000;
node.purge_idle_sessions(now_ms); node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 1, "Initiating session should not be purged by idle timeout"); assert_eq!(
node.session_count(),
1,
"Initiating session should not be purged by idle timeout"
);
} }
#[test] #[test]
@@ -1330,8 +1396,10 @@ fn test_purge_idle_sessions_cleans_pending_packets() {
node.purge_idle_sessions(now_ms); node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 0); assert_eq!(node.session_count(), 0);
assert!(!node.pending_tun_packets.contains_key(&remote_addr), assert!(
"Pending packets should be cleaned up with idle session"); !node.pending_tun_packets.contains_key(&remote_addr),
"Pending packets should be cleaned up with idle session"
);
} }
#[test] #[test]
@@ -1357,7 +1425,11 @@ fn test_purge_idle_sessions_disabled_when_zero() {
let now_ms = 1000 + 1_000_000; let now_ms = 1000 + 1_000_000;
node.purge_idle_sessions(now_ms); node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 1, "Sessions should not be purged when idle timeout is disabled"); assert_eq!(
node.session_count(),
1,
"Sessions should not be purged when idle timeout is disabled"
);
} }
#[test] #[test]
@@ -1386,8 +1458,11 @@ fn test_purge_idle_sessions_mmp_activity_does_not_prevent_purge() {
let now_ms = 92_000; let now_ms = 92_000;
node.purge_idle_sessions(now_ms); node.purge_idle_sessions(now_ms);
assert_eq!(node.session_count(), 0, assert_eq!(
"Session with MMP-only activity should be purged"); node.session_count(),
0,
"Session with MMP-only activity should be purged"
);
} }
// ============================================================================ // ============================================================================
@@ -1401,10 +1476,7 @@ fn test_coords_warmup_counter_default_zero_on_new() {
let identity_a = Identity::generate(); let identity_a = Identity::generate();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_initiator( let handshake = HandshakeState::new_initiator(identity_a.keypair(), identity_b.pubkey_full());
identity_a.keypair(),
identity_b.pubkey_full(),
);
let entry = crate::node::session::SessionEntry::new( let entry = crate::node::session::SessionEntry::new(
*identity_b.node_addr(), *identity_b.node_addr(),
@@ -1414,8 +1486,11 @@ fn test_coords_warmup_counter_default_zero_on_new() {
true, true,
); );
assert_eq!(entry.coords_warmup_remaining(), 0, assert_eq!(
"Counter should be 0 for non-Established sessions"); entry.coords_warmup_remaining(),
0,
"Counter should be 0 for non-Established sessions"
);
} }
#[test] #[test]
@@ -1466,15 +1541,20 @@ fn test_coords_warmup_counter_decrement() {
assert_eq!(entry.coords_warmup_remaining(), expected); assert_eq!(entry.coords_warmup_remaining(), expected);
} }
assert_eq!(entry.coords_warmup_remaining(), 0, assert_eq!(
"Counter should reach 0 after N decrements"); entry.coords_warmup_remaining(),
0,
"Counter should reach 0 after N decrements"
);
} }
#[test] #[test]
fn test_coords_warmup_config_default() { fn test_coords_warmup_config_default() {
let config = crate::config::Config::new(); let config = crate::config::Config::new();
assert_eq!(config.node.session.coords_warmup_packets, 5, assert_eq!(
"Default coords_warmup_packets should be 5"); config.node.session.coords_warmup_packets, 5,
"Default coords_warmup_packets should be 5"
);
} }
// ============================================================================ // ============================================================================
@@ -1493,11 +1573,13 @@ fn test_identity_cache_lru_eviction() {
// Insert first two with explicit timestamps to ensure deterministic ordering // Insert first two with explicit timestamps to ensure deterministic ordering
let mut prefix1 = [0u8; 15]; let mut prefix1 = [0u8; 15];
prefix1.copy_from_slice(&id1.node_addr().as_bytes()[0..15]); prefix1.copy_from_slice(&id1.node_addr().as_bytes()[0..15]);
node.identity_cache.insert(prefix1, (*id1.node_addr(), id1.pubkey_full(), 1000)); node.identity_cache
.insert(prefix1, (*id1.node_addr(), id1.pubkey_full(), 1000));
let mut prefix2 = [0u8; 15]; let mut prefix2 = [0u8; 15];
prefix2.copy_from_slice(&id2.node_addr().as_bytes()[0..15]); prefix2.copy_from_slice(&id2.node_addr().as_bytes()[0..15]);
node.identity_cache.insert(prefix2, (*id2.node_addr(), id2.pubkey_full(), 2000)); node.identity_cache
.insert(prefix2, (*id2.node_addr(), id2.pubkey_full(), 2000));
assert_eq!(node.identity_cache_len(), 2); assert_eq!(node.identity_cache_len(), 2);
@@ -1505,13 +1587,17 @@ fn test_identity_cache_lru_eviction() {
node.register_identity(*id3.node_addr(), id3.pubkey_full()); node.register_identity(*id3.node_addr(), id3.pubkey_full());
assert_eq!(node.identity_cache_len(), 2); assert_eq!(node.identity_cache_len(), 2);
assert!(node.lookup_by_fips_prefix(&prefix1).is_none(), assert!(
"Oldest entry should have been evicted"); node.lookup_by_fips_prefix(&prefix1).is_none(),
"Oldest entry should have been evicted"
);
let mut prefix3 = [0u8; 15]; let mut prefix3 = [0u8; 15];
prefix3.copy_from_slice(&id3.node_addr().as_bytes()[0..15]); prefix3.copy_from_slice(&id3.node_addr().as_bytes()[0..15]);
assert!(node.lookup_by_fips_prefix(&prefix3).is_some(), assert!(
"Newest entry should be present"); node.lookup_by_fips_prefix(&prefix3).is_some(),
"Newest entry should be present"
);
} }
#[test] #[test]
@@ -1546,10 +1632,7 @@ fn test_session_entry_handshake_payload_storage() {
let identity_a = Identity::generate(); let identity_a = Identity::generate();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_initiator( let handshake = HandshakeState::new_initiator(identity_a.keypair(), identity_b.pubkey_full());
identity_a.keypair(),
identity_b.pubkey_full(),
);
let mut entry = crate::node::session::SessionEntry::new( let mut entry = crate::node::session::SessionEntry::new(
*identity_b.node_addr(), *identity_b.node_addr(),
@@ -1581,10 +1664,7 @@ fn test_session_entry_resend_tracking() {
let identity_a = Identity::generate(); let identity_a = Identity::generate();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_initiator( let handshake = HandshakeState::new_initiator(identity_a.keypair(), identity_b.pubkey_full());
identity_a.keypair(),
identity_b.pubkey_full(),
);
let mut entry = crate::node::session::SessionEntry::new( let mut entry = crate::node::session::SessionEntry::new(
*identity_b.node_addr(), *identity_b.node_addr(),
@@ -1615,10 +1695,7 @@ fn test_session_entry_clear_handshake_payload() {
let identity_a = Identity::generate(); let identity_a = Identity::generate();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_initiator( let handshake = HandshakeState::new_initiator(identity_a.keypair(), identity_b.pubkey_full());
identity_a.keypair(),
identity_b.pubkey_full(),
);
let mut entry = crate::node::session::SessionEntry::new( let mut entry = crate::node::session::SessionEntry::new(
*identity_b.node_addr(), *identity_b.node_addr(),
@@ -1649,10 +1726,8 @@ async fn test_session_handshake_timeout() {
let mut node = make_node(); let mut node = make_node();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_initiator( let handshake =
node.identity.keypair(), HandshakeState::new_initiator(node.identity.keypair(), identity_b.pubkey_full());
identity_b.pubkey_full(),
);
let dest_addr = *identity_b.node_addr(); let dest_addr = *identity_b.node_addr();
@@ -1672,12 +1747,18 @@ async fn test_session_handshake_timeout() {
let timeout_secs = node.config.node.rate_limit.handshake_timeout_secs; let timeout_secs = node.config.node.rate_limit.handshake_timeout_secs;
let before_timeout = 1000 + timeout_secs * 1000 - 1; let before_timeout = 1000 + timeout_secs * 1000 - 1;
node.resend_pending_session_handshakes(before_timeout).await; node.resend_pending_session_handshakes(before_timeout).await;
assert!(node.sessions.contains_key(&dest_addr), "Session should survive before timeout"); assert!(
node.sessions.contains_key(&dest_addr),
"Session should survive before timeout"
);
// After timeout: session should be removed // After timeout: session should be removed
let after_timeout = 1000 + timeout_secs * 1000 + 1; let after_timeout = 1000 + timeout_secs * 1000 + 1;
node.resend_pending_session_handshakes(after_timeout).await; node.resend_pending_session_handshakes(after_timeout).await;
assert!(!node.sessions.contains_key(&dest_addr), "Timed-out session should be removed"); assert!(
!node.sessions.contains_key(&dest_addr),
"Timed-out session should be removed"
);
} }
/// Test that session handshake timeout removes stale AwaitingMsg3 sessions. /// Test that session handshake timeout removes stale AwaitingMsg3 sessions.
@@ -1690,9 +1771,7 @@ async fn test_session_awaiting_msg3_timeout() {
let identity_a = Identity::generate(); let identity_a = Identity::generate();
let identity_b = Identity::generate(); let identity_b = Identity::generate();
let handshake = HandshakeState::new_xk_responder( let handshake = HandshakeState::new_xk_responder(identity_b.keypair());
identity_b.keypair(),
);
let src_addr = *identity_a.node_addr(); let src_addr = *identity_a.node_addr();
@@ -1712,7 +1791,10 @@ async fn test_session_awaiting_msg3_timeout() {
let timeout_secs = node.config.node.rate_limit.handshake_timeout_secs; let timeout_secs = node.config.node.rate_limit.handshake_timeout_secs;
let after_timeout = 1000 + timeout_secs * 1000 + 1; let after_timeout = 1000 + timeout_secs * 1000 + 1;
node.resend_pending_session_handshakes(after_timeout).await; node.resend_pending_session_handshakes(after_timeout).await;
assert!(!node.sessions.contains_key(&src_addr), "Timed-out AwaitingMsg3 session should be removed"); assert!(
!node.sessions.contains_key(&src_addr),
"Timed-out AwaitingMsg3 session should be removed"
);
} }
#[tokio::test] #[tokio::test]
@@ -1734,7 +1816,11 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
let dst_fips = crate::FipsAddress::from_node_addr(&node1_addr); let dst_fips = crate::FipsAddress::from_node_addr(&node1_addr);
// Establish session (XK: 3 messages — Setup, Ack, Msg3) // Establish session (XK: 3 messages — Setup, Ack, Msg3)
nodes[0].node.initiate_session(node1_addr, node1_pubkey).await.unwrap(); nodes[0]
.node
.initiate_session(node1_addr, node1_pubkey)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await; process_available_packets(&mut nodes).await;
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
@@ -1742,7 +1828,14 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
tokio::time::sleep(Duration::from_millis(20)).await; tokio::time::sleep(Duration::from_millis(20)).await;
process_available_packets(&mut nodes).await; process_available_packets(&mut nodes).await;
assert!(nodes[0].node.get_session(&node1_addr).unwrap().state().is_established()); assert!(
nodes[0]
.node
.get_session(&node1_addr)
.unwrap()
.state()
.is_established()
);
// Simulate receipt of MtuExceeded by reducing PathMtuState to a value // Simulate receipt of MtuExceeded by reducing PathMtuState to a value
// lower than the local transport MTU. // lower than the local transport MTU.
@@ -1751,7 +1844,8 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
{ {
let entry = nodes[0].node.get_session_mut(&node1_addr).unwrap(); let entry = nodes[0].node.get_session_mut(&node1_addr).unwrap();
let mmp = entry.mmp_mut().unwrap(); let mmp = entry.mmp_mut().unwrap();
mmp.path_mtu.apply_notification(reduced_mtu, std::time::Instant::now()); mmp.path_mtu
.apply_notification(reduced_mtu, std::time::Instant::now());
assert_eq!(mmp.path_mtu.current_mtu(), reduced_mtu); assert_eq!(mmp.path_mtu.current_mtu(), reduced_mtu);
} }
@@ -1764,14 +1858,24 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
let local_ipv6_mtu = nodes[0].node.effective_ipv6_mtu() as usize; let local_ipv6_mtu = nodes[0].node.effective_ipv6_mtu() as usize;
let oversized_payload = vec![0u8; reduced_ipv6_mtu - 39]; // 40-byte hdr + payload > reduced MTU let oversized_payload = vec![0u8; reduced_ipv6_mtu - 39]; // 40-byte hdr + payload > reduced MTU
let ipv6_packet = build_ipv6_packet(&src_fips, &dst_fips, &oversized_payload); let ipv6_packet = build_ipv6_packet(&src_fips, &dst_fips, &oversized_payload);
assert!(ipv6_packet.len() > reduced_ipv6_mtu, "packet must exceed path MTU"); assert!(
assert!(ipv6_packet.len() <= local_ipv6_mtu, "packet must fit local MTU"); ipv6_packet.len() > reduced_ipv6_mtu,
"packet must exceed path MTU"
);
assert!(
ipv6_packet.len() <= local_ipv6_mtu,
"packet must fit local MTU"
);
nodes[0].node.handle_tun_outbound(ipv6_packet).await; nodes[0].node.handle_tun_outbound(ipv6_packet).await;
// Verify ICMPv6 Packet Too Big was generated // Verify ICMPv6 Packet Too Big was generated
let ptb_messages: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect(); let ptb_messages: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx.try_recv().ok()).collect();
assert_eq!(ptb_messages.len(), 1, "Should generate exactly one ICMPv6 PTB"); assert_eq!(
ptb_messages.len(),
1,
"Should generate exactly one ICMPv6 PTB"
);
let ptb = &ptb_messages[0]; let ptb = &ptb_messages[0];
assert_eq!(ptb[0] >> 4, 6, "Should be IPv6"); assert_eq!(ptb[0] >> 4, 6, "Should be IPv6");
@@ -1784,12 +1888,23 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
// address, causing a PMTUD blackhole. // address, causing a PMTUD blackhole.
let ptb_src = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[8..24]).unwrap()); let ptb_src = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[8..24]).unwrap());
let ptb_dst = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[24..40]).unwrap()); let ptb_dst = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[24..40]).unwrap());
assert_eq!(ptb_src, dst_fips.to_ipv6(), "PTB source must be remote peer (original dst), not local node"); assert_eq!(
assert_eq!(ptb_dst, src_fips.to_ipv6(), "PTB destination must be local node (original src)"); ptb_src,
dst_fips.to_ipv6(),
"PTB source must be remote peer (original dst), not local node"
);
assert_eq!(
ptb_dst,
src_fips.to_ipv6(),
"PTB destination must be local node (original src)"
);
// Verify reported MTU (32-bit field at ICMPv6 header bytes 4-7) // Verify reported MTU (32-bit field at ICMPv6 header bytes 4-7)
let reported_mtu = u32::from_be_bytes([ptb[44], ptb[45], ptb[46], ptb[47]]); let reported_mtu = u32::from_be_bytes([ptb[44], ptb[45], ptb[46], ptb[47]]);
assert_eq!(reported_mtu, reduced_ipv6_mtu as u32, "Reported MTU should match path IPv6 MTU"); assert_eq!(
reported_mtu, reduced_ipv6_mtu as u32,
"Reported MTU should match path IPv6 MTU"
);
// Verify a packet that fits within path MTU passes through (no PTB) // Verify a packet that fits within path MTU passes through (no PTB)
let (tun_tx2, tun_rx2) = std::sync::mpsc::channel(); let (tun_tx2, tun_rx2) = std::sync::mpsc::channel();
@@ -1802,7 +1917,11 @@ async fn test_tun_outbound_path_mtu_generates_ptb() {
// No PTB should be generated for a fitting packet // No PTB should be generated for a fitting packet
let ptb_messages2: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx2.try_recv().ok()).collect(); let ptb_messages2: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx2.try_recv().ok()).collect();
assert_eq!(ptb_messages2.len(), 0, "Should not generate PTB for fitting packet"); assert_eq!(
ptb_messages2.len(),
0,
"Should not generate PTB for fitting packet"
);
cleanup_nodes(&mut nodes).await; cleanup_nodes(&mut nodes).await;
} }
@@ -1845,10 +1964,19 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
nodes[0].node.register_identity(node2_addr, node2_pubkey); nodes[0].node.register_identity(node2_addr, node2_pubkey);
// Establish session A→C via B (triggers routing through tree) // Establish session A→C via B (triggers routing through tree)
nodes[0].node.initiate_session(node2_addr, node2_pubkey).await.unwrap(); nodes[0]
.node
.initiate_session(node2_addr, node2_pubkey)
.await
.unwrap();
drain_to_quiescence(&mut nodes).await; drain_to_quiescence(&mut nodes).await;
assert!( assert!(
nodes[0].node.get_session(&node2_addr).unwrap().state().is_established(), nodes[0]
.node
.get_session(&node2_addr)
.unwrap()
.state()
.is_established(),
"Session A→C should be established" "Session A→C should be established"
); );
@@ -1858,7 +1986,11 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
// With coords (~66 extra), the wire could exceed B's recv buffer. // With coords (~66 extra), the wire could exceed B's recv buffer.
for _ in 0..5 { for _ in 0..5 {
let small = build_ipv6_packet(&src_fips, &dst_fips, &[0u8; 10]); let small = build_ipv6_packet(&src_fips, &dst_fips, &[0u8; 10]);
nodes[0].node.send_ipv6_packet(&node2_addr, &small).await.unwrap(); nodes[0]
.node
.send_ipv6_packet(&node2_addr, &small)
.await
.unwrap();
} }
drain_to_quiescence(&mut nodes).await; drain_to_quiescence(&mut nodes).await;
@@ -1872,12 +2004,17 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
assert!( assert!(
ipv6_packet.len() <= local_effective_mtu, ipv6_packet.len() <= local_effective_mtu,
"packet ({}) must fit A's local MTU ({})", "packet ({}) must fit A's local MTU ({})",
ipv6_packet.len(), local_effective_mtu ipv6_packet.len(),
local_effective_mtu
); );
// Send the oversized packet — B should fail to forward and send // Send the oversized packet — B should fail to forward and send
// MtuExceeded signal back. // MtuExceeded signal back.
nodes[0].node.send_ipv6_packet(&node2_addr, &ipv6_packet).await.unwrap(); nodes[0]
.node
.send_ipv6_packet(&node2_addr, &ipv6_packet)
.await
.unwrap();
drain_to_quiescence(&mut nodes).await; drain_to_quiescence(&mut nodes).await;
// Verify PathMtuState was updated on A // Verify PathMtuState was updated on A
@@ -1902,7 +2039,8 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
let ptb_messages: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx2.try_recv().ok()).collect(); let ptb_messages: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx2.try_recv().ok()).collect();
assert_eq!( assert_eq!(
ptb_messages.len(), 1, ptb_messages.len(),
1,
"Should generate ICMPv6 PTB for oversized packet after PathMtuState update" "Should generate ICMPv6 PTB for oversized packet after PathMtuState update"
); );
@@ -1917,8 +2055,16 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
// address, causing a PMTUD blackhole. // address, causing a PMTUD blackhole.
let ptb_src = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[8..24]).unwrap()); let ptb_src = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[8..24]).unwrap());
let ptb_dst = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[24..40]).unwrap()); let ptb_dst = std::net::Ipv6Addr::from(<[u8; 16]>::try_from(&ptb[24..40]).unwrap());
assert_eq!(ptb_src, dst_fips.to_ipv6(), "PTB source must be remote peer (original dst), not local node"); assert_eq!(
assert_eq!(ptb_dst, src_fips.to_ipv6(), "PTB destination must be local node (original src)"); ptb_src,
dst_fips.to_ipv6(),
"PTB source must be remote peer (original dst), not local node"
);
assert_eq!(
ptb_dst,
src_fips.to_ipv6(),
"PTB destination must be local node (original src)"
);
// Verify reported MTU is the path MTU (not local MTU) // Verify reported MTU is the path MTU (not local MTU)
let reported_mtu = u32::from_be_bytes([ptb[44], ptb[45], ptb[46], ptb[47]]); let reported_mtu = u32::from_be_bytes([ptb[44], ptb[45], ptb[46], ptb[47]]);
@@ -1941,7 +2087,8 @@ async fn test_multihop_pmtud_heterogeneous_mtu() {
let ptb_messages3: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx3.try_recv().ok()).collect(); let ptb_messages3: Vec<Vec<u8>> = std::iter::from_fn(|| tun_rx3.try_recv().ok()).collect();
assert_eq!( assert_eq!(
ptb_messages3.len(), 0, ptb_messages3.len(),
0,
"Should not generate PTB for packet fitting within path MTU" "Should not generate PTB for packet fitting within path MTU"
); );

View File

@@ -69,7 +69,9 @@ pub(super) async fn initiate_handshake(nodes: &mut [TestNode], i: usize, j: usiz
let our_index = initiator.node.index_allocator.allocate().unwrap(); let our_index = initiator.node.index_allocator.allocate().unwrap();
let our_keypair = initiator.node.identity().keypair(); let our_keypair = initiator.node.identity().keypair();
let noise_msg1 = conn.start_handshake(our_keypair, initiator.node.startup_epoch, 1000).unwrap(); let noise_msg1 = conn
.start_handshake(our_keypair, initiator.node.startup_epoch, 1000)
.unwrap();
conn.set_our_index(our_index); conn.set_our_index(our_index);
conn.set_transport_id(transport_id); conn.set_transport_id(transport_id);
conn.set_source_addr(responder_addr.clone()); conn.set_source_addr(responder_addr.clone());
@@ -184,7 +186,12 @@ pub(super) fn print_tree_snapshot(label: &str, nodes: &[TestNode]) {
.count(); .count();
eprintln!( eprintln!(
" node[{}] root=node[{}] depth={} parent=node[{}] peers={} pending={}", " node[{}] root=node[{}] depth={} parent=node[{}] peers={} pending={}",
i, root_idx, ts.my_coords().depth(), parent_idx, tn.node.peer_count(), pending, i,
root_idx,
ts.my_coords().depth(),
parent_idx,
tn.node.peer_count(),
pending,
); );
} }
} else if correct_root_count < nodes.len() { } else if correct_root_count < nodes.len() {
@@ -209,7 +216,9 @@ pub(super) fn print_tree_snapshot(label: &str, nodes: &[TestNode]) {
/// ///
/// Returns the number of packets processed. /// Returns the number of packets processed.
pub(super) async fn process_available_packets(nodes: &mut [TestNode]) -> usize { pub(super) async fn process_available_packets(nodes: &mut [TestNode]) -> usize {
use crate::node::wire::{CommonPrefix, FMP_VERSION, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2, COMMON_PREFIX_SIZE}; use crate::node::wire::{
COMMON_PREFIX_SIZE, CommonPrefix, FMP_VERSION, PHASE_ESTABLISHED, PHASE_MSG1, PHASE_MSG2,
};
let mut count = 0; let mut count = 0;
for node in nodes.iter_mut() { for node in nodes.iter_mut() {
@@ -224,9 +233,7 @@ pub(super) async fn process_available_packets(nodes: &mut [TestNode]) -> usize {
match prefix.phase { match prefix.phase {
PHASE_MSG1 => node.node.handle_msg1(packet).await, PHASE_MSG1 => node.node.handle_msg1(packet).await,
PHASE_MSG2 => node.node.handle_msg2(packet).await, PHASE_MSG2 => node.node.handle_msg2(packet).await,
PHASE_ESTABLISHED => { PHASE_ESTABLISHED => node.node.handle_encrypted_frame(packet).await,
node.node.handle_encrypted_frame(packet).await
}
_ => {} _ => {}
} }
count += 1; count += 1;
@@ -319,7 +326,11 @@ pub(super) async fn drain_all_packets(nodes: &mut [TestNode], verbose: bool) ->
/// ///
/// First builds a random spanning tree to ensure connectivity, /// First builds a random spanning tree to ensure connectivity,
/// then adds extra edges up to the target count. /// then adds extra edges up to the target count.
pub(super) fn generate_random_edges(n: usize, target_edges: usize, seed: u64) -> Vec<(usize, usize)> { pub(super) fn generate_random_edges(
n: usize,
target_edges: usize,
seed: u64,
) -> Vec<(usize, usize)> {
use rand::rngs::StdRng; use rand::rngs::StdRng;
use rand::{RngExt, SeedableRng}; use rand::{RngExt, SeedableRng};
@@ -373,11 +384,7 @@ pub(super) fn verify_tree_convergence(nodes: &[TestNode]) {
assert!(n > 0); assert!(n > 0);
// Find the expected root (smallest NodeAddr across all nodes) // Find the expected root (smallest NodeAddr across all nodes)
let expected_root = nodes let expected_root = nodes.iter().map(|tn| *tn.node.node_addr()).min().unwrap();
.iter()
.map(|tn| *tn.node.node_addr())
.min()
.unwrap();
// All nodes should agree on the root // All nodes should agree on the root
for (i, tn) in nodes.iter().enumerate() { for (i, tn) in nodes.iter().enumerate() {
@@ -627,12 +634,16 @@ pub(super) async fn run_tree_test_with_mtus(
assert!( assert!(
nodes[i].node.get_peer(&j_addr).is_some(), nodes[i].node.get_peer(&j_addr).is_some(),
"Node {} should have peer {} (node {})", "Node {} should have peer {} (node {})",
i, j_addr, j i,
j_addr,
j
); );
assert!( assert!(
nodes[j].node.get_peer(&i_addr).is_some(), nodes[j].node.get_peer(&i_addr).is_some(),
"Node {} should have peer {} (node {})", "Node {} should have peer {} (node {})",
j, i_addr, i j,
i_addr,
i
); );
} }

View File

@@ -8,9 +8,9 @@
use super::*; use super::*;
use crate::config::TcpConfig; use crate::config::TcpConfig;
use crate::transport::tcp::TcpTransport; use crate::transport::tcp::TcpTransport;
use crate::transport::{packet_channel, TransportAddr, TransportHandle, TransportId}; use crate::transport::{TransportAddr, TransportHandle, TransportId, packet_channel};
use spanning_tree::{ use spanning_tree::{
cleanup_nodes, drain_all_packets, initiate_handshake, verify_tree_convergence, TestNode, TestNode, cleanup_nodes, drain_all_packets, initiate_handshake, verify_tree_convergence,
}; };
use std::time::Duration; use std::time::Duration;

View File

@@ -96,7 +96,10 @@ fn test_node_link_management() {
assert_eq!(node.link_count(), 0); assert_eq!(node.link_count(), 0);
// Lookup should be gone // Lookup should be gone
assert!(node.find_link_by_addr(TransportId::new(1), &TransportAddr::from_string("test")).is_none()); assert!(
node.find_link_by_addr(TransportId::new(1), &TransportAddr::from_string("test"))
.is_none()
);
} }
#[test] #[test]
@@ -183,8 +186,14 @@ fn test_node_promote_connection() {
let peer = node.get_peer(&node_addr).unwrap(); let peer = node.get_peer(&node_addr).unwrap();
assert_eq!(peer.authenticated_at(), 2000); assert_eq!(peer.authenticated_at(), 2000);
assert!(peer.has_session(), "Promoted peer should have NoiseSession"); assert!(peer.has_session(), "Promoted peer should have NoiseSession");
assert!(peer.our_index().is_some(), "Promoted peer should have our_index"); assert!(
assert!(peer.their_index().is_some(), "Promoted peer should have their_index"); peer.our_index().is_some(),
"Promoted peer should have our_index"
);
assert!(
peer.their_index().is_some(),
"Promoted peer should have their_index"
);
// Verify peers_by_index is populated // Verify peers_by_index is populated
let our_index = peer.our_index().unwrap(); let our_index = peer.our_index().unwrap();
@@ -201,8 +210,7 @@ fn test_node_cross_connection_resolution() {
// First connection and promotion (becomes active peer) // First connection and promotion (becomes active peer)
let link_id1 = LinkId::new(1); let link_id1 = LinkId::new(1);
let (conn1, identity) = let (conn1, identity) = make_completed_connection(&mut node, link_id1, transport_id, 1000);
make_completed_connection(&mut node, link_id1, transport_id, 1000);
let node_addr = *identity.node_addr(); let node_addr = *identity.node_addr();
node.add_connection(conn1).unwrap(); node.add_connection(conn1).unwrap();
@@ -236,8 +244,7 @@ fn test_node_peer_limit() {
// Add two peers via promotion // Add two peers via promotion
for i in 0..2 { for i in 0..2 {
let link_id = LinkId::new(i as u64 + 1); let link_id = LinkId::new(i as u64 + 1);
let (conn, identity) = let (conn, identity) = make_completed_connection(&mut node, link_id, transport_id, 1000);
make_completed_connection(&mut node, link_id, transport_id, 1000);
node.add_connection(conn).unwrap(); node.add_connection(conn).unwrap();
node.promote_connection(link_id, identity, 2000).unwrap(); node.promote_connection(link_id, identity, 2000).unwrap();
} }
@@ -246,8 +253,7 @@ fn test_node_peer_limit() {
// Third should fail // Third should fail
let link_id = LinkId::new(3); let link_id = LinkId::new(3);
let (conn, identity) = let (conn, identity) = make_completed_connection(&mut node, link_id, transport_id, 3000);
make_completed_connection(&mut node, link_id, transport_id, 3000);
node.add_connection(conn).unwrap(); node.add_connection(conn).unwrap();
let result = node.promote_connection(link_id, identity, 4000); let result = node.promote_connection(link_id, identity, 4000);
@@ -296,23 +302,20 @@ fn test_node_sendable_peers() {
// Add a healthy peer // Add a healthy peer
let link_id1 = LinkId::new(1); let link_id1 = LinkId::new(1);
let (conn1, identity1) = let (conn1, identity1) = make_completed_connection(&mut node, link_id1, transport_id, 1000);
make_completed_connection(&mut node, link_id1, transport_id, 1000);
let node_addr1 = *identity1.node_addr(); let node_addr1 = *identity1.node_addr();
node.add_connection(conn1).unwrap(); node.add_connection(conn1).unwrap();
node.promote_connection(link_id1, identity1, 2000).unwrap(); node.promote_connection(link_id1, identity1, 2000).unwrap();
// Add another peer and mark it stale (still sendable) // Add another peer and mark it stale (still sendable)
let link_id2 = LinkId::new(2); let link_id2 = LinkId::new(2);
let (conn2, identity2) = let (conn2, identity2) = make_completed_connection(&mut node, link_id2, transport_id, 1000);
make_completed_connection(&mut node, link_id2, transport_id, 1000);
node.add_connection(conn2).unwrap(); node.add_connection(conn2).unwrap();
node.promote_connection(link_id2, identity2, 2000).unwrap(); node.promote_connection(link_id2, identity2, 2000).unwrap();
// Add a third peer and mark it disconnected (not sendable) // Add a third peer and mark it disconnected (not sendable)
let link_id3 = LinkId::new(3); let link_id3 = LinkId::new(3);
let (conn3, identity3) = let (conn3, identity3) = make_completed_connection(&mut node, link_id3, transport_id, 1000);
make_completed_connection(&mut node, link_id3, transport_id, 1000);
let node_addr3 = *identity3.node_addr(); let node_addr3 = *identity3.node_addr();
node.add_connection(conn3).unwrap(); node.add_connection(conn3).unwrap();
node.promote_connection(link_id3, identity3, 2000).unwrap(); node.promote_connection(link_id3, identity3, 2000).unwrap();
@@ -345,14 +348,16 @@ fn test_node_pending_outbound_tracking() {
let index = node.index_allocator.allocate().unwrap(); let index = node.index_allocator.allocate().unwrap();
// Track in pending_outbound // Track in pending_outbound
node.pending_outbound.insert((transport_id, index.as_u32()), link_id); node.pending_outbound
.insert((transport_id, index.as_u32()), link_id);
// Verify we can look it up // Verify we can look it up
let found = node.pending_outbound.get(&(transport_id, index.as_u32())); let found = node.pending_outbound.get(&(transport_id, index.as_u32()));
assert_eq!(found, Some(&link_id)); assert_eq!(found, Some(&link_id));
// Clean up // Clean up
node.pending_outbound.remove(&(transport_id, index.as_u32())); node.pending_outbound
.remove(&(transport_id, index.as_u32()));
let _ = node.index_allocator.free(index); let _ = node.index_allocator.free(index);
assert_eq!(node.index_allocator.count(), 0); assert_eq!(node.index_allocator.count(), 0);
@@ -369,7 +374,8 @@ fn test_node_peers_by_index_tracking() {
let index = node.index_allocator.allocate().unwrap(); let index = node.index_allocator.allocate().unwrap();
// Track in peers_by_index // Track in peers_by_index
node.peers_by_index.insert((transport_id, index.as_u32()), node_addr); node.peers_by_index
.insert((transport_id, index.as_u32()), node_addr);
// Verify lookup // Verify lookup
let found = node.peers_by_index.get(&(transport_id, index.as_u32())); let found = node.peers_by_index.get(&(transport_id, index.as_u32()));
@@ -450,7 +456,9 @@ fn test_promote_cleans_up_pending_outbound_to_same_peer() {
PeerConnection::outbound(pending_link_id, peer_b_identity, pending_time_ms); PeerConnection::outbound(pending_link_id, peer_b_identity, pending_time_ms);
let our_keypair = node.identity.keypair(); let our_keypair = node.identity.keypair();
let _msg1 = pending_conn.start_handshake(our_keypair, node.startup_epoch, pending_time_ms).unwrap(); let _msg1 = pending_conn
.start_handshake(our_keypair, node.startup_epoch, pending_time_ms)
.unwrap();
let pending_index = node.index_allocator.allocate().unwrap(); let pending_index = node.index_allocator.allocate().unwrap();
pending_conn.set_our_index(pending_index); pending_conn.set_our_index(pending_index);
@@ -483,11 +491,8 @@ fn test_promote_cleans_up_pending_outbound_to_same_peer() {
let completing_link_id = LinkId::new(2); let completing_link_id = LinkId::new(2);
let completing_time_ms = 2000; let completing_time_ms = 2000;
let mut completing_conn = PeerConnection::outbound( let mut completing_conn =
completing_link_id, PeerConnection::outbound(completing_link_id, peer_b_identity, completing_time_ms);
peer_b_identity,
completing_time_ms,
);
let our_keypair = node.identity.keypair(); let our_keypair = node.identity.keypair();
let msg1 = completing_conn let msg1 = completing_conn
@@ -573,7 +578,10 @@ fn test_schedule_retry_creates_entry() {
assert_eq!(node.retry_pending.len(), 1); assert_eq!(node.retry_pending.len(), 1);
let state = node.retry_pending.get(&peer_node_addr).unwrap(); let state = node.retry_pending.get(&peer_node_addr).unwrap();
assert_eq!(state.retry_count, 1); assert_eq!(state.retry_count, 1);
assert!(state.reconnect, "Auto-connect peers always get reconnect=true"); assert!(
state.reconnect,
"Auto-connect peers always get reconnect=true"
);
// Default base = 5s, 2^1 = 10s, but first retry is 2^0... let me check: // Default base = 5s, 2^1 = 10s, but first retry is 2^0... let me check:
// retry_count is set to 1, backoff_ms(5000) = 5000 * 2^1 = 10000 // retry_count is set to 1, backoff_ms(5000) = 5000 * 2^1 = 10000
assert_eq!(state.retry_after_ms, 1000 + 10_000); assert_eq!(state.retry_after_ms, 1000 + 10_000);
@@ -597,7 +605,10 @@ fn test_schedule_retry_increments() {
// First failure // First failure
node.schedule_retry(peer_node_addr, 1000); node.schedule_retry(peer_node_addr, 1000);
assert_eq!(node.retry_pending.get(&peer_node_addr).unwrap().retry_count, 1); assert_eq!(
node.retry_pending.get(&peer_node_addr).unwrap().retry_count,
1
);
// Second failure // Second failure
node.schedule_retry(peer_node_addr, 11_000); node.schedule_retry(peer_node_addr, 11_000);
@@ -637,7 +648,10 @@ fn test_schedule_retry_auto_connect_never_exhausts() {
node.retry_pending.contains_key(&peer_node_addr), node.retry_pending.contains_key(&peer_node_addr),
"Auto-connect peers should never exhaust retries" "Auto-connect peers should never exhaust retries"
); );
assert_eq!(node.retry_pending.get(&peer_node_addr).unwrap().retry_count, 3); assert_eq!(
node.retry_pending.get(&peer_node_addr).unwrap().retry_count,
3
);
} }
/// Test that schedule_retry does nothing when max_retries is 0. /// Test that schedule_retry does nothing when max_retries is 0.
@@ -725,7 +739,7 @@ fn test_schedule_reconnect_preserves_backoff() {
let mut node = Node::new(config).unwrap(); let mut node = Node::new(config).unwrap();
// Simulate two stale handshake timeouts incrementing the retry count. // Simulate two stale handshake timeouts incrementing the retry count.
node.schedule_retry(peer_node_addr, 1_000); // count=1, delay=10s node.schedule_retry(peer_node_addr, 1_000); // count=1, delay=10s
node.schedule_retry(peer_node_addr, 11_000); // count=2, delay=20s node.schedule_retry(peer_node_addr, 11_000); // count=2, delay=20s
{ {
let state = node.retry_pending.get(&peer_node_addr).unwrap(); let state = node.retry_pending.get(&peer_node_addr).unwrap();
@@ -738,10 +752,7 @@ fn test_schedule_reconnect_preserves_backoff() {
node.schedule_reconnect(peer_node_addr, 31_000); node.schedule_reconnect(peer_node_addr, 31_000);
let state = node.retry_pending.get(&peer_node_addr).unwrap(); let state = node.retry_pending.get(&peer_node_addr).unwrap();
assert!( assert!(state.reconnect, "Entry should be marked as reconnect");
state.reconnect,
"Entry should be marked as reconnect"
);
assert_eq!( assert_eq!(
state.retry_count, 3, state.retry_count, 3,
"schedule_reconnect should increment existing count (was 2), not reset to 0 (regression: issue #5)" "schedule_reconnect should increment existing count (was 2), not reset to 0 (regression: issue #5)"
@@ -752,7 +763,8 @@ fn test_schedule_reconnect_preserves_backoff() {
let max_ms = node.config.node.retry.max_backoff_secs * 1000; let max_ms = node.config.node.retry.max_backoff_secs * 1000;
let expected_delay = state.backoff_ms(base_ms, max_ms); let expected_delay = state.backoff_ms(base_ms, max_ms);
assert_eq!( assert_eq!(
state.retry_after_ms, 31_000 + expected_delay, state.retry_after_ms,
31_000 + expected_delay,
"retry_after_ms should reflect count=3 backoff" "retry_after_ms should reflect count=3 backoff"
); );
} }
@@ -778,7 +790,10 @@ fn test_schedule_reconnect_fresh_state() {
let state = node.retry_pending.get(&peer_node_addr).unwrap(); let state = node.retry_pending.get(&peer_node_addr).unwrap();
assert!(state.reconnect, "Entry should be marked as reconnect"); assert!(state.reconnect, "Entry should be marked as reconnect");
assert_eq!(state.retry_count, 0, "Fresh reconnect should start at count=0"); assert_eq!(
state.retry_count, 0,
"Fresh reconnect should start at count=0"
);
// Base delay: 5s * 2^0 = 5s // Base delay: 5s * 2^0 = 5s
let base_ms = node.config.node.retry.base_interval_secs * 1000; let base_ms = node.config.node.retry.base_interval_secs * 1000;
let max_ms = node.config.node.retry.max_backoff_secs * 1000; let max_ms = node.config.node.retry.max_backoff_secs * 1000;

View File

@@ -5,8 +5,8 @@
use std::collections::HashMap; use std::collections::HashMap;
use crate::protocol::TreeAnnounce;
use crate::NodeAddr; use crate::NodeAddr;
use crate::protocol::TreeAnnounce;
use super::{Node, NodeError}; use super::{Node, NodeError};
use tracing::{debug, info, trace, warn}; use tracing::{debug, info, trace, warn};
@@ -105,7 +105,9 @@ impl Node {
let ready: Vec<NodeAddr> = self let ready: Vec<NodeAddr> = self
.peers .peers
.iter() .iter()
.filter(|(_, peer)| peer.has_pending_tree_announce() && peer.can_send_tree_announce(now_ms)) .filter(|(_, peer)| {
peer.has_pending_tree_announce() && peer.can_send_tree_announce(now_ms)
})
.map(|(addr, _)| *addr) .map(|(addr, _)| *addr)
.collect(); .collect();
@@ -185,10 +187,9 @@ impl Node {
} }
// Update in TreeState // Update in TreeState
let updated = self.tree_state.update_peer( let updated = self
announce.declaration.clone(), .tree_state
announce.ancestry.clone(), .update_peer(announce.declaration.clone(), announce.ancestry.clone());
);
if !updated { if !updated {
self.stats_mut().tree.stale += 1; self.stats_mut().tree.stale += 1;
@@ -214,7 +215,9 @@ impl Node {
// Re-evaluate parent selection with current link costs. // Re-evaluate parent selection with current link costs.
// Exclude peers without MMP RTT data — they are not yet eligible // Exclude peers without MMP RTT data — they are not yet eligible
// as parent candidates (prevents oscillation from optimistic defaults). // as parent candidates (prevents oscillation from optimistic defaults).
let peer_costs: HashMap<NodeAddr, f64> = self.peers.iter() let peer_costs: HashMap<NodeAddr, f64> = self
.peers
.iter()
.filter(|(_, peer)| peer.has_srtt()) .filter(|(_, peer)| peer.has_srtt())
.map(|(addr, peer)| (*addr, peer.link_cost())) .map(|(addr, peer)| (*addr, peer.link_cost()))
.collect(); .collect();
@@ -266,7 +269,9 @@ impl Node {
parent = %self.peer_display_name(from), parent = %self.peer_display_name(from),
"Parent ancestry contains us — loop detected, dropping parent" "Parent ancestry contains us — loop detected, dropping parent"
); );
let peer_costs: HashMap<NodeAddr, f64> = self.peers.iter() let peer_costs: HashMap<NodeAddr, f64> = self
.peers
.iter()
.filter(|(_, peer)| peer.has_srtt()) .filter(|(_, peer)| peer.has_srtt())
.map(|(addr, peer)| (*addr, peer.link_cost())) .map(|(addr, peer)| (*addr, peer.link_cost()))
.collect(); .collect();
@@ -276,7 +281,7 @@ impl Node {
return; return;
} }
self.coord_cache.clear(); self.coord_cache.clear();
self.reset_discovery_backoff(); self.reset_discovery_backoff();
self.send_tree_announce_to_all().await; self.send_tree_announce_to_all().await;
} }
return; return;
@@ -360,7 +365,9 @@ impl Node {
self.last_parent_reeval = Some(now); self.last_parent_reeval = Some(now);
let peer_costs: HashMap<NodeAddr, f64> = self.peers.iter() let peer_costs: HashMap<NodeAddr, f64> = self
.peers
.iter()
.filter(|(_, peer)| peer.has_srtt()) .filter(|(_, peer)| peer.has_srtt())
.map(|(addr, peer)| (*addr, peer.link_cost())) .map(|(addr, peer)| (*addr, peer.link_cost()))
.collect(); .collect();
@@ -411,14 +418,16 @@ impl Node {
/// ///
/// Returns `true` if our tree state changed (caller should announce). /// Returns `true` if our tree state changed (caller should announce).
pub(super) fn handle_peer_removal_tree_cleanup(&mut self, node_addr: &NodeAddr) -> bool { pub(super) fn handle_peer_removal_tree_cleanup(&mut self, node_addr: &NodeAddr) -> bool {
let was_parent = !self.tree_state.is_root() let was_parent =
&& self.tree_state.my_declaration().parent_id() == node_addr; !self.tree_state.is_root() && self.tree_state.my_declaration().parent_id() == node_addr;
self.tree_state.remove_peer(node_addr); self.tree_state.remove_peer(node_addr);
if was_parent { if was_parent {
self.stats_mut().tree.parent_losses += 1; self.stats_mut().tree.parent_losses += 1;
let peer_costs: HashMap<NodeAddr, f64> = self.peers.iter() let peer_costs: HashMap<NodeAddr, f64> = self
.peers
.iter()
.filter(|(_, peer)| peer.has_srtt()) .filter(|(_, peer)| peer.has_srtt())
.map(|(addr, peer)| (*addr, peer.link_cost())) .map(|(addr, peer)| (*addr, peer.link_cost()))
.collect(); .collect();

View File

@@ -17,8 +17,8 @@
//! | 0x1 | Noise IK msg1 | 114 bytes | Handshake initiation | //! | 0x1 | Noise IK msg1 | 114 bytes | Handshake initiation |
//! | 0x2 | Noise IK msg2 | 69 bytes | Handshake response | //! | 0x2 | Noise IK msg2 | 69 bytes | Handshake response |
use crate::utils::index::SessionIndex;
use crate::noise::{HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, TAG_SIZE}; use crate::noise::{HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, TAG_SIZE};
use crate::utils::index::SessionIndex;
// ============================================================================ // ============================================================================
// Constants // Constants
@@ -164,8 +164,7 @@ impl EncryptedHeader {
let payload_len = u16::from_le_bytes([data[2], data[3]]); let payload_len = u16::from_le_bytes([data[2], data[3]]);
let receiver_idx = SessionIndex::from_le_bytes([data[4], data[5], data[6], data[7]]); let receiver_idx = SessionIndex::from_le_bytes([data[4], data[5], data[6], data[7]]);
let counter = u64::from_le_bytes([ let counter = u64::from_le_bytes([
data[8], data[9], data[10], data[11], data[8], data[9], data[10], data[11], data[12], data[13], data[14], data[15],
data[12], data[13], data[14], data[15],
]); ]);
let mut header_bytes = [0u8; ESTABLISHED_HEADER_SIZE]; let mut header_bytes = [0u8; ESTABLISHED_HEADER_SIZE];
@@ -328,7 +327,11 @@ pub fn build_msg1(sender_idx: SessionIndex, noise_msg1: &[u8]) -> Vec<u8> {
/// Build a wire-format msg2 packet. /// Build a wire-format msg2 packet.
/// ///
/// Format: `[0x02][0x00][payload_len:2 LE][sender_idx:4 LE][receiver_idx:4 LE][noise_msg2:57]` /// Format: `[0x02][0x00][payload_len:2 LE][sender_idx:4 LE][receiver_idx:4 LE][noise_msg2:57]`
pub fn build_msg2(sender_idx: SessionIndex, receiver_idx: SessionIndex, noise_msg2: &[u8]) -> Vec<u8> { pub fn build_msg2(
sender_idx: SessionIndex,
receiver_idx: SessionIndex,
noise_msg2: &[u8],
) -> Vec<u8> {
debug_assert_eq!(noise_msg2.len(), HANDSHAKE_MSG2_SIZE); debug_assert_eq!(noise_msg2.len(), HANDSHAKE_MSG2_SIZE);
let payload_len = (4 + 4 + noise_msg2.len()) as u16; // sender + receiver + noise let payload_len = (4 + 4 + noise_msg2.len()) as u16; // sender + receiver + noise
@@ -542,8 +545,8 @@ mod tests {
#[test] #[test]
fn test_wire_sizes() { fn test_wire_sizes() {
assert_eq!(MSG1_WIRE_SIZE, 114); // 4 + 4 + 106 assert_eq!(MSG1_WIRE_SIZE, 114); // 4 + 4 + 106
assert_eq!(MSG2_WIRE_SIZE, 69); // 4 + 4 + 4 + 57 assert_eq!(MSG2_WIRE_SIZE, 69); // 4 + 4 + 4 + 57
assert_eq!(ENCRYPTED_MIN_SIZE, 32); // 16 + 16 assert_eq!(ENCRYPTED_MIN_SIZE, 32); // 16 + 16
assert_eq!(COMMON_PREFIX_SIZE, 4); assert_eq!(COMMON_PREFIX_SIZE, 4);
assert_eq!(ESTABLISHED_HEADER_SIZE, 16); assert_eq!(ESTABLISHED_HEADER_SIZE, 16);
@@ -585,22 +588,17 @@ mod tests {
#[test] #[test]
fn test_flags_byte() { fn test_flags_byte() {
let header = build_established_header( let header =
SessionIndex::new(1), build_established_header(SessionIndex::new(1), 0, FLAG_KEY_EPOCH | FLAG_SP, 100);
0,
FLAG_KEY_EPOCH | FLAG_SP,
100,
);
assert_eq!(header[1], 0x05); // bits 0 and 2 set assert_eq!(header[1], 0x05); // bits 0 and 2 set
let parsed = EncryptedHeader::parse(&[ let parsed = EncryptedHeader::parse(&[
header[0], header[1], header[2], header[3], header[0], header[1], header[2], header[3], header[4], header[5], header[6], header[7],
header[4], header[5], header[6], header[7], header[8], header[9], header[10], header[11], header[12], header[13], header[14],
header[8], header[9], header[10], header[11], header[15], // minimum: TAG_SIZE bytes of ciphertext
header[12], header[13], header[14], header[15],
// minimum: TAG_SIZE bytes of ciphertext
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
]).unwrap(); ])
.unwrap();
assert_eq!(parsed.flags & FLAG_KEY_EPOCH, FLAG_KEY_EPOCH); assert_eq!(parsed.flags & FLAG_KEY_EPOCH, FLAG_KEY_EPOCH);
assert_eq!(parsed.flags & FLAG_CE, 0); assert_eq!(parsed.flags & FLAG_CE, 0);
assert_eq!(parsed.flags & FLAG_SP, FLAG_SP); assert_eq!(parsed.flags & FLAG_SP, FLAG_SP);

View File

@@ -1,12 +1,12 @@
use super::{ use super::{
CipherState, HandshakeProgress, HandshakeRole, NoiseError, NoisePattern, NoiseSession, CipherState, EPOCH_ENCRYPTED_SIZE, EPOCH_SIZE, HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE,
EPOCH_ENCRYPTED_SIZE, EPOCH_SIZE, HANDSHAKE_MSG1_SIZE, HANDSHAKE_MSG2_SIZE, HandshakeProgress, HandshakeRole, NoiseError, NoisePattern, NoiseSession, PROTOCOL_NAME_IK,
PROTOCOL_NAME_IK, PROTOCOL_NAME_XK, PUBKEY_SIZE, PROTOCOL_NAME_XK, PUBKEY_SIZE, XK_HANDSHAKE_MSG1_SIZE, XK_HANDSHAKE_MSG2_SIZE,
XK_HANDSHAKE_MSG1_SIZE, XK_HANDSHAKE_MSG2_SIZE, XK_HANDSHAKE_MSG3_SIZE, XK_HANDSHAKE_MSG3_SIZE,
}; };
use hkdf::Hkdf; use hkdf::Hkdf;
use rand::Rng; use rand::Rng;
use secp256k1::{ecdh::shared_secret_point, Keypair, PublicKey, Secp256k1, SecretKey}; use secp256k1::{Keypair, PublicKey, Secp256k1, SecretKey, ecdh::shared_secret_point};
use sha2::{Digest, Sha256}; use sha2::{Digest, Sha256};
use std::fmt; use std::fmt;
@@ -343,8 +343,12 @@ impl HandshakeState {
}); });
} }
let remote_static = self.remote_static.expect("initiator must have remote static"); let remote_static = self
let epoch = self.local_epoch.expect("local epoch must be set before write_message_1"); .remote_static
.expect("initiator must have remote static");
let epoch = self
.local_epoch
.expect("local epoch must be set before write_message_1");
// Generate ephemeral keypair // Generate ephemeral keypair
self.generate_ephemeral(); self.generate_ephemeral();
@@ -462,7 +466,9 @@ impl HandshakeState {
} }
let re = self.remote_ephemeral.expect("should have remote ephemeral"); let re = self.remote_ephemeral.expect("should have remote ephemeral");
let epoch = self.local_epoch.expect("local epoch must be set before write_message_2"); let epoch = self
.local_epoch
.expect("local epoch must be set before write_message_2");
// Generate ephemeral keypair // Generate ephemeral keypair
self.generate_ephemeral(); self.generate_ephemeral();
@@ -572,7 +578,9 @@ impl HandshakeState {
}); });
} }
let remote_static = self.remote_static.expect("initiator must have remote static"); let remote_static = self
.remote_static
.expect("initiator must have remote static");
// Generate ephemeral keypair // Generate ephemeral keypair
self.generate_ephemeral(); self.generate_ephemeral();
@@ -657,7 +665,9 @@ impl HandshakeState {
} }
let re = self.remote_ephemeral.expect("should have remote ephemeral"); let re = self.remote_ephemeral.expect("should have remote ephemeral");
let epoch = self.local_epoch.expect("local epoch must be set before write_xk_message_2"); let epoch = self
.local_epoch
.expect("local epoch must be set before write_xk_message_2");
// Generate ephemeral keypair // Generate ephemeral keypair
self.generate_ephemeral(); self.generate_ephemeral();
@@ -755,8 +765,12 @@ impl HandshakeState {
}); });
} }
let re = self.remote_ephemeral.expect("should have remote ephemeral after msg2"); let re = self
let epoch = self.local_epoch.expect("local epoch must be set before write_xk_message_3"); .remote_ephemeral
.expect("should have remote ephemeral after msg2");
let epoch = self
.local_epoch
.expect("local epoch must be set before write_xk_message_3");
let mut message = Vec::with_capacity(XK_HANDSHAKE_MSG3_SIZE); let mut message = Vec::with_capacity(XK_HANDSHAKE_MSG3_SIZE);
@@ -813,7 +827,10 @@ impl HandshakeState {
// -> se: DH(e, rs), mix into key // -> se: DH(e, rs), mix into key
// (responder uses their ephemeral with initiator's now-known static) // (responder uses their ephemeral with initiator's now-known static)
let ephemeral = self.ephemeral_keypair.as_ref().expect("should have ephemeral after msg2"); let ephemeral = self
.ephemeral_keypair
.as_ref()
.expect("should have ephemeral after msg2");
let se = self.ecdh(&ephemeral.secret_key(), &rs); let se = self.ecdh(&ephemeral.secret_key(), &rs);
self.symmetric.mix_key(&se); self.symmetric.mix_key(&se);

View File

@@ -40,8 +40,8 @@ mod replay;
mod session; mod session;
use chacha20poly1305::{ use chacha20poly1305::{
aead::{Aead, KeyInit, Payload},
ChaCha20Poly1305, Nonce, ChaCha20Poly1305, Nonce,
aead::{Aead, KeyInit, Payload},
}; };
use std::fmt; use std::fmt;
use thiserror::Error; use thiserror::Error;
@@ -327,7 +327,13 @@ impl CipherState {
let nonce = self.next_nonce()?; let nonce = self.next_nonce()?;
let ciphertext = cipher let ciphertext = cipher
.encrypt(&nonce, Payload { msg: plaintext, aad }) .encrypt(
&nonce,
Payload {
msg: plaintext,
aad,
},
)
.map_err(|_| NoiseError::EncryptionFailed)?; .map_err(|_| NoiseError::EncryptionFailed)?;
Ok(ciphertext) Ok(ciphertext)
@@ -360,7 +366,13 @@ impl CipherState {
let nonce = Self::counter_to_nonce(counter); let nonce = Self::counter_to_nonce(counter);
let plaintext = cipher let plaintext = cipher
.decrypt(&nonce, Payload { msg: ciphertext, aad }) .decrypt(
&nonce,
Payload {
msg: ciphertext,
aad,
},
)
.map_err(|_| NoiseError::DecryptionFailed)?; .map_err(|_| NoiseError::DecryptionFailed)?;
Ok(plaintext) Ok(plaintext)

View File

@@ -133,7 +133,9 @@ impl NoiseSession {
} }
// Attempt decryption with AAD (expensive) // Attempt decryption with AAD (expensive)
let plaintext = self.recv_cipher.decrypt_with_counter_and_aad(ciphertext, counter, aad)?; let plaintext = self
.recv_cipher
.decrypt_with_counter_and_aad(ciphertext, counter, aad)?;
// Only accept into window after successful decryption // Only accept into window after successful decryption
self.replay_window.accept(counter); self.replay_window.accept(counter);

View File

@@ -129,9 +129,11 @@ fn test_wrong_role_errors() {
initiator.set_local_epoch(generate_epoch()); initiator.set_local_epoch(generate_epoch());
// Initiator can't read message 1 // Initiator can't read message 1
assert!(initiator assert!(
.read_message_1(&[0u8; HANDSHAKE_MSG1_SIZE]) initiator
.is_err()); .read_message_1(&[0u8; HANDSHAKE_MSG1_SIZE])
.is_err()
);
// Initiator can't write message 2 before message 1 // Initiator can't write message 2 before message 1
assert!(initiator.write_message_2().is_err()); assert!(initiator.write_message_2().is_err());
@@ -337,7 +339,11 @@ fn test_replay_window_sequential() {
// All should be marked as seen // All should be marked as seen
for i in 0..1000 { for i in 0..1000 {
assert!(!window.check(i), "Counter {} should be rejected as replay", i); assert!(
!window.check(i),
"Counter {} should be rejected as replay",
i
);
} }
assert_eq!(window.highest(), 999); assert_eq!(window.highest(), 999);
@@ -403,18 +409,20 @@ fn test_handshake_with_odd_parity_responder() {
// Node B (responder) - odd parity key // Node B (responder) - odd parity key
let sk_b = secp256k1::SecretKey::from_slice( let sk_b = secp256k1::SecretKey::from_slice(
&hex::decode("b102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1fb0") &hex::decode("b102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1fb0").unwrap(),
.unwrap(),
) )
.unwrap(); .unwrap();
let kp_b = secp256k1::Keypair::from_secret_key(&secp, &sk_b); let kp_b = secp256k1::Keypair::from_secret_key(&secp, &sk_b);
let (xonly_b, parity_b) = kp_b.public_key().x_only_public_key(); let (xonly_b, parity_b) = kp_b.public_key().x_only_public_key();
assert_eq!(parity_b, Parity::Odd, "Test requires odd-parity responder key"); assert_eq!(
parity_b,
Parity::Odd,
"Test requires odd-parity responder key"
);
// Node A (initiator) - even parity key // Node A (initiator) - even parity key
let sk_a = secp256k1::SecretKey::from_slice( let sk_a = secp256k1::SecretKey::from_slice(
&hex::decode("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20") &hex::decode("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20").unwrap(),
.unwrap(),
) )
.unwrap(); .unwrap();
let kp_a = secp256k1::Keypair::from_secret_key(&secp, &sk_a); let kp_a = secp256k1::Keypair::from_secret_key(&secp, &sk_a);
@@ -423,7 +431,8 @@ fn test_handshake_with_odd_parity_responder() {
// (x-only -> assumed even parity) // (x-only -> assumed even parity)
let assumed_even_b = xonly_b.public_key(Parity::Even); let assumed_even_b = xonly_b.public_key(Parity::Even);
assert_ne!( assert_ne!(
assumed_even_b, kp_b.public_key(), assumed_even_b,
kp_b.public_key(),
"Even assumption should differ from actual odd key" "Even assumption should differ from actual odd key"
); );
@@ -554,7 +563,8 @@ fn test_xk_identity_timing() {
let initiator_keypair = generate_keypair(); let initiator_keypair = generate_keypair();
let responder_keypair = generate_keypair(); let responder_keypair = generate_keypair();
let mut initiator = HandshakeState::new_xk_initiator(initiator_keypair, responder_keypair.public_key()); let mut initiator =
HandshakeState::new_xk_initiator(initiator_keypair, responder_keypair.public_key());
initiator.set_local_epoch(generate_epoch()); initiator.set_local_epoch(generate_epoch());
let mut responder = HandshakeState::new_xk_responder(responder_keypair); let mut responder = HandshakeState::new_xk_responder(responder_keypair);
responder.set_local_epoch(generate_epoch()); responder.set_local_epoch(generate_epoch());
@@ -565,18 +575,30 @@ fn test_xk_identity_timing() {
// After msg1 // After msg1
let msg1 = initiator.write_xk_message_1().unwrap(); let msg1 = initiator.write_xk_message_1().unwrap();
responder.read_xk_message_1(&msg1).unwrap(); responder.read_xk_message_1(&msg1).unwrap();
assert!(responder.remote_static().is_none(), "XK: responder should NOT know identity after msg1"); assert!(
responder.remote_static().is_none(),
"XK: responder should NOT know identity after msg1"
);
// After msg2 // After msg2
let msg2 = responder.write_xk_message_2().unwrap(); let msg2 = responder.write_xk_message_2().unwrap();
initiator.read_xk_message_2(&msg2).unwrap(); initiator.read_xk_message_2(&msg2).unwrap();
assert!(responder.remote_static().is_none(), "XK: responder should NOT know identity after msg2"); assert!(
responder.remote_static().is_none(),
"XK: responder should NOT know identity after msg2"
);
// After msg3 // After msg3
let msg3 = initiator.write_xk_message_3().unwrap(); let msg3 = initiator.write_xk_message_3().unwrap();
responder.read_xk_message_3(&msg3).unwrap(); responder.read_xk_message_3(&msg3).unwrap();
assert!(responder.remote_static().is_some(), "XK: responder should know identity after msg3"); assert!(
assert_eq!(responder.remote_static().unwrap(), &initiator_keypair.public_key()); responder.remote_static().is_some(),
"XK: responder should know identity after msg3"
);
assert_eq!(
responder.remote_static().unwrap(),
&initiator_keypair.public_key()
);
} }
#[test] #[test]
@@ -587,7 +609,11 @@ fn test_xk_wrong_state_errors() {
// Initiator can't read XK msg1 // Initiator can't read XK msg1
let mut initiator = HandshakeState::new_xk_initiator(keypair1, keypair2.public_key()); let mut initiator = HandshakeState::new_xk_initiator(keypair1, keypair2.public_key());
initiator.set_local_epoch(generate_epoch()); initiator.set_local_epoch(generate_epoch());
assert!(initiator.read_xk_message_1(&[0u8; XK_HANDSHAKE_MSG1_SIZE]).is_err()); assert!(
initiator
.read_xk_message_1(&[0u8; XK_HANDSHAKE_MSG1_SIZE])
.is_err()
);
// Initiator can't write msg2 // Initiator can't write msg2
assert!(initiator.write_xk_message_2().is_err()); assert!(initiator.write_xk_message_2().is_err());
@@ -601,7 +627,11 @@ fn test_xk_wrong_state_errors() {
assert!(responder.write_xk_message_1().is_err()); assert!(responder.write_xk_message_1().is_err());
// Responder can't read msg3 before msg2 // Responder can't read msg3 before msg2
assert!(responder.read_xk_message_3(&[0u8; XK_HANDSHAKE_MSG3_SIZE]).is_err()); assert!(
responder
.read_xk_message_3(&[0u8; XK_HANDSHAKE_MSG3_SIZE])
.is_err()
);
} }
#[test] #[test]
@@ -636,7 +666,10 @@ fn test_xk_handshake_hash_differs_from_ik() {
xk_resp.read_xk_message_3(&msg3).unwrap(); xk_resp.read_xk_message_3(&msg3).unwrap();
let xk_hash = xk_init.handshake_hash(); let xk_hash = xk_init.handshake_hash();
assert_ne!(ik_hash, xk_hash, "IK and XK should produce different handshake hashes"); assert_ne!(
ik_hash, xk_hash,
"IK and XK should produce different handshake hashes"
);
} }
#[test] #[test]
@@ -677,18 +710,20 @@ fn test_xk_with_odd_parity_responder() {
// Node B (responder) - odd parity key // Node B (responder) - odd parity key
let sk_b = secp256k1::SecretKey::from_slice( let sk_b = secp256k1::SecretKey::from_slice(
&hex::decode("b102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1fb0") &hex::decode("b102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1fb0").unwrap(),
.unwrap(),
) )
.unwrap(); .unwrap();
let kp_b = secp256k1::Keypair::from_secret_key(&secp, &sk_b); let kp_b = secp256k1::Keypair::from_secret_key(&secp, &sk_b);
let (xonly_b, parity_b) = kp_b.public_key().x_only_public_key(); let (xonly_b, parity_b) = kp_b.public_key().x_only_public_key();
assert_eq!(parity_b, Parity::Odd, "Test requires odd-parity responder key"); assert_eq!(
parity_b,
Parity::Odd,
"Test requires odd-parity responder key"
);
// Node A (initiator) // Node A (initiator)
let sk_a = secp256k1::SecretKey::from_slice( let sk_a = secp256k1::SecretKey::from_slice(
&hex::decode("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20") &hex::decode("0102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f20").unwrap(),
.unwrap(),
) )
.unwrap(); .unwrap();
let kp_a = secp256k1::Keypair::from_secret_key(&secp, &sk_a); let kp_a = secp256k1::Keypair::from_secret_key(&secp, &sk_a);
@@ -716,7 +751,9 @@ fn test_xk_with_odd_parity_responder() {
let counter = sender.current_send_counter(); let counter = sender.current_send_counter();
let ciphertext = sender.encrypt(b"xk parity test").unwrap(); let ciphertext = sender.encrypt(b"xk parity test").unwrap();
let plaintext = receiver.decrypt_with_replay_check(&ciphertext, counter).unwrap(); let plaintext = receiver
.decrypt_with_replay_check(&ciphertext, counter)
.unwrap();
assert_eq!(plaintext, b"xk parity test"); assert_eq!(plaintext, b"xk parity test");
} }
@@ -727,7 +764,11 @@ fn test_xk_invalid_msg1_size() {
responder.set_local_epoch(generate_epoch()); responder.set_local_epoch(generate_epoch());
// Wrong size (IK msg1 size instead of XK) // Wrong size (IK msg1 size instead of XK)
assert!(responder.read_xk_message_1(&[0u8; HANDSHAKE_MSG1_SIZE]).is_err()); assert!(
responder
.read_xk_message_1(&[0u8; HANDSHAKE_MSG1_SIZE])
.is_err()
);
// Too short // Too short
assert!(responder.read_xk_message_1(&[0u8; 10]).is_err()); assert!(responder.read_xk_message_1(&[0u8; 10]).is_err());
} }
@@ -748,5 +789,9 @@ fn test_xk_invalid_msg3_size() {
// Responder is now in Message2Done, try wrong-size msg3 // Responder is now in Message2Done, try wrong-size msg3
assert!(responder.read_xk_message_3(&[0u8; 10]).is_err()); assert!(responder.read_xk_message_3(&[0u8; 10]).is_err());
assert!(responder.read_xk_message_3(&[0u8; XK_HANDSHAKE_MSG3_SIZE + 1]).is_err()); assert!(
responder
.read_xk_message_3(&[0u8; XK_HANDSHAKE_MSG3_SIZE + 1])
.is_err()
);
} }

View File

@@ -5,10 +5,10 @@
use crate::bloom::BloomFilter; use crate::bloom::BloomFilter;
use crate::mmp::{MmpConfig, MmpPeerState}; use crate::mmp::{MmpConfig, MmpPeerState};
use crate::utils::index::SessionIndex;
use crate::noise::{HandshakeState as NoiseHandshakeState, NoiseError, NoiseSession}; use crate::noise::{HandshakeState as NoiseHandshakeState, NoiseError, NoiseSession};
use crate::transport::{LinkId, LinkStats, TransportAddr, TransportId}; use crate::transport::{LinkId, LinkStats, TransportAddr, TransportId};
use crate::tree::{ParentDeclaration, TreeCoordinate}; use crate::tree::{ParentDeclaration, TreeCoordinate};
use crate::utils::index::SessionIndex;
use crate::{FipsAddress, NodeAddr, PeerIdentity}; use crate::{FipsAddress, NodeAddr, PeerIdentity};
use secp256k1::XOnlyPublicKey; use secp256k1::XOnlyPublicKey;
use std::fmt; use std::fmt;
@@ -32,7 +32,10 @@ pub enum ConnectivityState {
impl ConnectivityState { impl ConnectivityState {
/// Check if the peer is usable for sending traffic. /// Check if the peer is usable for sending traffic.
pub fn can_send(&self) -> bool { pub fn can_send(&self) -> bool {
matches!(self, ConnectivityState::Connected | ConnectivityState::Stale) matches!(
self,
ConnectivityState::Connected | ConnectivityState::Stale
)
} }
/// Check if this is a terminal state requiring cleanup. /// Check if this is a terminal state requiring cleanup.
@@ -744,12 +747,7 @@ impl ActivePeer {
// === Filter Updates === // === Filter Updates ===
/// Update peer's inbound filter. /// Update peer's inbound filter.
pub fn update_filter( pub fn update_filter(&mut self, filter: BloomFilter, sequence: u64, current_time_ms: u64) {
&mut self,
filter: BloomFilter,
sequence: u64,
current_time_ms: u64,
) {
self.inbound_filter = Some(filter); self.inbound_filter = Some(filter);
self.filter_sequence = sequence; self.filter_sequence = sequence;
self.filter_received_at = current_time_ms; self.filter_received_at = current_time_ms;
@@ -964,12 +962,11 @@ impl ActivePeer {
self.rekey_msg1_next_resend = 0; self.rekey_msg1_next_resend = 0;
self.rekey_in_progress = false; self.rekey_in_progress = false;
// Return whichever index needs freeing // Return whichever index needs freeing
self.rekey_our_index.take() self.rekey_our_index.take().or_else(|| {
.or_else(|| { self.pending_new_session = None;
self.pending_new_session = None; self.pending_their_index = None;
self.pending_their_index = None; self.pending_our_index.take()
self.pending_our_index.take() })
})
} }
// === Rekey Handshake State (Initiator) === // === Rekey Handshake State (Initiator) ===
@@ -999,11 +996,9 @@ impl ActivePeer {
/// Takes the stored handshake state, reads msg2, and returns the /// Takes the stored handshake state, reads msg2, and returns the
/// completed NoiseSession. Clears the handshake-related fields but /// completed NoiseSession. Clears the handshake-related fields but
/// leaves rekey_our_index for set_pending_session to use. /// leaves rekey_our_index for set_pending_session to use.
pub fn complete_rekey_msg2( pub fn complete_rekey_msg2(&mut self, msg2_bytes: &[u8]) -> Result<NoiseSession, NoiseError> {
&mut self, let mut hs = self
msg2_bytes: &[u8], .rekey_handshake
) -> Result<NoiseSession, NoiseError> {
let mut hs = self.rekey_handshake
.take() .take()
.ok_or_else(|| NoiseError::WrongState { .ok_or_else(|| NoiseError::WrongState {
expected: "rekey handshake in progress".to_string(), expected: "rekey handshake in progress".to_string(),
@@ -1022,9 +1017,7 @@ impl ActivePeer {
/// Check if msg1 needs resending. /// Check if msg1 needs resending.
pub fn needs_msg1_resend(&self, now_ms: u64) -> bool { pub fn needs_msg1_resend(&self, now_ms: u64) -> bool {
self.rekey_in_progress self.rekey_in_progress && self.rekey_msg1.is_some() && now_ms >= self.rekey_msg1_next_resend
&& self.rekey_msg1.is_some()
&& now_ms >= self.rekey_msg1_next_resend
} }
/// Get msg1 bytes for resend (without consuming). /// Get msg1 bytes for resend (without consuming).

View File

@@ -4,10 +4,10 @@
//! PeerConnection tracks the Noise IK handshake state and transitions to //! PeerConnection tracks the Noise IK handshake state and transitions to
//! ActivePeer upon successful authentication. //! ActivePeer upon successful authentication.
use crate::utils::index::SessionIndex; use crate::PeerIdentity;
use crate::noise::{self, NoiseError, NoiseSession}; use crate::noise::{self, NoiseError, NoiseSession};
use crate::transport::{LinkDirection, LinkId, LinkStats, TransportAddr, TransportId}; use crate::transport::{LinkDirection, LinkId, LinkStats, TransportAddr, TransportId};
use crate::PeerIdentity; use crate::utils::index::SessionIndex;
use secp256k1::Keypair; use secp256k1::Keypair;
use std::fmt; use std::fmt;
@@ -558,7 +558,6 @@ impl PeerConnection {
pub fn is_timed_out(&self, current_time_ms: u64, timeout_ms: u64) -> bool { pub fn is_timed_out(&self, current_time_ms: u64, timeout_ms: u64) -> bool {
self.idle_time(current_time_ms) > timeout_ms self.idle_time(current_time_ms) > timeout_ms
} }
} }
impl fmt::Debug for PeerConnection { impl fmt::Debug for PeerConnection {
@@ -651,12 +650,13 @@ mod tests {
let responder_peer_id = PeerIdentity::from_pubkey_full(responder_identity.pubkey_full()); let responder_peer_id = PeerIdentity::from_pubkey_full(responder_identity.pubkey_full());
// Create connections // Create connections
let mut initiator_conn = let mut initiator_conn = PeerConnection::outbound(LinkId::new(1), responder_peer_id, 1000);
PeerConnection::outbound(LinkId::new(1), responder_peer_id, 1000);
let mut responder_conn = PeerConnection::inbound(LinkId::new(2), 1000); let mut responder_conn = PeerConnection::inbound(LinkId::new(2), 1000);
// Initiator starts handshake // Initiator starts handshake
let msg1 = initiator_conn.start_handshake(initiator_keypair, initiator_epoch, 1100).unwrap(); let msg1 = initiator_conn
.start_handshake(initiator_keypair, initiator_epoch, 1100)
.unwrap();
assert_eq!(initiator_conn.handshake_state(), HandshakeState::SentMsg1); assert_eq!(initiator_conn.handshake_state(), HandshakeState::SentMsg1);
// Responder processes msg1 and sends msg2 // Responder processes msg1 and sends msg2
@@ -723,12 +723,18 @@ mod tests {
// Outbound can't receive_handshake_init // Outbound can't receive_handshake_init
let mut outbound = PeerConnection::outbound(LinkId::new(1), identity, 1000); let mut outbound = PeerConnection::outbound(LinkId::new(1), identity, 1000);
assert!(outbound assert!(
.receive_handshake_init(keypair, make_epoch(), &[0u8; 106], 1100) outbound
.is_err()); .receive_handshake_init(keypair, make_epoch(), &[0u8; 106], 1100)
.is_err()
);
// Inbound can't start_handshake // Inbound can't start_handshake
let mut inbound = PeerConnection::inbound(LinkId::new(2), 1000); let mut inbound = PeerConnection::inbound(LinkId::new(2), 1000);
assert!(inbound.start_handshake(keypair, make_epoch(), 1100).is_err()); assert!(
inbound
.start_handshake(keypair, make_epoch(), 1100)
.is_err()
);
} }
} }

View File

@@ -13,8 +13,8 @@ mod connection;
pub use active::{ActivePeer, ConnectivityState}; pub use active::{ActivePeer, ConnectivityState};
pub use connection::{HandshakeState, PeerConnection}; pub use connection::{HandshakeState, PeerConnection};
use crate::transport::LinkId;
use crate::NodeAddr; use crate::NodeAddr;
use crate::transport::LinkId;
use std::fmt; use std::fmt;
use thiserror::Error; use thiserror::Error;
@@ -44,7 +44,10 @@ pub enum PeerError {
HandshakeTimeout, HandshakeTimeout,
#[error("identity mismatch: expected {expected:?}, got {actual:?}")] #[error("identity mismatch: expected {expected:?}, got {actual:?}")]
IdentityMismatch { expected: NodeAddr, actual: NodeAddr }, IdentityMismatch {
expected: NodeAddr,
actual: NodeAddr,
},
#[error("peer disconnected")] #[error("peer disconnected")]
Disconnected, Disconnected,
@@ -240,10 +243,20 @@ impl fmt::Display for PeerSlot {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self { match self {
PeerSlot::Connecting(conn) => { PeerSlot::Connecting(conn) => {
write!(f, "connecting(link={}, state={})", conn.link_id(), conn.handshake_state()) write!(
f,
"connecting(link={}, state={})",
conn.link_id(),
conn.handshake_state()
)
} }
PeerSlot::Active(peer) => { PeerSlot::Active(peer) => {
write!(f, "active(node={:?}, link={})", peer.node_addr(), peer.link_id()) write!(
f,
"active(node={:?}, link={})",
peer.node_addr(),
peer.link_id()
)
} }
} }
} }

View File

@@ -1,9 +1,9 @@
//! Discovery messages: LookupRequest and LookupResponse. //! Discovery messages: LookupRequest and LookupResponse.
use crate::NodeAddr;
use crate::protocol::error::ProtocolError; use crate::protocol::error::ProtocolError;
use crate::protocol::session::{decode_coords, encode_coords}; use crate::protocol::session::{decode_coords, encode_coords};
use crate::tree::TreeCoordinate; use crate::tree::TreeCoordinate;
use crate::NodeAddr;
use secp256k1::schnorr::Signature; use secp256k1::schnorr::Signature;
/// Request to discover a node's coordinates. /// Request to discover a node's coordinates.
@@ -192,7 +192,11 @@ impl LookupResponse {
/// Get the bytes that should be signed as proof. /// Get the bytes that should be signed as proof.
/// ///
/// Format: request_id (8) || target (16) || coords_encoding (2 + 16×n) /// Format: request_id (8) || target (16) || coords_encoding (2 + 16×n)
pub fn proof_bytes(request_id: u64, target: &NodeAddr, target_coords: &TreeCoordinate) -> Vec<u8> { pub fn proof_bytes(
request_id: u64,
target: &NodeAddr,
target_coords: &TreeCoordinate,
) -> Vec<u8> {
let coord_size = 2 + target_coords.entries().len() * 16; let coord_size = 2 + target_coords.entries().len() * 16;
let mut bytes = Vec::with_capacity(24 + coord_size); let mut bytes = Vec::with_capacity(24 + coord_size);
bytes.extend_from_slice(&request_id.to_le_bytes()); bytes.extend_from_slice(&request_id.to_le_bytes());

View File

@@ -42,11 +42,7 @@ impl FilterAnnounce {
} }
/// Create with explicit size_class (for testing or future protocol versions). /// Create with explicit size_class (for testing or future protocol versions).
pub fn with_size_class( pub fn with_size_class(filter: BloomFilter, sequence: u64, size_class: u8) -> Self {
filter: BloomFilter,
sequence: u64,
size_class: u8,
) -> Self {
Self { Self {
hash_count: filter.hash_count(), hash_count: filter.hash_count(),
size_class, size_class,
@@ -164,10 +160,8 @@ impl FilterAnnounce {
} }
// Construct BloomFilter from bytes // Construct BloomFilter from bytes
let filter = let filter = crate::bloom::BloomFilter::from_slice(&payload[pos..], hash_count)
crate::bloom::BloomFilter::from_slice(&payload[pos..], hash_count).map_err(|e| { .map_err(|e| ProtocolError::Malformed(format!("invalid bloom filter: {e}")))?;
ProtocolError::Malformed(format!("invalid bloom filter: {e}"))
})?;
let announce = Self { let announce = Self {
filter, filter,
@@ -253,7 +247,12 @@ mod tests {
let result = FilterAnnounce::decode(&encoded[1..]); let result = FilterAnnounce::decode(&encoded[1..]);
assert!(result.is_err()); assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("invalid size_class")); assert!(
result
.unwrap_err()
.to_string()
.contains("invalid size_class")
);
} }
#[test] #[test]
@@ -265,10 +264,12 @@ mod tests {
let result = FilterAnnounce::decode(&encoded[1..]); let result = FilterAnnounce::decode(&encoded[1..]);
assert!(result.is_err()); assert!(result.is_err());
assert!(result assert!(
.unwrap_err() result
.to_string() .unwrap_err()
.contains("unsupported size_class")); .to_string()
.contains("unsupported size_class")
);
} }
#[test] #[test]

View File

@@ -533,8 +533,7 @@ mod tests {
let src = make_node_addr(0xAA); let src = make_node_addr(0xAA);
let dest = make_node_addr(0xBB); let dest = make_node_addr(0xBB);
let payload = vec![0x10, 0x00, 0x05, 0x00, 1, 2, 3, 4, 5]; // session payload let payload = vec![0x10, 0x00, 0x05, 0x00, 1, 2, 3, 4, 5]; // session payload
let dg = SessionDatagram::new(src, dest, payload.clone()) let dg = SessionDatagram::new(src, dest, payload.clone()).with_ttl(32);
.with_ttl(32);
let encoded = dg.encode(); let encoded = dg.encode();
assert_eq!(encoded[0], 0x00); // msg_type (SessionDatagram) assert_eq!(encoded[0], 0x00); // msg_type (SessionDatagram)

View File

@@ -28,21 +28,21 @@ mod session;
mod tree; mod tree;
// Re-export all public types at protocol:: level // Re-export all public types at protocol:: level
pub use error::ProtocolError;
pub use link::{
Disconnect, DisconnectReason, HandshakeMessageType, LinkMessageType, SessionDatagram,
SESSION_DATAGRAM_HEADER_SIZE,
};
pub use tree::TreeAnnounce;
pub use filter::FilterAnnounce;
pub use discovery::{LookupRequest, LookupResponse}; pub use discovery::{LookupRequest, LookupResponse};
pub use error::ProtocolError;
pub use filter::FilterAnnounce;
pub use link::{
Disconnect, DisconnectReason, HandshakeMessageType, LinkMessageType,
SESSION_DATAGRAM_HEADER_SIZE, SessionDatagram,
};
pub use session::{ pub use session::{
CoordsRequired, FspFlags, FspInnerFlags, MtuExceeded, PathBroken, PathMtuNotification, COORDS_REQUIRED_SIZE, CoordsRequired, FspFlags, FspInnerFlags, MTU_EXCEEDED_SIZE, MtuExceeded,
SessionAck, SessionFlags, SessionMessageType, SessionMsg3, SessionReceiverReport, PATH_MTU_NOTIFICATION_SIZE, PathBroken, PathMtuNotification, SESSION_RECEIVER_REPORT_SIZE,
SessionSenderReport, SessionSetup, COORDS_REQUIRED_SIZE, MTU_EXCEEDED_SIZE, SESSION_SENDER_REPORT_SIZE, SessionAck, SessionFlags, SessionMessageType, SessionMsg3,
PATH_MTU_NOTIFICATION_SIZE, SESSION_RECEIVER_REPORT_SIZE, SESSION_SENDER_REPORT_SIZE, SessionReceiverReport, SessionSenderReport, SessionSetup,
}; };
pub(crate) use session::{coords_wire_size, decode_optional_coords, encode_coords}; pub(crate) use session::{coords_wire_size, decode_optional_coords, encode_coords};
pub use tree::TreeAnnounce;
/// Protocol version for message compatibility. /// Protocol version for message compatibility.
pub const PROTOCOL_VERSION: u8 = 1; pub const PROTOCOL_VERSION: u8 = 1;

View File

@@ -1,8 +1,8 @@
//! Session-layer message types: setup, ack, data, and error messages. //! Session-layer message types: setup, ack, data, and error messages.
use super::ProtocolError; use super::ProtocolError;
use crate::tree::TreeCoordinate;
use crate::NodeAddr; use crate::NodeAddr;
use crate::tree::TreeCoordinate;
use std::fmt; use std::fmt;
// ============================================================================ // ============================================================================
@@ -141,15 +141,17 @@ pub(crate) fn decode_coords(data: &[u8]) -> Result<(TreeCoordinate, usize), Prot
bytes.copy_from_slice(&data[offset..offset + 16]); bytes.copy_from_slice(&data[offset..offset + 16]);
addrs.push(NodeAddr::from_bytes(bytes)); addrs.push(NodeAddr::from_bytes(bytes));
} }
let coord = TreeCoordinate::from_addrs(addrs) let coord =
.map_err(|e| ProtocolError::Malformed(e.to_string()))?; TreeCoordinate::from_addrs(addrs).map_err(|e| ProtocolError::Malformed(e.to_string()))?;
Ok((coord, needed)) Ok((coord, needed))
} }
/// Decode an optional coordinate field (count may be 0). /// Decode an optional coordinate field (count may be 0).
/// ///
/// Returns None if count is 0, Some(coord) otherwise, plus bytes consumed. /// Returns None if count is 0, Some(coord) otherwise, plus bytes consumed.
pub(crate) fn decode_optional_coords(data: &[u8]) -> Result<(Option<TreeCoordinate>, usize), ProtocolError> { pub(crate) fn decode_optional_coords(
data: &[u8],
) -> Result<(Option<TreeCoordinate>, usize), ProtocolError> {
if data.len() < 2 { if data.len() < 2 {
return Err(ProtocolError::MessageTooShort { return Err(ProtocolError::MessageTooShort {
expected: 2, expected: 2,
@@ -174,8 +176,8 @@ pub(crate) fn decode_optional_coords(data: &[u8]) -> Result<(Option<TreeCoordina
bytes.copy_from_slice(&data[offset..offset + 16]); bytes.copy_from_slice(&data[offset..offset + 16]);
addrs.push(NodeAddr::from_bytes(bytes)); addrs.push(NodeAddr::from_bytes(bytes));
} }
let coord = TreeCoordinate::from_addrs(addrs) let coord =
.map_err(|e| ProtocolError::Malformed(e.to_string()))?; TreeCoordinate::from_addrs(addrs).map_err(|e| ProtocolError::Malformed(e.to_string()))?;
Ok((Some(coord), needed)) Ok((Some(coord), needed))
} }
@@ -909,7 +911,10 @@ pub const COORDS_REQUIRED_SIZE: usize = 34;
impl CoordsRequired { impl CoordsRequired {
/// Create a new CoordsRequired error. /// Create a new CoordsRequired error.
pub fn new(dest_addr: NodeAddr, reporter: NodeAddr) -> Self { pub fn new(dest_addr: NodeAddr, reporter: NodeAddr) -> Self {
Self { dest_addr, reporter } Self {
dest_addr,
reporter,
}
} }
/// Encode as wire format (4-byte FSP prefix + msg_type + body). /// Encode as wire format (4-byte FSP prefix + msg_type + body).
@@ -1081,7 +1086,11 @@ pub const MTU_EXCEEDED_SIZE: usize = 36;
impl MtuExceeded { impl MtuExceeded {
/// Create a new MtuExceeded error. /// Create a new MtuExceeded error.
pub fn new(dest_addr: NodeAddr, reporter: NodeAddr, mtu: u16) -> Self { pub fn new(dest_addr: NodeAddr, reporter: NodeAddr, mtu: u16) -> Self {
Self { dest_addr, reporter, mtu } Self {
dest_addr,
reporter,
mtu,
}
} }
/// Encode as wire format (4-byte FSP prefix + msg_type + body). /// Encode as wire format (4-byte FSP prefix + msg_type + body).
@@ -1359,8 +1368,7 @@ mod tests {
let addrs: Vec<u8> = (0..11).collect(); let addrs: Vec<u8> = (0..11).collect();
let src = make_coords(&addrs); let src = make_coords(&addrs);
let dest = make_coords(&[20, 21, 22, 23, 24]); let dest = make_coords(&[20, 21, 22, 23, 24]);
let setup = SessionSetup::new(src.clone(), dest.clone()) let setup = SessionSetup::new(src.clone(), dest.clone()).with_handshake(vec![0x55; 82]);
.with_handshake(vec![0x55; 82]);
let encoded = setup.encode(); let encoded = setup.encode();
let decoded = SessionSetup::decode(&encoded[4..]).unwrap(); let decoded = SessionSetup::decode(&encoded[4..]).unwrap();
@@ -1459,9 +1467,18 @@ mod tests {
#[test] #[test]
fn test_session_message_type_display() { fn test_session_message_type_display() {
assert_eq!(format!("{}", SessionMessageType::SenderReport), "SenderReport"); assert_eq!(
assert_eq!(format!("{}", SessionMessageType::ReceiverReport), "ReceiverReport"); format!("{}", SessionMessageType::SenderReport),
assert_eq!(format!("{}", SessionMessageType::PathMtuNotification), "PathMtuNotification"); "SenderReport"
);
assert_eq!(
format!("{}", SessionMessageType::ReceiverReport),
"ReceiverReport"
);
assert_eq!(
format!("{}", SessionMessageType::PathMtuNotification),
"PathMtuNotification"
);
} }
// ===== SessionSenderReport Tests ===== // ===== SessionSenderReport Tests =====
@@ -1639,7 +1656,10 @@ mod tests {
#[test] #[test]
fn test_mtu_exceeded_display() { fn test_mtu_exceeded_display() {
assert_eq!(format!("{}", SessionMessageType::MtuExceeded), "MtuExceeded"); assert_eq!(
format!("{}", SessionMessageType::MtuExceeded),
"MtuExceeded"
);
} }
// ===== SessionMsg3 Tests ===== // ===== SessionMsg3 Tests =====

View File

@@ -2,8 +2,8 @@
use super::error::ProtocolError; use super::error::ProtocolError;
use super::link::LinkMessageType; use super::link::LinkMessageType;
use crate::tree::{CoordEntry, ParentDeclaration, TreeCoordinate};
use crate::NodeAddr; use crate::NodeAddr;
use crate::tree::{CoordEntry, ParentDeclaration, TreeCoordinate};
use secp256k1::schnorr::Signature; use secp256k1::schnorr::Signature;
/// Spanning tree announcement carrying parent declaration and ancestry. /// Spanning tree announcement carrying parent declaration and ancestry.
@@ -166,8 +166,8 @@ impl TreeAnnounce {
let sig_bytes: [u8; 64] = payload[pos..pos + 64] let sig_bytes: [u8; 64] = payload[pos..pos + 64]
.try_into() .try_into()
.map_err(|_| ProtocolError::Malformed("bad signature".into()))?; .map_err(|_| ProtocolError::Malformed("bad signature".into()))?;
let signature = Signature::from_slice(&sig_bytes) let signature =
.map_err(|_| ProtocolError::InvalidSignature)?; Signature::from_slice(&sig_bytes).map_err(|_| ProtocolError::InvalidSignature)?;
// The first entry's node_addr is the declaring node // The first entry's node_addr is the declaring node
if entries.is_empty() { if entries.is_empty() {
@@ -324,7 +324,10 @@ mod tests {
encoded[1] = 0xFF; encoded[1] = 0xFF;
let result = TreeAnnounce::decode(&encoded[1..]); let result = TreeAnnounce::decode(&encoded[1..]);
assert!(matches!(result, Err(ProtocolError::UnsupportedVersion(0xFF)))); assert!(matches!(
result,
Err(ProtocolError::UnsupportedVersion(0xFF))
));
} }
#[test] #[test]
@@ -365,10 +368,7 @@ mod tests {
encoded[35] = 0; encoded[35] = 0;
let result = TreeAnnounce::decode(&encoded[1..]); let result = TreeAnnounce::decode(&encoded[1..]);
assert!(matches!( assert!(matches!(result, Err(ProtocolError::MessageTooShort { .. })));
result,
Err(ProtocolError::MessageTooShort { .. })
));
} }
#[test] #[test]

View File

@@ -13,10 +13,8 @@ use super::{
TransportId, TransportState, TransportType, TransportId, TransportState, TransportType,
}; };
use crate::config::EthernetConfig; use crate::config::EthernetConfig;
use discovery::{ use discovery::{DiscoveryBuffer, FRAME_TYPE_BEACON, FRAME_TYPE_DATA, build_beacon, parse_beacon};
build_beacon, parse_beacon, DiscoveryBuffer, FRAME_TYPE_BEACON, FRAME_TYPE_DATA, use socket::{AsyncPacketSocket, ETHERNET_BROADCAST, PacketSocket};
};
use socket::{AsyncPacketSocket, PacketSocket, ETHERNET_BROADCAST};
use stats::EthernetStats; use stats::EthernetStats;
use secp256k1::XOnlyPublicKey; use secp256k1::XOnlyPublicKey;
@@ -399,8 +397,7 @@ async fn ethernet_receive_loop(
trace!("Data frame too short ({len} bytes), ignoring"); trace!("Data frame too short ({len} bytes), ignoring");
continue; continue;
} }
let payload_len = let payload_len = u16::from_le_bytes([buf[1], buf[2]]) as usize;
u16::from_le_bytes([buf[1], buf[2]]) as usize;
if payload_len > len - 3 { if payload_len > len - 3 {
trace!( trace!(
"Data frame length field ({payload_len}) exceeds \ "Data frame length field ({payload_len}) exceeds \
@@ -431,9 +428,7 @@ async fn ethernet_receive_loop(
FRAME_TYPE_BEACON => { FRAME_TYPE_BEACON => {
stats.record_beacon_recv(); stats.record_beacon_recv();
if discovery_enabled if discovery_enabled && let Some(pubkey) = parse_beacon(&buf[..len]) {
&& let Some(pubkey) = parse_beacon(&buf[..len])
{
discovery_buffer.add_peer(src_mac, pubkey); discovery_buffer.add_peer(src_mac, pubkey);
trace!( trace!(
transport_id = %transport_id, transport_id = %transport_id,

View File

@@ -254,10 +254,7 @@ impl AsyncPacketSocket {
} }
/// Receive a payload and source MAC address. /// Receive a payload and source MAC address.
pub async fn recv_from( pub async fn recv_from(&self, buf: &mut [u8]) -> Result<(usize, [u8; 6]), TransportError> {
&self,
buf: &mut [u8],
) -> Result<(usize, [u8; 6]), TransportError> {
loop { loop {
let mut guard = self let mut guard = self
.inner .inner
@@ -308,7 +305,10 @@ fn get_mac_addr(fd: RawFd, if_index: i32) -> Result<[u8; 6], TransportError> {
// Use if_indextoname to get the name // Use if_indextoname to get the name
let mut name_buf = [0u8; libc::IFNAMSIZ]; let mut name_buf = [0u8; libc::IFNAMSIZ];
let ret = unsafe { let ret = unsafe {
libc::if_indextoname(if_index as libc::c_uint, name_buf.as_mut_ptr() as *mut libc::c_char) libc::if_indextoname(
if_index as libc::c_uint,
name_buf.as_mut_ptr() as *mut libc::c_char,
)
}; };
if ret.is_null() { if ret.is_null() {
return Err(TransportError::StartFailed(format!( return Err(TransportError::StartFailed(format!(
@@ -319,7 +319,10 @@ fn get_mac_addr(fd: RawFd, if_index: i32) -> Result<[u8; 6], TransportError> {
} }
// Copy name into ifreq // Copy name into ifreq
let name_len = name_buf.iter().position(|&b| b == 0).unwrap_or(name_buf.len()); let name_len = name_buf
.iter()
.position(|&b| b == 0)
.unwrap_or(name_buf.len());
let copy_len = name_len.min(libc::IFNAMSIZ - 1); let copy_len = name_len.min(libc::IFNAMSIZ - 1);
unsafe { unsafe {
std::ptr::copy_nonoverlapping( std::ptr::copy_nonoverlapping(
@@ -359,7 +362,10 @@ fn get_if_mtu(fd: RawFd, if_index: i32) -> Result<u16, TransportError> {
// Get the interface name from index // Get the interface name from index
let mut name_buf = [0u8; libc::IFNAMSIZ]; let mut name_buf = [0u8; libc::IFNAMSIZ];
let ret = unsafe { let ret = unsafe {
libc::if_indextoname(if_index as libc::c_uint, name_buf.as_mut_ptr() as *mut libc::c_char) libc::if_indextoname(
if_index as libc::c_uint,
name_buf.as_mut_ptr() as *mut libc::c_char,
)
}; };
if ret.is_null() { if ret.is_null() {
return Err(TransportError::StartFailed(format!( return Err(TransportError::StartFailed(format!(
@@ -369,7 +375,10 @@ fn get_if_mtu(fd: RawFd, if_index: i32) -> Result<u16, TransportError> {
))); )));
} }
let name_len = name_buf.iter().position(|&b| b == 0).unwrap_or(name_buf.len()); let name_len = name_buf
.iter()
.position(|&b| b == 0)
.unwrap_or(name_buf.len());
let copy_len = name_len.min(libc::IFNAMSIZ - 1); let copy_len = name_len.min(libc::IFNAMSIZ - 1);
unsafe { unsafe {
std::ptr::copy_nonoverlapping( std::ptr::copy_nonoverlapping(

View File

@@ -4,24 +4,24 @@
//! underlying communication mechanisms (UDP, Ethernet, Tor, etc.) over //! underlying communication mechanisms (UDP, Ethernet, Tor, etc.) over
//! which FIPS links are established. //! which FIPS links are established.
pub mod udp;
pub mod tcp; pub mod tcp;
pub mod tor; pub mod tor;
pub mod udp;
#[cfg(target_os = "linux")] #[cfg(target_os = "linux")]
pub mod ethernet; pub mod ethernet;
use secp256k1::XOnlyPublicKey;
use udp::UdpTransport;
use tcp::TcpTransport;
use tor::control::TorMonitoringInfo;
use tor::TorTransport;
#[cfg(target_os = "linux")] #[cfg(target_os = "linux")]
use ethernet::EthernetTransport; use ethernet::EthernetTransport;
use secp256k1::XOnlyPublicKey;
use std::fmt; use std::fmt;
use std::net::SocketAddr; use std::net::SocketAddr;
use std::time::{Duration, SystemTime, UNIX_EPOCH}; use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tcp::TcpTransport;
use thiserror::Error; use thiserror::Error;
use tor::TorTransport;
use tor::control::TorMonitoringInfo;
use udp::UdpTransport;
// ============================================================================ // ============================================================================
// Packet Channel Types // Packet Channel Types
@@ -1567,10 +1567,8 @@ mod tests {
let addr_b = TransportAddr::from_string("10.0.0.1:5000"); let addr_b = TransportAddr::from_string("10.0.0.1:5000");
let addr_unknown = TransportAddr::from_string("172.16.0.1:6000"); let addr_unknown = TransportAddr::from_string("172.16.0.1:6000");
let transport = PerLinkMtuTransport::new( let transport =
1280, PerLinkMtuTransport::new(1280, vec![(addr_a.clone(), 512), (addr_b.clone(), 247)]);
vec![(addr_a.clone(), 512), (addr_b.clone(), 247)],
);
// Known addresses return their per-link MTU // Known addresses return their per-link MTU
assert_eq!(transport.link_mtu(&addr_a), 512); assert_eq!(transport.link_mtu(&addr_a), 512);

View File

@@ -41,8 +41,8 @@ use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use tokio::io::AsyncWriteExt; use tokio::io::AsyncWriteExt;
use tokio::net::{TcpListener, TcpStream};
use tokio::net::tcp::OwnedWriteHalf; use tokio::net::tcp::OwnedWriteHalf;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::Mutex; use tokio::sync::Mutex;
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
use tokio::time::Instant; use tokio::time::Instant;
@@ -362,7 +362,8 @@ impl TcpTransport {
}; };
// Configure socket options via socket2 // Configure socket options via socket2
let std_stream = stream.into_std() let std_stream = stream
.into_std()
.map_err(|e| TransportError::StartFailed(format!("into_std: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("into_std: {}", e)))?;
configure_socket(&std_stream, &self.config)?; configure_socket(&std_stream, &self.config)?;
@@ -385,7 +386,16 @@ impl TcpTransport {
let mtu = mss_mtu; let mtu = mss_mtu;
let recv_task = tokio::spawn(async move { let recv_task = tokio::spawn(async move {
tcp_receive_loop(read_half, transport_id, remote_addr.clone(), packet_tx, pool, mtu, recv_stats).await; tcp_receive_loop(
read_half,
transport_id,
remote_addr.clone(),
packet_tx,
pool,
mtu,
recv_stats,
)
.await;
}); });
let conn = TcpConnection { let conn = TcpConnection {
@@ -521,7 +531,8 @@ impl TcpTransport {
}; };
// Configure socket options via socket2 // Configure socket options via socket2
let std_stream = stream.into_std() let std_stream = stream
.into_std()
.map_err(|e| TransportError::StartFailed(format!("into_std: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("into_std: {}", e)))?;
configure_socket(&std_stream, &config)?; configure_socket(&std_stream, &config)?;
@@ -589,9 +600,7 @@ impl TcpTransport {
self.promote_connection(addr, stream, mss_mtu); self.promote_connection(addr, stream, mss_mtu);
ConnectionState::Connected ConnectionState::Connected
} }
Some(Ok(Err(e))) => { Some(Ok(Err(e))) => ConnectionState::Failed(format!("{}", e)),
ConnectionState::Failed(format!("{}", e))
}
Some(Err(e)) => { Some(Err(e)) => {
// JoinError (panic or cancel) // JoinError (panic or cancel)
ConnectionState::Failed(format!("task failed: {}", e)) ConnectionState::Failed(format!("task failed: {}", e))
@@ -737,7 +746,14 @@ async fn accept_loop(
cfg: AcceptConfig, cfg: AcceptConfig,
stats: Arc<TcpStats>, stats: Arc<TcpStats>,
) { ) {
let AcceptConfig { mtu, max_inbound, nodelay, keepalive_secs, recv_buf, send_buf } = cfg; let AcceptConfig {
mtu,
max_inbound,
nodelay,
keepalive_secs,
recv_buf,
send_buf,
} = cfg;
debug!(transport_id = %transport_id, "TCP accept loop starting"); debug!(transport_id = %transport_id, "TCP accept loop starting");
loop { loop {
@@ -771,7 +787,13 @@ async fn accept_loop(
} }
}; };
if let Err(e) = configure_accepted_socket(&std_stream, nodelay, keepalive_secs, recv_buf, send_buf) { if let Err(e) = configure_accepted_socket(
&std_stream,
nodelay,
keepalive_secs,
recv_buf,
send_buf,
) {
warn!( warn!(
transport_id = %transport_id, transport_id = %transport_id,
peer_addr = %peer_addr, peer_addr = %peer_addr,
@@ -886,11 +908,7 @@ async fn tcp_receive_loop(
"TCP packet received" "TCP packet received"
); );
let packet = ReceivedPacket::new( let packet = ReceivedPacket::new(transport_id, remote_addr.clone(), data);
transport_id,
remote_addr.clone(),
data,
);
if packet_tx.send(packet).await.is_err() { if packet_tx.send(packet).await.is_err() {
info!( info!(
@@ -934,26 +952,30 @@ fn configure_socket(
stream: &std::net::TcpStream, stream: &std::net::TcpStream,
config: &TcpConfig, config: &TcpConfig,
) -> Result<(), TransportError> { ) -> Result<(), TransportError> {
let socket = socket2::SockRef::from(stream).try_clone() let socket = socket2::SockRef::from(stream)
.try_clone()
.map_err(|e| TransportError::StartFailed(format!("clone socket: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("clone socket: {}", e)))?;
// TCP_NODELAY // TCP_NODELAY
socket.set_tcp_nodelay(config.nodelay()) socket
.set_tcp_nodelay(config.nodelay())
.map_err(|e| TransportError::StartFailed(format!("set nodelay: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("set nodelay: {}", e)))?;
// Keepalive // Keepalive
let keepalive_secs = config.keepalive_secs(); let keepalive_secs = config.keepalive_secs();
if keepalive_secs > 0 { if keepalive_secs > 0 {
let keepalive = TcpKeepalive::new() let keepalive = TcpKeepalive::new().with_time(Duration::from_secs(keepalive_secs));
.with_time(Duration::from_secs(keepalive_secs)); socket
socket.set_tcp_keepalive(&keepalive) .set_tcp_keepalive(&keepalive)
.map_err(|e| TransportError::StartFailed(format!("set keepalive: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("set keepalive: {}", e)))?;
} }
// Buffer sizes // Buffer sizes
socket.set_recv_buffer_size(config.recv_buf_size()) socket
.set_recv_buffer_size(config.recv_buf_size())
.map_err(|e| TransportError::StartFailed(format!("set recv buffer: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("set recv buffer: {}", e)))?;
socket.set_send_buffer_size(config.send_buf_size()) socket
.set_send_buffer_size(config.send_buf_size())
.map_err(|e| TransportError::StartFailed(format!("set send buffer: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("set send buffer: {}", e)))?;
Ok(()) Ok(())
@@ -967,22 +989,26 @@ fn configure_accepted_socket(
recv_buf: usize, recv_buf: usize,
send_buf: usize, send_buf: usize,
) -> Result<(), TransportError> { ) -> Result<(), TransportError> {
let socket = socket2::SockRef::from(stream).try_clone() let socket = socket2::SockRef::from(stream)
.try_clone()
.map_err(|e| TransportError::StartFailed(format!("clone socket: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("clone socket: {}", e)))?;
socket.set_tcp_nodelay(nodelay) socket
.set_tcp_nodelay(nodelay)
.map_err(|e| TransportError::StartFailed(format!("set nodelay: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("set nodelay: {}", e)))?;
if keepalive_secs > 0 { if keepalive_secs > 0 {
let keepalive = TcpKeepalive::new() let keepalive = TcpKeepalive::new().with_time(Duration::from_secs(keepalive_secs));
.with_time(Duration::from_secs(keepalive_secs)); socket
socket.set_tcp_keepalive(&keepalive) .set_tcp_keepalive(&keepalive)
.map_err(|e| TransportError::StartFailed(format!("set keepalive: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("set keepalive: {}", e)))?;
} }
socket.set_recv_buffer_size(recv_buf) socket
.set_recv_buffer_size(recv_buf)
.map_err(|e| TransportError::StartFailed(format!("set recv buffer: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("set recv buffer: {}", e)))?;
socket.set_send_buffer_size(send_buf) socket
.set_send_buffer_size(send_buf)
.map_err(|e| TransportError::StartFailed(format!("set send buffer: {}", e)))?; .map_err(|e| TransportError::StartFailed(format!("set send buffer: {}", e)))?;
Ok(()) Ok(())
@@ -1028,7 +1054,7 @@ fn read_mss_mtu(stream: &std::net::TcpStream, default_mtu: u16) -> u16 {
mod tests { mod tests {
use super::*; use super::*;
use crate::transport::packet_channel; use crate::transport::packet_channel;
use tokio::time::{timeout, Duration}; use tokio::time::{Duration, timeout};
fn make_config() -> TcpConfig { fn make_config() -> TcpConfig {
TcpConfig { TcpConfig {
@@ -1064,7 +1090,8 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_start_outbound_only() { async fn test_start_outbound_only() {
let (tx, _rx) = packet_channel(100); let (tx, _rx) = packet_channel(100);
let mut transport = TcpTransport::new(TransportId::new(1), None, make_outbound_config(), tx); let mut transport =
TcpTransport::new(TransportId::new(1), None, make_outbound_config(), tx);
transport.start_async().await.unwrap(); transport.start_async().await.unwrap();
assert_eq!(transport.state(), TransportState::Up); assert_eq!(transport.state(), TransportState::Up);
@@ -1135,10 +1162,7 @@ mod tests {
} }
let bytes_sent = t1 let bytes_sent = t1
.send_async( .send_async(&TransportAddr::from_string(&addr2.to_string()), &frame)
&TransportAddr::from_string(&addr2.to_string()),
&frame,
)
.await .await
.unwrap(); .unwrap();
assert_eq!(bytes_sent, frame.len()); assert_eq!(bytes_sent, frame.len());
@@ -1176,12 +1200,9 @@ mod tests {
msg1_frame[2..4].copy_from_slice(&110u16.to_le_bytes()); // payload_len = 110 msg1_frame[2..4].copy_from_slice(&110u16.to_le_bytes()); // payload_len = 110
// Send from t1 to t2 // Send from t1 to t2
t1.send_async( t1.send_async(&TransportAddr::from_string(&addr2.to_string()), &msg1_frame)
&TransportAddr::from_string(&addr2.to_string()), .await
&msg1_frame, .unwrap();
)
.await
.unwrap();
let packet = timeout(Duration::from_secs(2), rx2.recv()) let packet = timeout(Duration::from_secs(2), rx2.recv())
.await .await
@@ -1196,12 +1217,9 @@ mod tests {
msg2_frame[2..4].copy_from_slice(&65u16.to_le_bytes()); // payload_len = 65 msg2_frame[2..4].copy_from_slice(&65u16.to_le_bytes()); // payload_len = 65
// Send from t2 to t1 // Send from t2 to t1
t2.send_async( t2.send_async(&TransportAddr::from_string(&addr1.to_string()), &msg2_frame)
&TransportAddr::from_string(&addr1.to_string()), .await
&msg2_frame, .unwrap();
)
.await
.unwrap();
let packet = timeout(Duration::from_secs(2), rx1.recv()) let packet = timeout(Duration::from_secs(2), rx1.recv())
.await .await
@@ -1226,7 +1244,10 @@ mod tests {
// Try to connect to a non-routable address (should timeout) // Try to connect to a non-routable address (should timeout)
let result = transport let result = transport
.send_async(&TransportAddr::from_string("192.0.2.1:2121"), b"\x00\x00\x04\x00test1234567890123456789012345678") .send_async(
&TransportAddr::from_string("192.0.2.1:2121"),
b"\x00\x00\x04\x00test1234567890123456789012345678",
)
.await; .await;
assert!(result.is_err()); assert!(result.is_err());
@@ -1410,8 +1431,7 @@ mod tests {
t1.send_async(&remote, &msg1).await.unwrap(); t1.send_async(&remote, &msg1).await.unwrap();
let packet = timeout(Duration::from_secs(2), rx1.recv()) let packet = timeout(Duration::from_secs(2), rx1.recv()).await;
.await;
// We receive on rx1 but that's the wrong receiver — t2's rx gets the packet // We receive on rx1 but that's the wrong receiver — t2's rx gets the packet
// Just verify send didn't error // Just verify send didn't error
drop(packet); drop(packet);
@@ -1472,7 +1492,10 @@ mod tests {
// Connect first time // Connect first time
t1.connect_async(&remote).await.unwrap(); t1.connect_async(&remote).await.unwrap();
tokio::time::sleep(Duration::from_millis(200)).await; tokio::time::sleep(Duration::from_millis(200)).await;
assert_eq!(t1.connection_state_sync(&remote), ConnectionState::Connected); assert_eq!(
t1.connection_state_sync(&remote),
ConnectionState::Connected
);
// Second connect should be a no-op (already connected) // Second connect should be a no-op (already connected)
t1.connect_async(&remote).await.unwrap(); t1.connect_async(&remote).await.unwrap();
@@ -1498,7 +1521,10 @@ mod tests {
// Connect first, then send // Connect first, then send
t1.connect_async(&remote).await.unwrap(); t1.connect_async(&remote).await.unwrap();
tokio::time::sleep(Duration::from_millis(200)).await; tokio::time::sleep(Duration::from_millis(200)).await;
assert_eq!(t1.connection_state_sync(&remote), ConnectionState::Connected); assert_eq!(
t1.connection_state_sync(&remote),
ConnectionState::Connected
);
// Build valid FMP msg1 frame // Build valid FMP msg1 frame
let mut msg1 = vec![0xAA; 114]; let mut msg1 = vec![0xAA; 114];
@@ -1524,9 +1550,7 @@ mod tests {
let (tx, _rx) = packet_channel(100); let (tx, _rx) = packet_channel(100);
let transport = TcpTransport::new(TransportId::new(1), None, make_config(), tx); let transport = TcpTransport::new(TransportId::new(1), None, make_config(), tx);
let state = transport.connection_state_sync( let state = transport.connection_state_sync(&TransportAddr::from_string("unknown:1234"));
&TransportAddr::from_string("unknown:1234"),
);
assert_eq!(state, ConnectionState::None); assert_eq!(state, ConnectionState::None);
} }
@@ -1605,10 +1629,7 @@ mod tests {
tokio::time::sleep(Duration::from_millis(100)).await; tokio::time::sleep(Duration::from_millis(100)).await;
} }
assert_eq!( assert_eq!(t1.connection_state_sync(&addr), ConnectionState::Connected,);
t1.connection_state_sync(&addr),
ConnectionState::Connected,
);
t1.stop_async().await.unwrap(); t1.stop_async().await.unwrap();
t2.stop_async().await.unwrap(); t2.stop_async().await.unwrap();

View File

@@ -31,7 +31,10 @@ pub enum StreamError {
/// Unknown FMP phase byte — protocol error, close connection. /// Unknown FMP phase byte — protocol error, close connection.
UnknownPhase(u8), UnknownPhase(u8),
/// Payload length exceeds the connection's MTU — corrupted or malicious. /// Payload length exceeds the connection's MTU — corrupted or malicious.
PayloadTooLarge { payload_len: u16, max_payload_len: u16 }, PayloadTooLarge {
payload_len: u16,
max_payload_len: u16,
},
/// Handshake packet has unexpected payload_len for its phase. /// Handshake packet has unexpected payload_len for its phase.
HandshakeSizeMismatch { phase: u8, expected: u16, got: u16 }, HandshakeSizeMismatch { phase: u8, expected: u16, got: u16 },
/// I/O error (including EOF). /// I/O error (including EOF).
@@ -43,11 +46,26 @@ impl std::fmt::Display for StreamError {
match self { match self {
StreamError::UnknownVersion(v) => write!(f, "unknown FMP version: {}", v), StreamError::UnknownVersion(v) => write!(f, "unknown FMP version: {}", v),
StreamError::UnknownPhase(p) => write!(f, "unknown FMP phase: 0x{:02x}", p), StreamError::UnknownPhase(p) => write!(f, "unknown FMP phase: 0x{:02x}", p),
StreamError::PayloadTooLarge { payload_len, max_payload_len } => { StreamError::PayloadTooLarge {
write!(f, "payload_len {} exceeds max {}", payload_len, max_payload_len) payload_len,
max_payload_len,
} => {
write!(
f,
"payload_len {} exceeds max {}",
payload_len, max_payload_len
)
} }
StreamError::HandshakeSizeMismatch { phase, expected, got } => { StreamError::HandshakeSizeMismatch {
write!(f, "handshake phase 0x{:x}: expected payload_len {}, got {}", phase, expected, got) phase,
expected,
got,
} => {
write!(
f,
"handshake phase 0x{:x}: expected payload_len {}, got {}",
phase, expected, got
)
} }
StreamError::Io(e) => write!(f, "io: {}", e), StreamError::Io(e) => write!(f, "io: {}", e),
} }
@@ -172,7 +190,8 @@ mod tests {
/// Build a minimal established frame with the given payload_len. /// Build a minimal established frame with the given payload_len.
/// Layout: [ver+phase:1][flags:1][payload_len:2 LE][12 bytes header][payload_len bytes][16 bytes tag] /// Layout: [ver+phase:1][flags:1][payload_len:2 LE][12 bytes header][payload_len bytes][16 bytes tag]
fn build_established_frame(payload_len: u16) -> Vec<u8> { fn build_established_frame(payload_len: u16) -> Vec<u8> {
let total = PREFIX_SIZE + ESTABLISHED_REMAINING_HEADER + payload_len as usize + AEAD_TAG_SIZE; let total =
PREFIX_SIZE + ESTABLISHED_REMAINING_HEADER + payload_len as usize + AEAD_TAG_SIZE;
let mut frame = vec![0u8; total]; let mut frame = vec![0u8; total];
frame[0] = 0x00; // ver=0, phase=0 (established) frame[0] = 0x00; // ver=0, phase=0 (established)
frame[1] = 0x00; // flags frame[1] = 0x00; // flags
@@ -304,7 +323,10 @@ mod tests {
let mut cursor = Cursor::new(frame); let mut cursor = Cursor::new(frame);
let err = read_fmp_packet(&mut cursor, 1400).await.unwrap_err(); let err = read_fmp_packet(&mut cursor, 1400).await.unwrap_err();
assert!(matches!(err, StreamError::HandshakeSizeMismatch { phase: 0x1, .. })); assert!(matches!(
err,
StreamError::HandshakeSizeMismatch { phase: 0x1, .. }
));
} }
#[tokio::test] #[tokio::test]
@@ -316,7 +338,10 @@ mod tests {
let mut cursor = Cursor::new(frame); let mut cursor = Cursor::new(frame);
let err = read_fmp_packet(&mut cursor, 1400).await.unwrap_err(); let err = read_fmp_packet(&mut cursor, 1400).await.unwrap_err();
assert!(matches!(err, StreamError::HandshakeSizeMismatch { phase: 0x2, .. })); assert!(matches!(
err,
StreamError::HandshakeSizeMismatch { phase: 0x2, .. }
));
} }
#[tokio::test] #[tokio::test]

View File

@@ -6,11 +6,11 @@
use std::fmt; use std::fmt;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use serde::Serialize;
use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader}; use tokio::io::{AsyncBufReadExt, AsyncRead, AsyncWrite, AsyncWriteExt, BufReader};
use tokio::net::TcpStream; use tokio::net::TcpStream;
#[cfg(unix)] #[cfg(unix)]
use tokio::net::UnixStream; use tokio::net::UnixStream;
use serde::Serialize;
use tracing::debug; use tracing::debug;
// ============================================================================ // ============================================================================
@@ -67,10 +67,7 @@ impl ControlAuth {
/// ///
/// - `"cookie"` or `"cookie:/path/to/cookie"` → Cookie auth /// - `"cookie"` or `"cookie:/path/to/cookie"` → Cookie auth
/// - `"password:secret"` → Password auth /// - `"password:secret"` → Password auth
pub fn from_config( pub fn from_config(auth_str: &str, default_cookie_path: &str) -> Result<Self, TorControlError> {
auth_str: &str,
default_cookie_path: &str,
) -> Result<Self, TorControlError> {
if auth_str == "cookie" { if auth_str == "cookie" {
Ok(Self::Cookie(PathBuf::from(default_cookie_path))) Ok(Self::Cookie(PathBuf::from(default_cookie_path)))
} else if let Some(path) = auth_str.strip_prefix("cookie:") { } else if let Some(path) = auth_str.strip_prefix("cookie:") {
@@ -286,10 +283,7 @@ impl TorControlClient {
pub async fn traffic_read(&mut self) -> Result<u64, TorControlError> { pub async fn traffic_read(&mut self) -> Result<u64, TorControlError> {
let value = self.getinfo("traffic/read").await?; let value = self.getinfo("traffic/read").await?;
value.trim().parse::<u64>().map_err(|_| { value.trim().parse::<u64>().map_err(|_| {
TorControlError::ProtocolError(format!( TorControlError::ProtocolError(format!("invalid traffic/read value: '{}'", value))
"invalid traffic/read value: '{}'",
value
))
}) })
} }
@@ -297,10 +291,7 @@ impl TorControlClient {
pub async fn traffic_written(&mut self) -> Result<u64, TorControlError> { pub async fn traffic_written(&mut self) -> Result<u64, TorControlError> {
let value = self.getinfo("traffic/written").await?; let value = self.getinfo("traffic/written").await?;
value.trim().parse::<u64>().map_err(|_| { value.trim().parse::<u64>().map_err(|_| {
TorControlError::ProtocolError(format!( TorControlError::ProtocolError(format!("invalid traffic/written value: '{}'", value))
"invalid traffic/written value: '{}'",
value
))
}) })
} }
@@ -327,7 +318,10 @@ impl TorControlClient {
/// Returns a list of addresses Tor is listening on for SOCKS connections. /// Returns a list of addresses Tor is listening on for SOCKS connections.
pub async fn socks_listeners(&mut self) -> Result<Vec<String>, TorControlError> { pub async fn socks_listeners(&mut self) -> Result<Vec<String>, TorControlError> {
let value = self.getinfo("net/listeners/socks").await?; let value = self.getinfo("net/listeners/socks").await?;
Ok(value.split_whitespace().map(|s| s.trim_matches('"').to_string()).collect()) Ok(value
.split_whitespace()
.map(|s| s.trim_matches('"').to_string())
.collect())
} }
/// Collect all monitoring info in a single batch of queries. /// Collect all monitoring info in a single batch of queries.
@@ -336,7 +330,10 @@ impl TorControlClient {
let circuit_established = self.is_circuit_established().await.unwrap_or(false); let circuit_established = self.is_circuit_established().await.unwrap_or(false);
let traffic_read = self.traffic_read().await.unwrap_or(0); let traffic_read = self.traffic_read().await.unwrap_or(0);
let traffic_written = self.traffic_written().await.unwrap_or(0); let traffic_written = self.traffic_written().await.unwrap_or(0);
let network_liveness = self.network_liveness().await.unwrap_or_else(|_| "unknown".into()); let network_liveness = self
.network_liveness()
.await
.unwrap_or_else(|_| "unknown".into());
let version = self.version().await.unwrap_or_else(|_| "unknown".into()); let version = self.version().await.unwrap_or_else(|_| "unknown".into());
let dormant = self.is_dormant().await.unwrap_or(false); let dormant = self.is_dormant().await.unwrap_or(false);
@@ -393,10 +390,7 @@ impl TorControlClient {
} }
let code: u16 = line[..3].parse().map_err(|_| { let code: u16 = line[..3].parse().map_err(|_| {
TorControlError::ProtocolError(format!( TorControlError::ProtocolError(format!("invalid response code in: '{}'", line))
"invalid response code in: '{}'",
line
))
})?; })?;
let separator = line.as_bytes()[3]; let separator = line.as_bytes()[3];
@@ -426,8 +420,7 @@ impl TorControlClient {
"connection closed during multi-line response".into(), "connection closed during multi-line response".into(),
)); ));
} }
let dot_line = let dot_line = line_buf.trim_end_matches(['\r', '\n']);
line_buf.trim_end_matches(['\r', '\n']);
if dot_line == "." { if dot_line == "." {
break; break;
} }
@@ -464,7 +457,11 @@ struct ControlResponse {
/// Read a Tor control cookie file (32 bytes of raw binary). /// Read a Tor control cookie file (32 bytes of raw binary).
fn read_cookie_file(path: &Path) -> Result<Vec<u8>, TorControlError> { fn read_cookie_file(path: &Path) -> Result<Vec<u8>, TorControlError> {
let data = std::fs::read(path).map_err(|e| { let data = std::fs::read(path).map_err(|e| {
TorControlError::AuthFailed(format!("failed to read cookie file '{}': {}", path.display(), e)) TorControlError::AuthFailed(format!(
"failed to read cookie file '{}': {}",
path.display(),
e
))
})?; })?;
if data.len() != 32 { if data.len() != 32 {
@@ -644,7 +641,9 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_authenticate_password() { async fn test_authenticate_password() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
.await
.unwrap();
let auth = ControlAuth::Password("testpass".to_string()); let auth = ControlAuth::Password("testpass".to_string());
client.authenticate(&auth).await.unwrap(); client.authenticate(&auth).await.unwrap();
@@ -659,7 +658,9 @@ mod tests {
let cookie_path = dir.path().join("cookie"); let cookie_path = dir.path().join("cookie");
std::fs::write(&cookie_path, [0xAA; 32]).unwrap(); std::fs::write(&cookie_path, [0xAA; 32]).unwrap();
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
.await
.unwrap();
let auth = ControlAuth::Cookie(cookie_path); let auth = ControlAuth::Cookie(cookie_path);
client.authenticate(&auth).await.unwrap(); client.authenticate(&auth).await.unwrap();
} }
@@ -667,7 +668,9 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_get_bootstrap_phase() { async fn test_get_bootstrap_phase() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
.await
.unwrap();
let auth = ControlAuth::Password("testpass".to_string()); let auth = ControlAuth::Password("testpass".to_string());
client.authenticate(&auth).await.unwrap(); client.authenticate(&auth).await.unwrap();
@@ -682,7 +685,9 @@ mod tests {
reject_auth: true, reject_auth: true,
}) })
.await; .await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
.await
.unwrap();
let auth = ControlAuth::Password("wrongpass".to_string()); let auth = ControlAuth::Password("wrongpass".to_string());
let result = client.authenticate(&auth).await; let result = client.authenticate(&auth).await;
@@ -705,8 +710,13 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_is_circuit_established() { async fn test_is_circuit_established() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
client.authenticate(&ControlAuth::Password("test".into())).await.unwrap(); .await
.unwrap();
client
.authenticate(&ControlAuth::Password("test".into()))
.await
.unwrap();
assert!(client.is_circuit_established().await.unwrap()); assert!(client.is_circuit_established().await.unwrap());
} }
@@ -714,8 +724,13 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_traffic_counters() { async fn test_traffic_counters() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
client.authenticate(&ControlAuth::Password("test".into())).await.unwrap(); .await
.unwrap();
client
.authenticate(&ControlAuth::Password("test".into()))
.await
.unwrap();
assert_eq!(client.traffic_read().await.unwrap(), 1048576); assert_eq!(client.traffic_read().await.unwrap(), 1048576);
assert_eq!(client.traffic_written().await.unwrap(), 524288); assert_eq!(client.traffic_written().await.unwrap(), 524288);
@@ -724,8 +739,13 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_network_liveness() { async fn test_network_liveness() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
client.authenticate(&ControlAuth::Password("test".into())).await.unwrap(); .await
.unwrap();
client
.authenticate(&ControlAuth::Password("test".into()))
.await
.unwrap();
assert_eq!(client.network_liveness().await.unwrap(), "up"); assert_eq!(client.network_liveness().await.unwrap(), "up");
} }
@@ -733,8 +753,13 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_version() { async fn test_version() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
client.authenticate(&ControlAuth::Password("test".into())).await.unwrap(); .await
.unwrap();
client
.authenticate(&ControlAuth::Password("test".into()))
.await
.unwrap();
assert_eq!(client.version().await.unwrap(), "0.4.8.10"); assert_eq!(client.version().await.unwrap(), "0.4.8.10");
} }
@@ -742,8 +767,13 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_dormant() { async fn test_dormant() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
client.authenticate(&ControlAuth::Password("test".into())).await.unwrap(); .await
.unwrap();
client
.authenticate(&ControlAuth::Password("test".into()))
.await
.unwrap();
assert!(!client.is_dormant().await.unwrap()); assert!(!client.is_dormant().await.unwrap());
} }
@@ -751,8 +781,13 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_socks_listeners() { async fn test_socks_listeners() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
client.authenticate(&ControlAuth::Password("test".into())).await.unwrap(); .await
.unwrap();
client
.authenticate(&ControlAuth::Password("test".into()))
.await
.unwrap();
let listeners = client.socks_listeners().await.unwrap(); let listeners = client.socks_listeners().await.unwrap();
assert_eq!(listeners, vec!["127.0.0.1:9050"]); assert_eq!(listeners, vec!["127.0.0.1:9050"]);
@@ -761,8 +796,13 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn test_monitoring_snapshot() { async fn test_monitoring_snapshot() {
let mock = MockTorControlServer::start().await; let mock = MockTorControlServer::start().await;
let mut client = TorControlClient::connect(&mock.addr().to_string()).await.unwrap(); let mut client = TorControlClient::connect(&mock.addr().to_string())
client.authenticate(&ControlAuth::Password("test".into())).await.unwrap(); .await
.unwrap();
client
.authenticate(&ControlAuth::Password("test".into()))
.await
.unwrap();
let info = client.monitoring_snapshot().await.unwrap(); let info = client.monitoring_snapshot().await.unwrap();
assert_eq!(info.bootstrap, 100); assert_eq!(info.bootstrap, 100);

View File

@@ -33,7 +33,9 @@ impl MockTorControlServer {
/// Start a mock control server with custom options. /// Start a mock control server with custom options.
pub async fn start_with_options(options: MockOptions) -> Self { pub async fn start_with_options(options: MockOptions) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind mock control"); let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind mock control");
let addr = listener.local_addr().expect("local addr"); let addr = listener.local_addr().expect("local addr");
let handle = tokio::spawn(async move { let handle = tokio::spawn(async move {
@@ -63,45 +65,39 @@ impl MockTorControlServer {
let _ = writer.write_all(b"250 OK\r\n").await; let _ = writer.write_all(b"250 OK\r\n").await;
} }
} else if !authenticated { } else if !authenticated {
let _ = writer let _ = writer.write_all(b"514 Authentication required\r\n").await;
.write_all(b"514 Authentication required\r\n")
.await;
} else if cmd.starts_with("GETINFO status/bootstrap-phase") { } else if cmd.starts_with("GETINFO status/bootstrap-phase") {
let _ = writer.write_all( let _ = writer.write_all(
b"250-status/bootstrap-phase=NOTICE BOOTSTRAP PROGRESS=100 TAG=done SUMMARY=\"Done\"\r\n250 OK\r\n", b"250-status/bootstrap-phase=NOTICE BOOTSTRAP PROGRESS=100 TAG=done SUMMARY=\"Done\"\r\n250 OK\r\n",
).await; ).await;
} else if cmd.starts_with("GETINFO status/circuit-established") { } else if cmd.starts_with("GETINFO status/circuit-established") {
let _ = writer.write_all(
b"250-status/circuit-established=1\r\n250 OK\r\n",
).await;
} else if cmd.starts_with("GETINFO traffic/read") {
let _ = writer.write_all(
b"250-traffic/read=1048576\r\n250 OK\r\n",
).await;
} else if cmd.starts_with("GETINFO traffic/written") {
let _ = writer.write_all(
b"250-traffic/written=524288\r\n250 OK\r\n",
).await;
} else if cmd.starts_with("GETINFO network-liveness") {
let _ = writer.write_all(
b"250-network-liveness=up\r\n250 OK\r\n",
).await;
} else if cmd.starts_with("GETINFO version") {
let _ = writer.write_all(
b"250-version=0.4.8.10\r\n250 OK\r\n",
).await;
} else if cmd.starts_with("GETINFO dormant") {
let _ = writer.write_all(
b"250-dormant=0\r\n250 OK\r\n",
).await;
} else if cmd.starts_with("GETINFO net/listeners/socks") {
let _ = writer.write_all(
b"250-net/listeners/socks=\"127.0.0.1:9050\"\r\n250 OK\r\n",
).await;
} else {
let _ = writer let _ = writer
.write_all(b"510 Unrecognized command\r\n") .write_all(b"250-status/circuit-established=1\r\n250 OK\r\n")
.await; .await;
} else if cmd.starts_with("GETINFO traffic/read") {
let _ = writer
.write_all(b"250-traffic/read=1048576\r\n250 OK\r\n")
.await;
} else if cmd.starts_with("GETINFO traffic/written") {
let _ = writer
.write_all(b"250-traffic/written=524288\r\n250 OK\r\n")
.await;
} else if cmd.starts_with("GETINFO network-liveness") {
let _ = writer
.write_all(b"250-network-liveness=up\r\n250 OK\r\n")
.await;
} else if cmd.starts_with("GETINFO version") {
let _ = writer
.write_all(b"250-version=0.4.8.10\r\n250 OK\r\n")
.await;
} else if cmd.starts_with("GETINFO dormant") {
let _ = writer.write_all(b"250-dormant=0\r\n250 OK\r\n").await;
} else if cmd.starts_with("GETINFO net/listeners/socks") {
let _ = writer
.write_all(b"250-net/listeners/socks=\"127.0.0.1:9050\"\r\n250 OK\r\n")
.await;
} else {
let _ = writer.write_all(b"510 Unrecognized command\r\n").await;
} }
let _ = writer.flush().await; let _ = writer.flush().await;

View File

@@ -70,7 +70,10 @@ impl MockSocks5Server {
// === Method negotiation === // === Method negotiation ===
// Client sends: [version, nmethods, methods...] // Client sends: [version, nmethods, methods...]
let mut ver_nmethods = [0u8; 2]; let mut ver_nmethods = [0u8; 2];
client.read_exact(&mut ver_nmethods).await.expect("read version+nmethods"); client
.read_exact(&mut ver_nmethods)
.await
.expect("read version+nmethods");
assert_eq!(ver_nmethods[0], SOCKS_VERSION, "expected SOCKS5"); assert_eq!(ver_nmethods[0], SOCKS_VERSION, "expected SOCKS5");
let nmethods = ver_nmethods[1] as usize; let nmethods = ver_nmethods[1] as usize;
@@ -87,14 +90,23 @@ impl MockSocks5Server {
}; };
// Reply: [version, selected_method] // Reply: [version, selected_method]
client.write_all(&[SOCKS_VERSION, selected]).await.expect("write method reply"); client
.write_all(&[SOCKS_VERSION, selected])
.await
.expect("write method reply");
// === Username/password sub-negotiation (RFC 1929) === // === Username/password sub-negotiation (RFC 1929) ===
if selected == AUTH_PASSWORD { if selected == AUTH_PASSWORD {
// Client sends: [ver(1), ulen(1), uname(ulen), plen(1), passwd(plen)] // Client sends: [ver(1), ulen(1), uname(ulen), plen(1), passwd(plen)]
let mut subneg_header = [0u8; 2]; let mut subneg_header = [0u8; 2];
client.read_exact(&mut subneg_header).await.expect("read subneg header"); client
assert_eq!(subneg_header[0], AUTH_SUBNEG_VERSION, "expected auth subneg v1"); .read_exact(&mut subneg_header)
.await
.expect("read subneg header");
assert_eq!(
subneg_header[0], AUTH_SUBNEG_VERSION,
"expected auth subneg v1"
);
let ulen = subneg_header[1] as usize; let ulen = subneg_header[1] as usize;
let mut uname = vec![0u8; ulen]; let mut uname = vec![0u8; ulen];
@@ -107,14 +119,19 @@ impl MockSocks5Server {
client.read_exact(&mut passwd).await.expect("read password"); client.read_exact(&mut passwd).await.expect("read password");
// Always accept (Tor uses these as isolation keys, not real auth) // Always accept (Tor uses these as isolation keys, not real auth)
client.write_all(&[AUTH_SUBNEG_VERSION, AUTH_SUBNEG_SUCCESS]) client
.await.expect("write subneg reply"); .write_all(&[AUTH_SUBNEG_VERSION, AUTH_SUBNEG_SUCCESS])
.await
.expect("write subneg reply");
} }
// === Connect request === // === Connect request ===
// Client sends: [version, cmd, rsv, atyp, addr..., port] // Client sends: [version, cmd, rsv, atyp, addr..., port]
let mut header = [0u8; 4]; let mut header = [0u8; 4];
client.read_exact(&mut header).await.expect("read connect header"); client
.read_exact(&mut header)
.await
.expect("read connect header");
assert_eq!(header[0], SOCKS_VERSION); assert_eq!(header[0], SOCKS_VERSION);
assert_eq!(header[1], CMD_CONNECT); assert_eq!(header[1], CMD_CONNECT);
@@ -122,14 +139,23 @@ impl MockSocks5Server {
match header[3] { match header[3] {
ATYP_IPV4 => { ATYP_IPV4 => {
let mut addr_port = [0u8; 6]; // 4 IP + 2 port let mut addr_port = [0u8; 6]; // 4 IP + 2 port
client.read_exact(&mut addr_port).await.expect("read IPv4 addr"); client
.read_exact(&mut addr_port)
.await
.expect("read IPv4 addr");
} }
ATYP_DOMAIN => { ATYP_DOMAIN => {
let mut len_buf = [0u8; 1]; let mut len_buf = [0u8; 1];
client.read_exact(&mut len_buf).await.expect("read domain len"); client
.read_exact(&mut len_buf)
.await
.expect("read domain len");
let domain_len = len_buf[0] as usize; let domain_len = len_buf[0] as usize;
let mut domain_port = vec![0u8; domain_len + 2]; // domain + 2 port let mut domain_port = vec![0u8; domain_len + 2]; // domain + 2 port
client.read_exact(&mut domain_port).await.expect("read domain addr"); client
.read_exact(&mut domain_port)
.await
.expect("read domain addr");
} }
other => panic!("unsupported ATYP: {}", other), other => panic!("unsupported ATYP: {}", other),
} }
@@ -145,8 +171,12 @@ impl MockSocks5Server {
REP_SUCCESS, REP_SUCCESS,
0x00, // RSV 0x00, // RSV
ATYP_IPV4, ATYP_IPV4,
0, 0, 0, 0, // bind addr 0,
0, 0, // bind port 0,
0,
0, // bind addr
0,
0, // bind port
]; ];
client.write_all(&reply).await.expect("write connect reply"); client.write_all(&reply).await.expect("write connect reply");

View File

@@ -42,8 +42,8 @@ use std::net::SocketAddr;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
use tokio::io::AsyncWriteExt; use tokio::io::AsyncWriteExt;
use tokio::net::{TcpListener, TcpStream};
use tokio::net::tcp::OwnedWriteHalf; use tokio::net::tcp::OwnedWriteHalf;
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::Mutex; use tokio::sync::Mutex;
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
use tokio::time::Instant; use tokio::time::Instant;
@@ -92,16 +92,15 @@ fn parse_tor_addr(addr: &TransportAddr) -> Result<TorAddr, TransportError> {
} else { } else {
// Hostname:port — pass through SOCKS5 for Tor-side DNS resolution // Hostname:port — pass through SOCKS5 for Tor-side DNS resolution
let (host, port_str) = s.rsplit_once(':').ok_or_else(|| { let (host, port_str) = s.rsplit_once(':').ok_or_else(|| {
TransportError::InvalidAddress(format!( TransportError::InvalidAddress(format!("invalid address (expected host:port): {}", s))
"invalid address (expected host:port): {}", s
))
})?;
let port: u16 = port_str.parse().map_err(|_| {
TransportError::InvalidAddress(format!("invalid port: {}", s))
})?; })?;
let port: u16 = port_str
.parse()
.map_err(|_| TransportError::InvalidAddress(format!("invalid port: {}", s)))?;
if !host.contains('.') { if !host.contains('.') {
return Err(TransportError::InvalidAddress(format!( return Err(TransportError::InvalidAddress(format!(
"hostname must be fully qualified (contain a dot): {}", host "hostname must be fully qualified (contain a dot): {}",
host
))); )));
} }
Ok(TorAddr::ClearnetHostname(host.to_string(), port)) Ok(TorAddr::ClearnetHostname(host.to_string(), port))
@@ -306,17 +305,16 @@ impl TorTransport {
} }
// Connect to Tor control port // Connect to Tor control port
let mut client = TorControlClient::connect(&control_addr).await.map_err(|e| { let mut client = TorControlClient::connect(&control_addr)
self.stats.record_control_error(); .await
TransportError::StartFailed(format!("Tor control port: {}", e)) .map_err(|e| {
})?; self.stats.record_control_error();
TransportError::StartFailed(format!("Tor control port: {}", e))
})?;
// Authenticate // Authenticate
let auth = ControlAuth::from_config( let auth = ControlAuth::from_config(self.config.control_auth(), self.config.cookie_path())
self.config.control_auth(), .map_err(|e| TransportError::StartFailed(format!("Tor auth config: {}", e)))?;
self.config.cookie_path(),
)
.map_err(|e| TransportError::StartFailed(format!("Tor auth config: {}", e)))?;
client.authenticate(&auth).await.map_err(|e| { client.authenticate(&auth).await.map_err(|e| {
self.stats.record_control_error(); self.stats.record_control_error();
@@ -370,9 +368,9 @@ impl TorTransport {
bind_addr, e bind_addr, e
)) ))
})?; })?;
let local_addr = listener.local_addr().map_err(|e| { let local_addr = listener
TransportError::StartFailed(format!("failed to get local addr: {}", e)) .local_addr()
})?; .map_err(|e| TransportError::StartFailed(format!("failed to get local addr: {}", e)))?;
info!( info!(
onion_address = %onion_addr, onion_address = %onion_addr,
@@ -417,7 +415,8 @@ impl TorTransport {
/// Non-fatal: logs a warning on failure and continues without monitoring. /// Non-fatal: logs a warning on failure and continues without monitoring.
async fn try_connect_control_port(&mut self) { async fn try_connect_control_port(&mut self) {
let control_addr = self.config.control_addr().to_string(); let control_addr = self.config.control_addr().to_string();
if !control_addr.starts_with('/') && !control_addr.starts_with("./") if !control_addr.starts_with('/')
&& !control_addr.starts_with("./")
&& let Err(e) = validate_host_port(&control_addr, "control_addr") && let Err(e) = validate_host_port(&control_addr, "control_addr")
{ {
warn!( warn!(
@@ -441,20 +440,18 @@ impl TorTransport {
} }
}; };
let auth = match ControlAuth::from_config( let auth =
self.config.control_auth(), match ControlAuth::from_config(self.config.control_auth(), self.config.cookie_path()) {
self.config.cookie_path(), Ok(a) => a,
) { Err(e) => {
Ok(a) => a, warn!(
Err(e) => { transport_id = %self.transport_id,
warn!( error = %e,
transport_id = %self.transport_id, "Tor control auth config error, monitoring disabled"
error = %e, );
"Tor control auth config error, monitoring disabled" return;
); }
return; };
}
};
let mut client = client; let mut client = client;
if let Err(e) = client.authenticate(&auth).await { if let Err(e) = client.authenticate(&auth).await {
@@ -567,9 +564,7 @@ impl TorTransport {
Ok(info) => { Ok(info) => {
// Log bootstrap milestones // Log bootstrap milestones
for &milestone in &[25u8, 50, 75, 100] { for &milestone in &[25u8, 50, 75, 100] {
if info.bootstrap >= milestone if info.bootstrap >= milestone && last_bootstrap < milestone {
&& last_bootstrap < milestone
{
info!( info!(
transport_id = %transport_id, transport_id = %transport_id,
bootstrap = info.bootstrap, bootstrap = info.bootstrap,
@@ -598,9 +593,7 @@ impl TorTransport {
last_bootstrap = info.bootstrap; last_bootstrap = info.bootstrap;
// Network liveness transitions // Network liveness transitions
if !last_liveness.is_empty() if !last_liveness.is_empty() && info.network_liveness != last_liveness {
&& info.network_liveness != last_liveness
{
warn!( warn!(
transport_id = %transport_id, transport_id = %transport_id,
from = %last_liveness, from = %last_liveness,
@@ -732,16 +725,23 @@ impl TorTransport {
let connect_start = Instant::now(); let connect_start = Instant::now();
let socks_result = tokio::time::timeout(Duration::from_millis(timeout_ms), async { let socks_result = tokio::time::timeout(Duration::from_millis(timeout_ms), async {
match &tor_addr { match &tor_addr {
TorAddr::Onion(host, port) TorAddr::Onion(host, port) | TorAddr::ClearnetHostname(host, port) => {
| TorAddr::ClearnetHostname(host, port) => {
Socks5Stream::connect_with_password( Socks5Stream::connect_with_password(
proxy_addr, (host.as_str(), *port), "fips", &isolation_key, proxy_addr,
).await (host.as_str(), *port),
"fips",
&isolation_key,
)
.await
} }
TorAddr::Clearnet(socket_addr) => { TorAddr::Clearnet(socket_addr) => {
Socks5Stream::connect_with_password( Socks5Stream::connect_with_password(
proxy_addr, *socket_addr, "fips", &isolation_key, proxy_addr,
).await *socket_addr,
"fips",
&isolation_key,
)
.await
} }
} }
}) })
@@ -877,26 +877,28 @@ impl TorTransport {
let task = tokio::spawn(async move { let task = tokio::spawn(async move {
// SOCKS5 CONNECT through proxy with timeout. // SOCKS5 CONNECT through proxy with timeout.
// Uses username/password auth for stream isolation (see connect()). // Uses username/password auth for stream isolation (see connect()).
let socks_result = tokio::time::timeout( let socks_result = tokio::time::timeout(Duration::from_millis(timeout_ms), async {
Duration::from_millis(timeout_ms), match &tor_addr {
async { TorAddr::Onion(host, port) | TorAddr::ClearnetHostname(host, port) => {
match &tor_addr { Socks5Stream::connect_with_password(
TorAddr::Onion(host, port) proxy_addr.as_str(),
| TorAddr::ClearnetHostname(host, port) => { (host.as_str(), *port),
Socks5Stream::connect_with_password( "fips",
proxy_addr.as_str(), (host.as_str(), *port), &isolation_key,
"fips", &isolation_key, )
).await .await
}
TorAddr::Clearnet(socket_addr) => {
Socks5Stream::connect_with_password(
proxy_addr.as_str(), *socket_addr,
"fips", &isolation_key,
).await
}
} }
}, TorAddr::Clearnet(socket_addr) => {
) Socks5Stream::connect_with_password(
proxy_addr.as_str(),
*socket_addr,
"fips",
&isolation_key,
)
.await
}
}
})
.await; .await;
let stream = match socks_result { let stream = match socks_result {
@@ -1400,7 +1402,10 @@ mod tests {
let tor_addr = parse_tor_addr(&addr).unwrap(); let tor_addr = parse_tor_addr(&addr).unwrap();
match tor_addr { match tor_addr {
TorAddr::Clearnet(socket_addr) => { TorAddr::Clearnet(socket_addr) => {
assert_eq!(socket_addr, "192.168.1.1:8080".parse::<SocketAddr>().unwrap()); assert_eq!(
socket_addr,
"192.168.1.1:8080".parse::<SocketAddr>().unwrap()
);
} }
_ => panic!("expected Clearnet variant"), _ => panic!("expected Clearnet variant"),
} }
@@ -1547,9 +1552,9 @@ mod tests {
// Integration tests using MockSocks5Server // Integration tests using MockSocks5Server
// ======================================================================== // ========================================================================
use mock_socks5::MockSocks5Server;
use crate::transport::tcp::TcpTransport;
use crate::config::TcpConfig; use crate::config::TcpConfig;
use crate::transport::tcp::TcpTransport;
use mock_socks5::MockSocks5Server;
/// msg1 wire size: 4 prefix + 4 sender_idx + 106 noise_msg1 = 114 bytes. /// msg1 wire size: 4 prefix + 4 sender_idx + 106 noise_msg1 = 114 bytes.
const MSG1_WIRE_SIZE: usize = 114; const MSG1_WIRE_SIZE: usize = 114;
@@ -1597,13 +1602,10 @@ mod tests {
tor.send_async(&target, &frame).await.unwrap(); tor.send_async(&target, &frame).await.unwrap();
// Receive it on the destination // Receive it on the destination
let received = tokio::time::timeout( let received = tokio::time::timeout(Duration::from_secs(5), dest_rx.recv())
Duration::from_secs(5), .await
dest_rx.recv(), .expect("timeout waiting for packet")
) .expect("channel closed");
.await
.expect("timeout waiting for packet")
.expect("channel closed");
assert_eq!(received.data, frame); assert_eq!(received.data, frame);
@@ -1741,7 +1743,10 @@ mod tests {
#[test] #[test]
fn test_directory_service_config_defaults() { fn test_directory_service_config_defaults() {
let config = DirectoryServiceConfig::default(); let config = DirectoryServiceConfig::default();
assert_eq!(config.hostname_file(), "/var/lib/tor/fips_onion_service/hostname"); assert_eq!(
config.hostname_file(),
"/var/lib/tor/fips_onion_service/hostname"
);
assert_eq!(config.bind_addr(), "127.0.0.1:8443"); assert_eq!(config.bind_addr(), "127.0.0.1:8443");
} }
@@ -1865,5 +1870,4 @@ mod tests {
let err = format!("{}", result.unwrap_err()); let err = format!("{}", result.unwrap_err());
assert!(err.contains("directory")); assert!(err.contains("directory"));
} }
} }

View File

@@ -8,10 +8,10 @@ use super::{
}; };
mod socket; mod socket;
mod stats; mod stats;
use socket::{AsyncUdpSocket, UdpRawSocket};
use stats::UdpStats;
use super::resolve_socket_addr; use super::resolve_socket_addr;
use crate::config::UdpConfig; use crate::config::UdpConfig;
use socket::{AsyncUdpSocket, UdpRawSocket};
use stats::UdpStats;
use std::collections::HashMap; use std::collections::HashMap;
use std::net::SocketAddr; use std::net::SocketAddr;
use std::sync::{Arc, Mutex as StdMutex}; use std::sync::{Arc, Mutex as StdMutex};
@@ -125,7 +125,11 @@ impl UdpTransport {
/// Query transport-local congestion indicators. /// Query transport-local congestion indicators.
pub fn congestion(&self) -> super::TransportCongestion { pub fn congestion(&self) -> super::TransportCongestion {
super::TransportCongestion { super::TransportCongestion {
recv_drops: Some(self.stats.kernel_drops.load(std::sync::atomic::Ordering::Relaxed)), recv_drops: Some(
self.stats
.kernel_drops
.load(std::sync::atomic::Ordering::Relaxed),
),
} }
} }
@@ -367,7 +371,7 @@ async fn udp_receive_loop(
mod tests { mod tests {
use super::*; use super::*;
use crate::transport::packet_channel; use crate::transport::packet_channel;
use tokio::time::{timeout, Duration}; use tokio::time::{Duration, timeout};
fn make_config(port: u16) -> UdpConfig { fn make_config(port: u16) -> UdpConfig {
UdpConfig { UdpConfig {
@@ -444,7 +448,10 @@ mod tests {
.expect("channel closed"); .expect("channel closed");
assert_eq!(packet.data, data); assert_eq!(packet.data, data);
assert_eq!(packet.remote_addr.as_str(), Some(addr1.to_string().as_str())); assert_eq!(
packet.remote_addr.as_str(),
Some(addr1.to_string().as_str())
);
t1.stop_async().await.unwrap(); t1.stop_async().await.unwrap();
t2.stop_async().await.unwrap(); t2.stop_async().await.unwrap();
@@ -615,7 +622,10 @@ mod tests {
// Send using IP string address // Send using IP string address
let data = b"hello via ip string"; let data = b"hello via ip string";
let bytes_sent = t1 let bytes_sent = t1
.send_async(&TransportAddr::from_string(&format!("127.0.0.1:{}", port2)), data) .send_async(
&TransportAddr::from_string(&format!("127.0.0.1:{}", port2)),
data,
)
.await .await
.unwrap(); .unwrap();
assert_eq!(bytes_sent, data.len()); assert_eq!(bytes_sent, data.len());

View File

@@ -178,8 +178,7 @@ impl UdpRawSocket {
unsafe { unsafe {
let mut cmsg = libc::CMSG_FIRSTHDR(&msg); let mut cmsg = libc::CMSG_FIRSTHDR(&msg);
while !cmsg.is_null() { while !cmsg.is_null() {
if (*cmsg).cmsg_level == libc::SOL_SOCKET if (*cmsg).cmsg_level == libc::SOL_SOCKET && (*cmsg).cmsg_type == libc::SO_RXQ_OVFL
&& (*cmsg).cmsg_type == libc::SO_RXQ_OVFL
{ {
let data = libc::CMSG_DATA(cmsg); let data = libc::CMSG_DATA(cmsg);
drops = std::ptr::read_unaligned(data as *const u32); drops = std::ptr::read_unaligned(data as *const u32);
@@ -218,11 +217,7 @@ pub struct AsyncUdpSocket {
impl AsyncUdpSocket { impl AsyncUdpSocket {
/// Send a payload to a destination address. /// Send a payload to a destination address.
pub async fn send_to( pub async fn send_to(&self, data: &[u8], dest: &SocketAddr) -> Result<usize, TransportError> {
&self,
data: &[u8],
dest: &SocketAddr,
) -> Result<usize, TransportError> {
loop { loop {
let mut guard = self let mut guard = self
.inner .inner
@@ -262,9 +257,7 @@ impl AsyncUdpSocket {
} }
/// Convert a `libc::sockaddr_storage` to `std::net::SocketAddr`. /// Convert a `libc::sockaddr_storage` to `std::net::SocketAddr`.
fn sockaddr_to_socket_addr( fn sockaddr_to_socket_addr(storage: &libc::sockaddr_storage) -> std::io::Result<SocketAddr> {
storage: &libc::sockaddr_storage,
) -> std::io::Result<SocketAddr> {
match storage.ss_family as libc::c_int { match storage.ss_family as libc::c_int {
libc::AF_INET => { libc::AF_INET => {
let addr: &libc::sockaddr_in = let addr: &libc::sockaddr_in =

View File

@@ -80,12 +80,7 @@ impl TreeCoordinate {
if addrs.is_empty() { if addrs.is_empty() {
return Err(TreeError::EmptyCoordinate); return Err(TreeError::EmptyCoordinate);
} }
Ok(Self( Ok(Self(addrs.into_iter().map(CoordEntry::addr_only).collect()))
addrs
.into_iter()
.map(CoordEntry::addr_only)
.collect(),
))
} }
/// Create a coordinate for a root node. /// Create a coordinate for a root node.

View File

@@ -1,7 +1,7 @@
//! Parent declarations for the spanning tree. //! Parent declarations for the spanning tree.
use secp256k1::schnorr::Signature;
use secp256k1::XOnlyPublicKey; use secp256k1::XOnlyPublicKey;
use secp256k1::schnorr::Signature;
use std::fmt; use std::fmt;
use super::TreeError; use super::TreeError;

View File

@@ -157,7 +157,8 @@ impl TreeState {
/// Only records a flap when the parent actually changes. /// Only records a flap when the parent actually changes.
pub fn set_parent(&mut self, parent_id: NodeAddr, sequence: u64, timestamp: u64) -> bool { pub fn set_parent(&mut self, parent_id: NodeAddr, sequence: u64, timestamp: u64) -> bool {
let parent_changed = self.is_root() || *self.my_declaration.parent_id() != parent_id; let parent_changed = self.is_root() || *self.my_declaration.parent_id() != parent_id;
self.my_declaration = ParentDeclaration::new(self.my_node_addr, parent_id, sequence, timestamp); self.my_declaration =
ParentDeclaration::new(self.my_node_addr, parent_id, sequence, timestamp);
self.last_parent_switch = Some(Instant::now()); self.last_parent_switch = Some(Instant::now());
// Record switch for flap detection only when parent actually changes; // Record switch for flap detection only when parent actually changes;
// coordinates will be recomputed when ancestry is available // coordinates will be recomputed when ancestry is available
@@ -228,8 +229,7 @@ impl TreeState {
let dominated = match &best { let dominated = match &best {
None => true, None => true,
Some((best_id, best_dist)) => { Some((best_id, best_dist)) => {
distance < *best_dist distance < *best_dist || (distance == *best_dist && peer_id < best_id)
|| (distance == *best_dist && peer_id < best_id)
} }
}; };
@@ -353,9 +353,7 @@ impl TreeState {
match &best_peer { match &best_peer {
None => best_peer = Some((*peer_id, eff_depth)), None => best_peer = Some((*peer_id, eff_depth)),
Some((best_id, best_eff)) => { Some((best_id, best_eff)) => {
if eff_depth < *best_eff if eff_depth < *best_eff || (eff_depth == *best_eff && peer_id < best_id) {
|| (eff_depth == *best_eff && peer_id < best_id)
{
best_peer = Some((*peer_id, eff_depth)); best_peer = Some((*peer_id, eff_depth));
} }
} }
@@ -372,7 +370,11 @@ impl TreeState {
// --- Mandatory switches (bypass hold-down and hysteresis) --- // --- Mandatory switches (bypass hold-down and hysteresis) ---
// If our current parent is gone from peer_ancestry, our path is broken — always switch // If our current parent is gone from peer_ancestry, our path is broken — always switch
if !self.is_root() && !self.peer_ancestry.contains_key(self.my_declaration.parent_id()) { if !self.is_root()
&& !self
.peer_ancestry
.contains_key(self.my_declaration.parent_id())
{
return Some(best_peer_id); return Some(best_peer_id);
} }
@@ -411,7 +413,11 @@ impl TreeState {
let current_parent_cost = peer_costs let current_parent_cost = peer_costs
.get(self.my_declaration.parent_id()) .get(self.my_declaration.parent_id())
.copied() .copied()
.unwrap_or(if peer_costs.is_empty() { 1.0 } else { f64::INFINITY }); .unwrap_or(if peer_costs.is_empty() {
1.0
} else {
f64::INFINITY
});
let current_parent_coords = self.peer_ancestry.get(self.my_declaration.parent_id()); let current_parent_coords = self.peer_ancestry.get(self.my_declaration.parent_id());
let current_parent_eff = match current_parent_coords { let current_parent_eff = match current_parent_coords {
Some(coords) => coords.depth() as f64 + current_parent_cost, Some(coords) => coords.depth() as f64 + current_parent_cost,
@@ -451,8 +457,7 @@ impl TreeState {
.map(|d| d.as_secs()) .map(|d| d.as_secs())
.unwrap_or(0); .unwrap_or(0);
let new_seq = self.my_declaration.sequence() + 1; let new_seq = self.my_declaration.sequence() + 1;
self.my_declaration = self.my_declaration = ParentDeclaration::self_root(self.my_node_addr, new_seq, timestamp);
ParentDeclaration::self_root(self.my_node_addr, new_seq, timestamp);
self.recompute_coords(); self.recompute_coords();
true true
} }

View File

@@ -12,8 +12,8 @@
use crate::upper::hosts::{HostMap, HostMapReloader}; use crate::upper::hosts::{HostMap, HostMapReloader};
use crate::{NodeAddr, PeerIdentity}; use crate::{NodeAddr, PeerIdentity};
use simple_dns::rdata::{RData, AAAA}; use simple_dns::rdata::{AAAA, RData};
use simple_dns::{Packet, Name, ResourceRecord, CLASS, RCODE, PacketFlag, QTYPE, TYPE}; use simple_dns::{CLASS, Name, Packet, PacketFlag, QTYPE, RCODE, ResourceRecord, TYPE};
use std::net::Ipv6Addr; use std::net::Ipv6Addr;
use tracing::{debug, warn}; use tracing::{debug, warn};
@@ -109,16 +109,10 @@ pub fn handle_dns_packet(
let mut response = query.into_reply(); let mut response = query.into_reply();
response.set_flags(PacketFlag::AUTHORITATIVE_ANSWER); response.set_flags(PacketFlag::AUTHORITATIVE_ANSWER);
if is_aaaa if is_aaaa && let Some((ipv6, node_addr, pubkey)) = resolve_fips_query_with_hosts(&qname, hosts)
&& let Some((ipv6, node_addr, pubkey)) = resolve_fips_query_with_hosts(&qname, hosts)
{ {
let name = Name::new_unchecked(&qname).into_owned(); let name = Name::new_unchecked(&qname).into_owned();
let record = ResourceRecord::new( let record = ResourceRecord::new(name, CLASS::IN, ttl, RData::AAAA(AAAA::from(ipv6)));
name,
CLASS::IN,
ttl,
RData::AAAA(AAAA::from(ipv6)),
);
response.answers.push(record); response.answers.push(record);
let identity = DnsResolvedIdentity { node_addr, pubkey }; let identity = DnsResolvedIdentity { node_addr, pubkey };
@@ -359,7 +353,10 @@ mod tests {
assert!(result.is_some(), "should handle hostname AAAA query"); assert!(result.is_some(), "should handle hostname AAAA query");
let (response_bytes, identity_opt) = result.unwrap(); let (response_bytes, identity_opt) = result.unwrap();
assert!(identity_opt.is_some(), "should produce identity for hostname"); assert!(
identity_opt.is_some(),
"should produce identity for hostname"
);
let response = Packet::parse(&response_bytes).unwrap(); let response = Packet::parse(&response_bytes).unwrap();
assert_eq!(response.answers.len(), 1); assert_eq!(response.answers.len(), 1);
@@ -380,7 +377,10 @@ mod tests {
assert!(result.is_some()); assert!(result.is_some());
let (response_bytes, identity_opt) = result.unwrap(); let (response_bytes, identity_opt) = result.unwrap();
assert!(identity_opt.is_none(), "should not produce identity for unknown"); assert!(
identity_opt.is_none(),
"should not produce identity for unknown"
);
let response = Packet::parse(&response_bytes).unwrap(); let response = Packet::parse(&response_bytes).unwrap();
assert_eq!(response.rcode(), RCODE::NameError); assert_eq!(response.rcode(), RCODE::NameError);
@@ -426,12 +426,8 @@ mod tests {
let (identity_tx, mut identity_rx) = tokio::sync::mpsc::channel(16); let (identity_tx, mut identity_rx) = tokio::sync::mpsc::channel(16);
// Spawn the responder // Spawn the responder
let responder_handle = tokio::spawn(run_dns_responder( let responder_handle =
server_socket, tokio::spawn(run_dns_responder(server_socket, identity_tx, 300, reloader));
identity_tx,
300,
reloader,
));
// Send a query // Send a query
let client_socket = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap(); let client_socket = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
@@ -457,13 +453,10 @@ mod tests {
} }
// Verify identity was sent through channel // Verify identity was sent through channel
let resolved = tokio::time::timeout( let resolved = tokio::time::timeout(std::time::Duration::from_secs(1), identity_rx.recv())
std::time::Duration::from_secs(1), .await
identity_rx.recv(), .unwrap()
) .unwrap();
.await
.unwrap()
.unwrap();
assert_eq!(resolved.node_addr, *identity.node_addr()); assert_eq!(resolved.node_addr, *identity.node_addr());
responder_handle.abort(); responder_handle.abort();
@@ -486,12 +479,8 @@ mod tests {
let (identity_tx, mut identity_rx) = tokio::sync::mpsc::channel(16); let (identity_tx, mut identity_rx) = tokio::sync::mpsc::channel(16);
let responder_handle = tokio::spawn(run_dns_responder( let responder_handle =
server_socket, tokio::spawn(run_dns_responder(server_socket, identity_tx, 300, reloader));
identity_tx,
300,
reloader,
));
// Query by hostname instead of npub // Query by hostname instead of npub
let client_socket = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap(); let client_socket = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
@@ -516,13 +505,10 @@ mod tests {
} }
// Verify identity registration // Verify identity registration
let resolved = tokio::time::timeout( let resolved = tokio::time::timeout(std::time::Duration::from_secs(1), identity_rx.recv())
std::time::Duration::from_secs(1), .await
identity_rx.recv(), .unwrap()
) .unwrap();
.await
.unwrap()
.unwrap();
assert_eq!(resolved.node_addr, *identity.node_addr()); assert_eq!(resolved.node_addr, *identity.node_addr());
responder_handle.abort(); responder_handle.abort();
@@ -545,12 +531,8 @@ mod tests {
let server_addr = server_socket.local_addr().unwrap(); let server_addr = server_socket.local_addr().unwrap();
let (identity_tx, _identity_rx) = tokio::sync::mpsc::channel(16); let (identity_tx, _identity_rx) = tokio::sync::mpsc::channel(16);
let responder_handle = tokio::spawn(run_dns_responder( let responder_handle =
server_socket, tokio::spawn(run_dns_responder(server_socket, identity_tx, 300, reloader));
identity_tx,
300,
reloader,
));
let client_socket = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap(); let client_socket = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
@@ -561,16 +543,23 @@ mod tests {
let (len, _) = tokio::time::timeout( let (len, _) = tokio::time::timeout(
std::time::Duration::from_secs(2), std::time::Duration::from_secs(2),
client_socket.recv_from(&mut buf), client_socket.recv_from(&mut buf),
).await.unwrap().unwrap(); )
.await
.unwrap()
.unwrap();
let response = Packet::parse(&buf[..len]).unwrap(); let response = Packet::parse(&buf[..len]).unwrap();
assert!(response.answers.is_empty(), "server2 should not resolve before reload"); assert!(
response.answers.is_empty(),
"server2 should not resolve before reload"
);
// Update the hosts file to add server2 // Update the hosts file to add server2
std::thread::sleep(std::time::Duration::from_millis(50)); std::thread::sleep(std::time::Duration::from_millis(50));
std::fs::write( std::fs::write(
&hosts_path, &hosts_path,
format!("gateway {}\nserver2 {}\n", id1.npub(), id2.npub()), format!("gateway {}\nserver2 {}\n", id1.npub(), id2.npub()),
).unwrap(); )
.unwrap();
// Next query should trigger reload — query server2 again // Next query should trigger reload — query server2 again
let query = build_test_query("server2.fips", TYPE::AAAA); let query = build_test_query("server2.fips", TYPE::AAAA);
@@ -578,9 +567,16 @@ mod tests {
let (len, _) = tokio::time::timeout( let (len, _) = tokio::time::timeout(
std::time::Duration::from_secs(2), std::time::Duration::from_secs(2),
client_socket.recv_from(&mut buf), client_socket.recv_from(&mut buf),
).await.unwrap().unwrap(); )
.await
.unwrap()
.unwrap();
let response = Packet::parse(&buf[..len]).unwrap(); let response = Packet::parse(&buf[..len]).unwrap();
assert_eq!(response.answers.len(), 1, "server2 should resolve after reload"); assert_eq!(
response.answers.len(),
1,
"server2 should resolve after reload"
);
if let RData::AAAA(aaaa) = &response.answers[0].rdata { if let RData::AAAA(aaaa) = &response.answers[0].rdata {
assert_eq!(Ipv6Addr::from(aaaa.address), expected_ipv6_2); assert_eq!(Ipv6Addr::from(aaaa.address), expected_ipv6_2);
} else { } else {

View File

@@ -80,7 +80,9 @@ impl HostMap {
/// Look up the npub for a hostname (case-insensitive). /// Look up the npub for a hostname (case-insensitive).
pub fn lookup_npub(&self, hostname: &str) -> Option<&str> { pub fn lookup_npub(&self, hostname: &str) -> Option<&str> {
self.by_name.get(&hostname.to_ascii_lowercase()).map(|s| s.as_str()) self.by_name
.get(&hostname.to_ascii_lowercase())
.map(|s| s.as_str())
} }
/// Look up the hostname for a NodeAddr (reverse lookup for display). /// Look up the hostname for a NodeAddr (reverse lookup for display).
@@ -279,7 +281,9 @@ pub fn validate_hostname(hostname: &str) -> Result<(), HostMapError> {
} }
if hostname.to_ascii_lowercase().starts_with("npub1") { if hostname.to_ascii_lowercase().starts_with("npub1") {
return Err(err("must not start with 'npub1' (ambiguous with npub resolution)")); return Err(err(
"must not start with 'npub1' (ambiguous with npub resolution)",
));
} }
if hostname.starts_with('-') { if hostname.starts_with('-') {
@@ -338,7 +342,10 @@ mod tests {
("NPUB1bar", "npub1 prefix case"), ("NPUB1bar", "npub1 prefix case"),
]; ];
for (h, desc) in cases { for (h, desc) in cases {
assert!(validate_hostname(h).is_err(), "should be invalid ({desc}): {h}"); assert!(
validate_hostname(h).is_err(),
"should be invalid ({desc}): {h}"
);
} }
} }
@@ -574,10 +581,7 @@ mod tests {
let mut base = HostMap::new(); let mut base = HostMap::new();
base.insert("core", &id.npub()).unwrap(); base.insert("core", &id.npub()).unwrap();
let reloader = HostMapReloader::new( let reloader = HostMapReloader::new(base, std::path::PathBuf::from("/nonexistent/hosts"));
base,
std::path::PathBuf::from("/nonexistent/hosts"),
);
// Only base entries present // Only base entries present
assert_eq!(reloader.hosts().len(), 1); assert_eq!(reloader.hosts().len(), 1);
assert!(reloader.hosts().lookup_npub("core").is_some()); assert!(reloader.hosts().lookup_npub("core").is_some());
@@ -605,7 +609,11 @@ mod tests {
// Modify the file — bump mtime by writing new content // Modify the file — bump mtime by writing new content
// Sleep briefly to ensure mtime changes (filesystem granularity) // Sleep briefly to ensure mtime changes (filesystem granularity)
std::thread::sleep(std::time::Duration::from_millis(50)); std::thread::sleep(std::time::Duration::from_millis(50));
std::fs::write(&path, format!("gateway {}\nnew-host {}\n", id1.npub(), id2.npub())).unwrap(); std::fs::write(
&path,
format!("gateway {}\nnew-host {}\n", id1.npub(), id2.npub()),
)
.unwrap();
assert!(reloader.check_reload()); assert!(reloader.check_reload());
assert_eq!(reloader.hosts().len(), 2); assert_eq!(reloader.hosts().len(), 2);

View File

@@ -543,7 +543,8 @@ mod tests {
let short_packet = vec![0u8; 20]; let short_packet = vec![0u8; 20];
let our_addr: Ipv6Addr = "fd00::ffff".parse().unwrap(); let our_addr: Ipv6Addr = "fd00::ffff".parse().unwrap();
let response = build_dest_unreachable(&short_packet, DestUnreachableCode::NoRoute, our_addr); let response =
build_dest_unreachable(&short_packet, DestUnreachableCode::NoRoute, our_addr);
assert!(response.is_none()); assert!(response.is_none());
} }

View File

@@ -68,9 +68,8 @@ impl IcmpRateLimiter {
/// Remove entries older than max_age. /// Remove entries older than max_age.
fn cleanup(&mut self, now: Instant) { fn cleanup(&mut self, now: Instant) {
self.last_sent.retain(|_, &mut last| { self.last_sent
now.duration_since(last) < self.max_age .retain(|_, &mut last| now.duration_since(last) < self.max_age);
});
} }
/// Get the number of tracked sources. /// Get the number of tracked sources.
@@ -179,4 +178,4 @@ mod tests {
limiter.cleanup(Instant::now()); limiter.cleanup(Instant::now());
assert_eq!(limiter.len(), 1); assert_eq!(limiter.len(), 1);
} }
} }

View File

@@ -156,13 +156,17 @@ mod tests {
} }
fn sample_src() -> [u8; 16] { fn sample_src() -> [u8; 16] {
[0xfd, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, [
0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f] 0xfd, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d,
0x0e, 0x0f,
]
} }
fn sample_dst() -> [u8; 16] { fn sample_dst() -> [u8; 16] {
[0xfd, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, [
0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f] 0xfd, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d,
0x1e, 0x1f,
]
} }
// ===== Round-trip fidelity ===== // ===== Round-trip fidelity =====
@@ -334,10 +338,14 @@ mod tests {
#[test] #[test]
fn test_addresses_from_context() { fn test_addresses_from_context() {
let original_src = [0xfd, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, let original_src = [
0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA]; 0xfd, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA, 0xAA,
let original_dst = [0xfd, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xAA, 0xAA,
0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB]; ];
let original_dst = [
0xfd, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB, 0xBB,
0xBB, 0xBB,
];
let pkt = build_ipv6_packet(0, 0, 17, 64, original_src, original_dst, &[1, 2]); let pkt = build_ipv6_packet(0, 0, 17, 64, original_src, original_dst, &[1, 2]);
let compressed = compress_ipv6(&pkt).unwrap(); let compressed = compress_ipv6(&pkt).unwrap();

View File

@@ -74,7 +74,7 @@ pub fn clamp_tcp_mss(ipv6_packet: &mut [u8], max_mss: u16) -> bool {
// Parse TCP options // Parse TCP options
let options_start = tcp_start + TCP_HEADER_MIN_LEN; let options_start = tcp_start + TCP_HEADER_MIN_LEN;
let options_end = tcp_start + tcp_header_len; let options_end = tcp_start + tcp_header_len;
if options_end > ipv6_packet.len() { if options_end > ipv6_packet.len() {
return false; return false;
} }
@@ -114,10 +114,10 @@ pub fn clamp_tcp_mss(ipv6_packet: &mut [u8], max_mss: u16) -> bool {
// Clamp if needed // Clamp if needed
if current_mss > max_mss { if current_mss > max_mss {
ipv6_packet[i + 2..i + 4].copy_from_slice(&max_mss.to_be_bytes()); ipv6_packet[i + 2..i + 4].copy_from_slice(&max_mss.to_be_bytes());
// Recalculate TCP checksum // Recalculate TCP checksum
recalculate_tcp_checksum(ipv6_packet, tcp_start); recalculate_tcp_checksum(ipv6_packet, tcp_start);
modified = true; modified = true;
} }
break; // MSS option found, no need to continue break; // MSS option found, no need to continue
@@ -258,7 +258,7 @@ mod tests {
let src = [0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1]; let src = [0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1];
let dst = [0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2]; let dst = [0xfd, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 2];
let mut packet = make_tcp_syn_packet(src, dst, 1460); let mut packet = make_tcp_syn_packet(src, dst, 1460);
// Clear SYN flag // Clear SYN flag
packet[40 + 13] = 0x10; // ACK only packet[40 + 13] = 0x10; // ACK only
@@ -277,4 +277,4 @@ mod tests {
assert!(!modified); assert!(!modified);
} }
} }

View File

@@ -6,7 +6,7 @@
use crate::{FipsAddress, TunConfig}; use crate::{FipsAddress, TunConfig};
use futures::TryStreamExt; use futures::TryStreamExt;
use rtnetlink::{new_connection, Handle, LinkUnspec, RouteMessageBuilder}; use rtnetlink::{Handle, LinkUnspec, RouteMessageBuilder, new_connection};
use std::fs::File; use std::fs::File;
use std::io::{Read, Write}; use std::io::{Read, Write};
use std::net::Ipv6Addr; use std::net::Ipv6Addr;
@@ -151,7 +151,9 @@ impl TunDevice {
/// Returns the number of bytes read into the buffer, or an error. /// Returns the number of bytes read into the buffer, or an error.
/// The buffer should be at least MTU + header size (typically 1500+ bytes). /// The buffer should be at least MTU + header size (typically 1500+ bytes).
pub fn read_packet(&mut self, buf: &mut [u8]) -> Result<usize, TunError> { pub fn read_packet(&mut self, buf: &mut [u8]) -> Result<usize, TunError> {
self.device.read(buf).map_err(|e| TunError::Configure(format!("read failed: {}", e))) self.device
.read(buf)
.map_err(|e| TunError::Configure(format!("read failed: {}", e)))
} }
/// Shutdown and delete the TUN device. /// Shutdown and delete the TUN device.
@@ -263,7 +265,9 @@ pub fn run_tun_reader(
outbound_tx: TunOutboundTx, outbound_tx: TunOutboundTx,
transport_mtu: u16, transport_mtu: u16,
) { ) {
use super::icmp::{build_dest_unreachable, effective_ipv6_mtu, should_send_icmp_error, DestUnreachableCode}; use super::icmp::{
DestUnreachableCode, build_dest_unreachable, effective_ipv6_mtu, should_send_icmp_error,
};
use super::tcp_mss::clamp_tcp_mss; use super::tcp_mss::clamp_tcp_mss;
let name = device.name().to_string(); let name = device.name().to_string();
@@ -470,7 +474,13 @@ async fn configure_interface(name: &str, addr: Ipv6Addr, mtu: u16) -> Result<(),
// Add ip6 rule to ensure fd00::/8 uses the main table, preventing other // Add ip6 rule to ensure fd00::/8 uses the main table, preventing other
// routing software (e.g. Tailscale) from intercepting FIPS traffic via // routing software (e.g. Tailscale) from intercepting FIPS traffic via
// catch-all rules in auxiliary routing tables. // catch-all rules in auxiliary routing tables.
let mut rule_req = handle.rule().add().v6().destination_prefix(fd_prefix, 8).table_id(254).priority(5265); let mut rule_req = handle
.rule()
.add()
.v6()
.destination_prefix(fd_prefix, 8)
.table_id(254)
.priority(5265);
rule_req.message_mut().header.action = 1.into(); // FR_ACT_TO_TBL rule_req.message_mut().header.action = 1.into(); // FR_ACT_TO_TBL
if let Err(e) = rule_req.execute().await { if let Err(e) = rule_req.execute().await {
debug!("ip6 rule for fd00::/8 not added (may already exist): {e}"); debug!("ip6 rule for fd00::/8 not added (may already exist): {e}");

Some files were not shown because too many files have changed in this diff Show More