diff --git a/src/node/handlers/discovery.rs b/src/node/handlers/discovery.rs new file mode 100644 index 0000000..138e803 --- /dev/null +++ b/src/node/handlers/discovery.rs @@ -0,0 +1,299 @@ +//! LookupRequest/LookupResponse discovery protocol handlers. +//! +//! Handles coordinate discovery requests: flood-based lookup with TTL, +//! visited filter for loop prevention, and reverse-path forwarding for +//! responses. + +use crate::node::{Node, RecentRequest}; +use crate::protocol::{LookupRequest, LookupResponse}; +use crate::NodeAddr; +use tracing::{debug, trace}; + +impl Node { + /// Handle an incoming LookupRequest from a peer. + /// + /// Processing steps: + /// 1. Decode and validate + /// 2. Check request_id for duplicates (dedup) + /// 3. Record request for reverse-path forwarding + /// 4. Lazy purge expired entries + /// 5. Check visited filter (loop prevention) + /// 6. If we're the target, generate and send response + /// 7. If TTL > 0, forward to peers not in visited filter + pub(in crate::node) async fn handle_lookup_request( + &mut self, + from: &NodeAddr, + payload: &[u8], + ) { + let request = match LookupRequest::decode(payload) { + Ok(req) => req, + Err(e) => { + debug!(from = %from, error = %e, "Malformed LookupRequest"); + return; + } + }; + + let now_ms = Self::now_ms(); + + // Dedup: drop if we've already seen this request_id + if self.recent_requests.contains_key(&request.request_id) { + trace!( + request_id = request.request_id, + from = %from, + "Duplicate LookupRequest, dropping" + ); + return; + } + + // Record for reverse-path forwarding and dedup + self.recent_requests.insert( + request.request_id, + RecentRequest::new(*from, now_ms), + ); + + // Lazy purge expired entries + self.purge_expired_requests(now_ms); + + // Loop prevention: drop if we've already been visited + if request.was_visited(self.node_addr()) { + trace!( + request_id = request.request_id, + target = %request.target, + "Already visited, dropping LookupRequest" + ); + return; + } + + // Are we the target? + if request.target == *self.node_addr() { + debug!( + request_id = request.request_id, + origin = %request.origin, + "We are the lookup target, generating response" + ); + self.send_lookup_response(&request).await; + return; + } + + // Forward if TTL permits + if request.can_forward() { + self.forward_lookup_request(request).await; + } else { + trace!( + request_id = request.request_id, + target = %request.target, + "LookupRequest TTL exhausted, not forwarding" + ); + } + } + + /// Handle an incoming LookupResponse from a peer. + /// + /// Processing steps: + /// 1. Decode and validate + /// 2. Check recent_requests to determine if we originated or are forwarding + /// 3. If originator: cache target_coords in route_cache + /// 4. If transit: reverse-path forward to from_peer + pub(in crate::node) async fn handle_lookup_response( + &mut self, + from: &NodeAddr, + payload: &[u8], + ) { + let response = match LookupResponse::decode(payload) { + Ok(resp) => resp, + Err(e) => { + debug!(from = %from, error = %e, "Malformed LookupResponse"); + return; + } + }; + + let now_ms = Self::now_ms(); + + // Check if we forwarded this request (transit node) or originated it + if let Some(recent) = self.recent_requests.get(&response.request_id) { + // Transit node: reverse-path forward + let from_peer = recent.from_peer; + + debug!( + request_id = response.request_id, + target = %response.target, + next_hop = %from_peer, + "Reverse-path forwarding LookupResponse" + ); + + let encoded = response.encode(); + if let Err(e) = self.send_encrypted_link_message(&from_peer, &encoded).await { + debug!( + next_hop = %from_peer, + error = %e, + "Failed to forward LookupResponse" + ); + } + } else { + // We originated this request — cache the discovered coordinates + debug!( + request_id = response.request_id, + target = %response.target, + depth = response.target_coords.depth(), + "Received LookupResponse, caching route" + ); + + self.route_cache.insert( + response.target, + response.target_coords, + now_ms, + ); + } + } + + /// Generate and send a LookupResponse when we are the target. + /// + /// Signs a proof using our identity and routes the response toward + /// the origin. The first hop uses find_next_hop; subsequent hops use + /// reverse-path forwarding via recent_requests. + async fn send_lookup_response(&mut self, request: &LookupRequest) { + let our_coords = self.tree_state().my_coords().clone(); + + // Sign proof: Identity::sign hashes with SHA-256 internally + let proof_data = LookupResponse::proof_bytes(request.request_id, &request.target); + let proof = self.identity().sign(&proof_data); + + let response = LookupResponse::new( + request.request_id, + request.target, + our_coords, + proof, + ); + + // Route toward origin + let next_hop_addr = match self.find_next_hop(&request.origin) { + Some(peer) => *peer.node_addr(), + None => { + // Origin might be our direct peer that sent us the request + // Check if origin == the peer we received from + if let Some(recent) = self.recent_requests.get(&request.request_id) { + recent.from_peer + } else { + debug!( + origin = %request.origin, + "Cannot route LookupResponse: no path to origin" + ); + return; + } + } + }; + + debug!( + request_id = request.request_id, + origin = %request.origin, + next_hop = %next_hop_addr, + "Sending LookupResponse" + ); + + let encoded = response.encode(); + if let Err(e) = self.send_encrypted_link_message(&next_hop_addr, &encoded).await { + debug!( + next_hop = %next_hop_addr, + error = %e, + "Failed to send LookupResponse" + ); + } + } + + /// Forward a LookupRequest to peers not in the visited filter. + /// + /// Decrements TTL, adds self to visited, and sends to all eligible peers. + async fn forward_lookup_request(&mut self, mut request: LookupRequest) { + if !request.forward(self.node_addr()) { + return; + } + + // Collect peers not in visited filter + let forward_to: Vec = self + .peers + .keys() + .filter(|addr| !request.was_visited(addr)) + .copied() + .collect(); + + if forward_to.is_empty() { + trace!( + request_id = request.request_id, + "No eligible peers to forward LookupRequest" + ); + return; + } + + debug!( + request_id = request.request_id, + target = %request.target, + ttl = request.ttl, + peer_count = forward_to.len(), + "Forwarding LookupRequest" + ); + + let encoded = request.encode(); + + for peer_addr in forward_to { + if let Err(e) = self.send_encrypted_link_message(&peer_addr, &encoded).await { + debug!( + peer = %peer_addr, + error = %e, + "Failed to forward LookupRequest to peer" + ); + } + } + } + + /// Initiate a discovery lookup for a target node. + /// + /// Creates a LookupRequest and floods it to all peers. The originator + /// does NOT record the request_id in recent_requests, so when the + /// response arrives, it's recognized as "our request" and the + /// target's coordinates are cached in route_cache. + #[allow(dead_code)] // Called from integration tests; will be used from event loop + pub(in crate::node) async fn initiate_lookup(&mut self, target: &NodeAddr, ttl: u8) { + let origin = *self.node_addr(); + let origin_coords = self.tree_state().my_coords().clone(); + let mut request = LookupRequest::generate(*target, origin, origin_coords, ttl); + + // Add ourselves to the visited filter so forwarding nodes + // won't send the request back to us + request.visited.insert(&origin); + + debug!( + request_id = request.request_id, + target = %target, + ttl = ttl, + "Initiating LookupRequest" + ); + + // Send to all peers (flood) + let peer_addrs: Vec = self.peers.keys().copied().collect(); + let encoded = request.encode(); + + for peer_addr in peer_addrs { + if let Err(e) = self.send_encrypted_link_message(&peer_addr, &encoded).await { + debug!( + peer = %peer_addr, + error = %e, + "Failed to send LookupRequest to peer" + ); + } + } + } + + /// Remove expired entries from the recent_requests cache. + fn purge_expired_requests(&mut self, current_time_ms: u64) { + self.recent_requests + .retain(|_, entry| !entry.is_expired(current_time_ms)); + } + + /// Get current time in milliseconds since Unix epoch. + fn now_ms() -> u64 { + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_millis() as u64) + .unwrap_or(0) + } +} diff --git a/src/node/handlers/dispatch.rs b/src/node/handlers/dispatch.rs index b79ee05..a1397bf 100644 --- a/src/node/handlers/dispatch.rs +++ b/src/node/handlers/dispatch.rs @@ -16,7 +16,6 @@ impl Node { let msg_type = plaintext[0]; let payload = &plaintext[1..]; - // TODO: Implement remaining link message handlers match msg_type { 0x10 => { // TreeAnnounce @@ -28,11 +27,11 @@ impl Node { } 0x30 => { // LookupRequest - debug!("Received LookupRequest (not yet implemented)"); + self.handle_lookup_request(from, payload).await; } 0x31 => { // LookupResponse - debug!("Received LookupResponse (not yet implemented)"); + self.handle_lookup_response(from, payload).await; } 0x40 => { // SessionDatagram diff --git a/src/node/handlers/mod.rs b/src/node/handlers/mod.rs index c03d61d..de6338f 100644 --- a/src/node/handlers/mod.rs +++ b/src/node/handlers/mod.rs @@ -1,5 +1,6 @@ //! RX event loop and message handlers. +mod discovery; mod dispatch; mod encrypted; mod forwarding; diff --git a/src/node/mod.rs b/src/node/mod.rs index 10f0012..ab8db98 100644 --- a/src/node/mod.rs +++ b/src/node/mod.rs @@ -13,7 +13,7 @@ mod tree; mod tests; use crate::bloom::BloomState; -use crate::cache::CoordCache; +use crate::cache::{CoordCache, RouteCache}; use crate::index::IndexAllocator; use crate::peer::{ActivePeer, PeerConnection}; use crate::rate_limit::HandshakeRateLimiter; @@ -142,6 +142,33 @@ impl fmt::Display for NodeState { } } +/// Recent request tracking for dedup and reverse-path forwarding. +/// +/// When a LookupRequest is forwarded through a node, the node stores the +/// request_id and which peer sent it. When the corresponding LookupResponse +/// arrives, it's forwarded back to that peer (reverse-path forwarding). +#[derive(Clone, Debug)] +pub(crate) struct RecentRequest { + /// The peer who sent this request to us. + pub(crate) from_peer: NodeAddr, + /// When we received this request (Unix milliseconds). + pub(crate) timestamp_ms: u64, +} + +impl RecentRequest { + pub(crate) fn new(from_peer: NodeAddr, timestamp_ms: u64) -> Self { + Self { + from_peer, + timestamp_ms, + } + } + + /// Check if this entry has expired (older than 10 seconds). + pub(crate) fn is_expired(&self, current_time_ms: u64) -> bool { + current_time_ms.saturating_sub(self.timestamp_ms) > 10_000 + } +} + /// Key for addr_to_link reverse lookup. type AddrKey = (TransportId, TransportAddr); @@ -182,8 +209,13 @@ pub struct Node { bloom_state: BloomState, // === Routing === - /// Address -> coordinates cache. + /// Address -> coordinates cache (from session setup). coord_cache: CoordCache, + /// Discovered routes (from discovery protocol). + route_cache: RouteCache, + /// Recent discovery requests (dedup + reverse-path forwarding). + /// Maps request_id → RecentRequest. + recent_requests: HashMap, // === Transports & Links === /// Active transports (owned by Node). @@ -294,6 +326,8 @@ impl Node { tree_state, bloom_state, coord_cache: CoordCache::with_defaults(), + route_cache: RouteCache::with_defaults(), + recent_requests: HashMap::new(), transports: HashMap::new(), links: HashMap::new(), addr_to_link: HashMap::new(), @@ -343,6 +377,8 @@ impl Node { tree_state, bloom_state: BloomState::new(node_addr), coord_cache: CoordCache::with_defaults(), + route_cache: RouteCache::with_defaults(), + recent_requests: HashMap::new(), transports: HashMap::new(), links: HashMap::new(), addr_to_link: HashMap::new(), @@ -492,6 +528,18 @@ impl Node { &mut self.coord_cache } + // === Route Cache === + + /// Get the route cache (discovery protocol). + pub fn route_cache(&self) -> &RouteCache { + &self.route_cache + } + + /// Get mutable route cache. + pub fn route_cache_mut(&mut self) -> &mut RouteCache { + &mut self.route_cache + } + // === TUN Interface === /// Get the TUN state. diff --git a/src/node/tests/discovery.rs b/src/node/tests/discovery.rs new file mode 100644 index 0000000..590c9f8 --- /dev/null +++ b/src/node/tests/discovery.rs @@ -0,0 +1,361 @@ +//! Discovery protocol tests: LookupRequest and LookupResponse. +//! +//! Unit tests for handler logic (dedup, visited filter, TTL, response +//! caching) and integration tests for multi-node forwarding and +//! reverse-path response routing. + +use super::*; +use crate::node::RecentRequest; +use crate::protocol::{LookupRequest, LookupResponse}; +use crate::tree::TreeCoordinate; +use spanning_tree::{cleanup_nodes, process_available_packets, run_tree_test}; + +// ============================================================================ +// Unit Tests — LookupRequest Handler +// ============================================================================ + +#[tokio::test] +async fn test_request_decode_error() { + let mut node = make_node(); + let from = make_node_addr(0xAA); + // Too-short payload: should log error and return without panic + node.handle_lookup_request(&from, &[0x00; 5]).await; + assert!(node.recent_requests.is_empty()); +} + +#[tokio::test] +async fn test_request_dedup() { + let mut node = make_node(); + let from = make_node_addr(0xAA); + let target = make_node_addr(0xBB); + let origin = make_node_addr(0xCC); + let coords = TreeCoordinate::from_addrs(vec![origin, make_node_addr(0)]).unwrap(); + + let request = LookupRequest::new(999, target, origin, coords, 5); + let payload = &request.encode()[1..]; // skip msg_type byte + + // First request: accepted + node.handle_lookup_request(&from, payload).await; + assert_eq!(node.recent_requests.len(), 1); + + // Duplicate request: dropped + node.handle_lookup_request(&from, payload).await; + assert_eq!(node.recent_requests.len(), 1); +} + +#[tokio::test] +async fn test_request_visited_filter_self() { + let mut node = make_node(); + let from = make_node_addr(0xAA); + let target = make_node_addr(0xBB); + let origin = make_node_addr(0xCC); + let coords = TreeCoordinate::from_addrs(vec![origin, make_node_addr(0)]).unwrap(); + + let mut request = LookupRequest::new(888, target, origin, coords, 5); + // Mark ourselves as already visited + request.visited.insert(node.node_addr()); + + let payload = &request.encode()[1..]; + node.handle_lookup_request(&from, payload).await; + + // Request was recorded (dedup happens before visited check) + // but the handler should have stopped after detecting self in visited filter + assert!(node.recent_requests.contains_key(&888)); +} + +#[tokio::test] +async fn test_request_target_is_self() { + let mut node = make_node(); + let from = make_node_addr(0xAA); + let origin = make_node_addr(0xCC); + let my_addr = *node.node_addr(); + let coords = TreeCoordinate::from_addrs(vec![origin, make_node_addr(0)]).unwrap(); + + // Request targeting us + let request = LookupRequest::new(777, my_addr, origin, coords, 5); + let payload = &request.encode()[1..]; + + // Should succeed without panic (response send will fail silently + // since we have no peers to route toward origin) + node.handle_lookup_request(&from, payload).await; + assert!(node.recent_requests.contains_key(&777)); +} + +#[tokio::test] +async fn test_request_ttl_zero_not_forwarded() { + let mut node = make_node(); + let from = make_node_addr(0xAA); + let target = make_node_addr(0xBB); + let origin = make_node_addr(0xCC); + let coords = TreeCoordinate::from_addrs(vec![origin, make_node_addr(0)]).unwrap(); + + let request = LookupRequest::new(666, target, origin, coords, 0); + let payload = &request.encode()[1..]; + + node.handle_lookup_request(&from, payload).await; + // Request recorded, but not forwarded (TTL=0, and no peers anyway) + assert!(node.recent_requests.contains_key(&666)); +} + +// ============================================================================ +// Unit Tests — LookupResponse Handler +// ============================================================================ + +#[tokio::test] +async fn test_response_decode_error() { + let mut node = make_node(); + let from = make_node_addr(0xAA); + node.handle_lookup_response(&from, &[0x00; 10]).await; + // No panic, no route cached + assert!(node.route_cache.is_empty()); +} + +#[tokio::test] +async fn test_response_originator_caches_route() { + let mut node = make_node(); + let from = make_node_addr(0xAA); + let target = make_node_addr(0xBB); + let root = make_node_addr(0xF0); + let coords = TreeCoordinate::from_addrs(vec![target, root]).unwrap(); + + // Create a valid response with a real proof signature + let proof_data = LookupResponse::proof_bytes(555, &target); + let target_identity = Identity::generate(); + let proof = target_identity.sign(&proof_data); + + let response = LookupResponse::new(555, target, coords.clone(), proof); + let payload = &response.encode()[1..]; // skip msg_type + + // No entry in recent_requests for 555 → we're the originator + assert!(!node.recent_requests.contains_key(&555)); + + node.handle_lookup_response(&from, payload).await; + + // Route should be cached + assert!(node.route_cache.contains(&target)); + let cached = node.route_cache.get(&target).unwrap(); + assert_eq!(cached.coords(), &coords); +} + +#[tokio::test] +async fn test_response_transit_needs_recent_request() { + let mut node = make_node(); + let from = make_node_addr(0xAA); + let target = make_node_addr(0xBB); + let root = make_node_addr(0xF0); + let coords = TreeCoordinate::from_addrs(vec![target, root]).unwrap(); + + let proof_data = LookupResponse::proof_bytes(444, &target); + let target_identity = Identity::generate(); + let proof = target_identity.sign(&proof_data); + + let response = LookupResponse::new(444, target, coords, proof); + let payload = &response.encode()[1..]; + + // Simulate being a transit node: record a recent_request for this ID + let now_ms = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_millis() as u64; + node.recent_requests.insert( + 444, + RecentRequest::new(make_node_addr(0xDD), now_ms), + ); + + // Handle response — should try to reverse-path forward to 0xDD + // (will fail silently since 0xDD is not an actual peer) + node.handle_lookup_response(&from, payload).await; + + // Should NOT cache in route_cache (we're transit, not originator) + assert!(!node.route_cache.contains(&target)); +} + +// ============================================================================ +// Unit Tests — RecentRequest Expiry +// ============================================================================ + +#[tokio::test] +async fn test_recent_request_expiry() { + let mut node = make_node(); + + let now_ms = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_millis() as u64; + + // Insert an old request (11 seconds ago) + node.recent_requests.insert( + 123, + RecentRequest::new(make_node_addr(1), now_ms - 11_000), + ); + + // Insert a recent request + node.recent_requests.insert( + 456, + RecentRequest::new(make_node_addr(2), now_ms), + ); + + assert_eq!(node.recent_requests.len(), 2); + + // Trigger purge via a new lookup request + let target = make_node_addr(0xBB); + let origin = make_node_addr(0xCC); + let coords = TreeCoordinate::from_addrs(vec![origin, make_node_addr(0)]).unwrap(); + let request = LookupRequest::new(789, target, origin, coords, 3); + let payload = &request.encode()[1..]; + node.handle_lookup_request(&make_node_addr(0xAA), payload).await; + + // 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(&456)); + assert!(node.recent_requests.contains_key(&789)); +} + +// ============================================================================ +// Integration Tests — Multi-Node Forwarding +// ============================================================================ + +#[tokio::test] +async fn test_request_forwarding_two_node() { + // Set up a two-node topology: node0 — node1 + // Send a LookupRequest from node0 targeting some unknown node. + // Node1 should receive the forwarded request. + let edges = vec![(0, 1)]; + let mut nodes = run_tree_test(2, &edges, false).await; + + let node0_addr = *nodes[0].node.node_addr(); + let target = make_node_addr(0xEE); // unknown node + let root = make_node_addr(0); + + let coords = TreeCoordinate::from_addrs(vec![node0_addr, root]).unwrap(); + let request = LookupRequest::new(42, target, node0_addr, coords, 5); + let payload = &request.encode()[1..]; + + // Handle on node0 as if we received it from outside + nodes[0] + .node + .handle_lookup_request(&node0_addr, payload) + .await; + + // Process packets — node1 should receive the forwarded request + tokio::time::sleep(Duration::from_millis(50)).await; + let count = process_available_packets(&mut nodes).await; + assert!(count > 0, "Expected forwarded LookupRequest to arrive at node 1"); + + // Node1 should have recorded the request + assert!( + nodes[1].node.recent_requests.contains_key(&42), + "Node 1 should have recorded the forwarded request" + ); + + cleanup_nodes(&mut nodes).await; +} + +#[tokio::test] +async fn test_request_target_found_generates_response() { + // Set up a two-node topology: node0 — node1 + // Node0 initiates a lookup targeting node1. + // Node1 receives, detects it's the target, generates a LookupResponse. + // Response routes back to node0 which caches the coordinates. + let edges = vec![(0, 1)]; + let mut nodes = run_tree_test(2, &edges, false).await; + + let node1_addr = *nodes[1].node.node_addr(); + + // Node0 initiates lookup (doesn't record in recent_requests) + nodes[0].node.initiate_lookup(&node1_addr, 5).await; + + // Process packets in rounds to allow request + response + for _ in 0..4 { + tokio::time::sleep(Duration::from_millis(50)).await; + process_available_packets(&mut nodes).await; + } + + // Node0 should have cached node1's route (it originated the request) + assert!( + nodes[0].node.route_cache.contains(&node1_addr), + "Node 0 should have cached node 1's route from LookupResponse" + ); + + cleanup_nodes(&mut nodes).await; +} + +#[tokio::test] +async fn test_request_three_node_chain() { + // Topology: node0 — node1 — node2 + // Node0 initiates a lookup targeting node2. + // Request should propagate: node0 → node1 → node2. + // Node2 generates response, reverse-path: node2 → node1 → node0. + let edges = vec![(0, 1), (1, 2)]; + let mut nodes = run_tree_test(3, &edges, false).await; + + let node2_addr = *nodes[2].node.node_addr(); + + // Node0 initiates lookup (doesn't record in recent_requests) + nodes[0].node.initiate_lookup(&node2_addr, 8).await; + + // Process packets in rounds to allow multi-hop propagation + response + // Chain: node0→node1→node2 (request), node2→node1→node0 (response) + for _ in 0..10 { + tokio::time::sleep(Duration::from_millis(100)).await; + process_available_packets(&mut nodes).await; + } + + // Node1 should have been a transit node (has the request_id in recent_requests) + assert!( + !nodes[1].node.recent_requests.is_empty(), + "Node 1 should have recorded the forwarded request" + ); + + // Node2 should have received the request (it's the target) + assert!( + !nodes[2].node.recent_requests.is_empty(), + "Node 2 should have received the request" + ); + + // Node0 should have cached node2's route + assert!( + nodes[0].node.route_cache.contains(&node2_addr), + "Node 0 should have cached node 2's route through 3-node chain" + ); + + cleanup_nodes(&mut nodes).await; +} + +#[tokio::test] +async fn test_request_dedup_convergent_paths() { + // Topology: triangle (node0 — node1, node0 — node2, node1 — node2) + // A request from node0 reaches node2 via two paths: 0→1→2 and 0→2. + // The second arrival at node2 should be deduped. + let edges = vec![(0, 1), (0, 2), (1, 2)]; + let mut nodes = run_tree_test(3, &edges, false).await; + + let node0_addr = *nodes[0].node.node_addr(); + let target = make_node_addr(0xEE); + let root = make_node_addr(0); + + let coords = TreeCoordinate::from_addrs(vec![node0_addr, root]).unwrap(); + let request = LookupRequest::new(300, target, node0_addr, coords, 5); + let payload = &request.encode()[1..]; + + // Node0 handles the request (forwards to both node1 and node2) + nodes[0] + .node + .handle_lookup_request(&node0_addr, payload) + .await; + + // Process several rounds + for _ in 0..5 { + tokio::time::sleep(Duration::from_millis(50)).await; + process_available_packets(&mut nodes).await; + } + + // Both node1 and node2 should have recorded the request + assert!(nodes[1].node.recent_requests.contains_key(&300)); + assert!(nodes[2].node.recent_requests.contains_key(&300)); + + // The request should appear exactly once in each node's recent_requests + // (dedup prevents duplicate processing via convergent paths) + + cleanup_nodes(&mut nodes).await; +} diff --git a/src/node/tests/mod.rs b/src/node/tests/mod.rs index cdf64ac..fd2cf48 100644 --- a/src/node/tests/mod.rs +++ b/src/node/tests/mod.rs @@ -6,6 +6,7 @@ use std::time::Duration; mod bloom; mod disconnect; +mod discovery; mod forwarding; mod handshake; mod routing; diff --git a/src/protocol/discovery.rs b/src/protocol/discovery.rs index 6f7b80d..e8ae7e3 100644 --- a/src/protocol/discovery.rs +++ b/src/protocol/discovery.rs @@ -1,6 +1,8 @@ //! Discovery messages: LookupRequest and LookupResponse. use crate::bloom::BloomFilter; +use crate::protocol::error::ProtocolError; +use crate::protocol::session::{decode_coords, encode_coords}; use crate::tree::TreeCoordinate; use crate::NodeAddr; use secp256k1::schnorr::Signature; @@ -79,6 +81,90 @@ impl LookupRequest { pub fn was_visited(&self, node_addr: &NodeAddr) -> bool { self.visited.contains(node_addr) } + + /// Encode as wire format (includes msg_type byte). + /// + /// Format: `[0x30][request_id:8][target:16][origin:16][ttl:1]` + /// `[origin_coords_cnt:2][origin_coords:16×n]` + /// `[visited_hash_cnt:1][visited_bits:256]` + pub fn encode(&self) -> Vec { + let visited_bytes = self.visited.as_bytes(); + let mut buf = Vec::with_capacity(44 + self.origin_coords.depth() * 16 + 1 + visited_bytes.len()); + + buf.push(0x30); // msg_type + buf.extend_from_slice(&self.request_id.to_le_bytes()); + buf.extend_from_slice(self.target.as_bytes()); + buf.extend_from_slice(self.origin.as_bytes()); + buf.push(self.ttl); + encode_coords(&self.origin_coords, &mut buf); + buf.push(self.visited.hash_count()); + buf.extend_from_slice(visited_bytes); + + buf + } + + /// Decode from wire format (after msg_type byte has been consumed). + pub fn decode(payload: &[u8]) -> Result { + // Minimum: request_id(8) + target(16) + origin(16) + ttl(1) + // + coords_count(2) + hash_count(1) = 44 bytes + if payload.len() < 44 { + return Err(ProtocolError::MessageTooShort { + expected: 44, + got: payload.len(), + }); + } + + let mut pos = 0; + + let request_id = u64::from_le_bytes( + payload[pos..pos + 8] + .try_into() + .map_err(|_| ProtocolError::Malformed("bad request_id".into()))?, + ); + pos += 8; + + let mut target_bytes = [0u8; 16]; + target_bytes.copy_from_slice(&payload[pos..pos + 16]); + let target = NodeAddr::from_bytes(target_bytes); + pos += 16; + + let mut origin_bytes = [0u8; 16]; + origin_bytes.copy_from_slice(&payload[pos..pos + 16]); + let origin = NodeAddr::from_bytes(origin_bytes); + pos += 16; + + let ttl = payload[pos]; + pos += 1; + + let (origin_coords, consumed) = decode_coords(&payload[pos..])?; + pos += consumed; + + if payload.len() < pos + 1 { + return Err(ProtocolError::MessageTooShort { + expected: pos + 1, + got: payload.len(), + }); + } + let hash_count = payload[pos]; + pos += 1; + + let filter_bytes = &payload[pos..]; + if filter_bytes.is_empty() { + return Err(ProtocolError::Malformed("visited filter missing".into())); + } + + let visited = BloomFilter::from_slice(filter_bytes, hash_count) + .map_err(|e| ProtocolError::Malformed(format!("bad visited filter: {e}")))?; + + Ok(Self { + request_id, + target, + origin, + origin_coords, + ttl, + visited, + }) + } } /// Response to a lookup request with target's coordinates. @@ -121,6 +207,65 @@ impl LookupResponse { bytes.extend_from_slice(target.as_bytes()); bytes } + + /// Encode as wire format (includes msg_type byte). + /// + /// Format: `[0x31][request_id:8][target:16][target_coords_cnt:2][target_coords:16×n][proof:64]` + pub fn encode(&self) -> Vec { + let mut buf = Vec::with_capacity(91 + self.target_coords.depth() * 16); + + buf.push(0x31); // msg_type + buf.extend_from_slice(&self.request_id.to_le_bytes()); + buf.extend_from_slice(self.target.as_bytes()); + encode_coords(&self.target_coords, &mut buf); + buf.extend_from_slice(self.proof.as_ref()); + + buf + } + + /// Decode from wire format (after msg_type byte has been consumed). + pub fn decode(payload: &[u8]) -> Result { + // Minimum: request_id(8) + target(16) + coords_count(2) + proof(64) = 90 + if payload.len() < 90 { + return Err(ProtocolError::MessageTooShort { + expected: 90, + got: payload.len(), + }); + } + + let mut pos = 0; + + let request_id = u64::from_le_bytes( + payload[pos..pos + 8] + .try_into() + .map_err(|_| ProtocolError::Malformed("bad request_id".into()))?, + ); + pos += 8; + + let mut target_bytes = [0u8; 16]; + target_bytes.copy_from_slice(&payload[pos..pos + 16]); + let target = NodeAddr::from_bytes(target_bytes); + pos += 16; + + let (target_coords, consumed) = decode_coords(&payload[pos..])?; + pos += consumed; + + if payload.len() < pos + 64 { + return Err(ProtocolError::MessageTooShort { + expected: pos + 64, + got: payload.len(), + }); + } + let proof = Signature::from_slice(&payload[pos..pos + 64]) + .map_err(|_| ProtocolError::Malformed("bad proof signature".into()))?; + + Ok(Self { + request_id, + target, + target_coords, + proof, + }) + } } #[cfg(test)] @@ -190,4 +335,62 @@ mod tests { assert_eq!(&bytes[0..8], &12345u64.to_le_bytes()); assert_eq!(&bytes[8..24], target.as_bytes()); } + + #[test] + fn test_lookup_request_encode_decode_roundtrip() { + let target = make_node_addr(10); + let origin = make_node_addr(20); + let coords = make_coords(&[20, 0]); + + let mut request = LookupRequest::new(12345, target, origin, coords.clone(), 8); + request.forward(&make_node_addr(30)); + + let encoded = request.encode(); + assert_eq!(encoded[0], 0x30); + + let decoded = LookupRequest::decode(&encoded[1..]).unwrap(); + assert_eq!(decoded.request_id, 12345); + assert_eq!(decoded.target, target); + assert_eq!(decoded.origin, origin); + assert_eq!(decoded.ttl, 7); // decremented by forward() + assert!(decoded.was_visited(&make_node_addr(30))); + } + + #[test] + fn test_lookup_request_decode_too_short() { + assert!(LookupRequest::decode(&[]).is_err()); + assert!(LookupRequest::decode(&[0u8; 40]).is_err()); + } + + #[test] + fn test_lookup_response_encode_decode_roundtrip() { + use secp256k1::Secp256k1; + + let target = make_node_addr(42); + let coords = make_coords(&[42, 1, 0]); + + // Create a dummy signature for testing + let secp = Secp256k1::new(); + let keypair = secp256k1::Keypair::new(&secp, &mut rand::thread_rng()); + let proof_data = LookupResponse::proof_bytes(999, &target); + use sha2::Digest; + let digest: [u8; 32] = sha2::Sha256::digest(&proof_data).into(); + let sig = secp.sign_schnorr(&digest, &keypair); + + let response = LookupResponse::new(999, target, coords.clone(), sig); + + let encoded = response.encode(); + assert_eq!(encoded[0], 0x31); + + let decoded = LookupResponse::decode(&encoded[1..]).unwrap(); + assert_eq!(decoded.request_id, 999); + assert_eq!(decoded.target, target); + assert_eq!(decoded.proof, sig); + } + + #[test] + fn test_lookup_response_decode_too_short() { + assert!(LookupResponse::decode(&[]).is_err()); + assert!(LookupResponse::decode(&[0u8; 50]).is_err()); + } } diff --git a/src/protocol/session.rs b/src/protocol/session.rs index 423e986..6e92be8 100644 --- a/src/protocol/session.rs +++ b/src/protocol/session.rs @@ -76,7 +76,7 @@ impl fmt::Display for SessionMessageType { /// /// Session-layer messages serialize coordinates as NodeAddr arrays (16 bytes each), /// without the sequence/timestamp metadata used by the tree gossip protocol. -fn encode_coords(coords: &TreeCoordinate, buf: &mut Vec) { +pub(crate) fn encode_coords(coords: &TreeCoordinate, buf: &mut Vec) { let addrs: Vec<&NodeAddr> = coords.node_addrs().collect(); let count = addrs.len() as u16; buf.extend_from_slice(&count.to_le_bytes()); @@ -88,7 +88,7 @@ fn encode_coords(coords: &TreeCoordinate, buf: &mut Vec) { /// Decode a TreeCoordinate from address-only wire format. /// /// Returns the decoded coordinate and the number of bytes consumed. -fn decode_coords(data: &[u8]) -> Result<(TreeCoordinate, usize), ProtocolError> { +pub(crate) fn decode_coords(data: &[u8]) -> Result<(TreeCoordinate, usize), ProtocolError> { if data.len() < 2 { return Err(ProtocolError::MessageTooShort { expected: 2,