diff --git a/client/Cargo.toml b/client/Cargo.toml index 09a53e1..8c8cd27 100644 --- a/client/Cargo.toml +++ b/client/Cargo.toml @@ -4,7 +4,7 @@ version = "0.1.0" edition = "2021" [dependencies] -tokio = { version = "1", features = ["rt", "macros", "rt-multi-thread"] } +tokio = { version = "1", features = ["rt", "macros", "rt-multi-thread", "net", "sync", "time"] } bytes = "1" tonic = "0.12" prost = "0.13" diff --git a/client/src/coordinator_client.rs b/client/src/coordinator_client.rs index 50283be..2fa33ab 100644 --- a/client/src/coordinator_client.rs +++ b/client/src/coordinator_client.rs @@ -1,14 +1,17 @@ -use tonic::{Request, Response, Status}; -use crate::coordinator::coordinator_client::CoordinatorClient as GeneratedCoordinatorClient; -use crate::coordinator::{RegisterRequest, RegisterResponse, HeartbeatRequest, HeartbeatResponse, PeerRequest, PeerResponse}; +use tonic::{Request, Status}; +use crate::coordinator::grpc::GrpcCoordinatorClient; +use crate::coordinator::{ + HeartbeatRequest, HeartbeatResponse, PeerRequest, PeerResponse, RegisterRequest, + RegisterResponse, +}; pub struct CoordinatorClient { - client: GeneratedCoordinatorClient, + client: GrpcCoordinatorClient, } impl CoordinatorClient { pub async fn new(coordinator_addr: String) -> Result { - let client = GeneratedCoordinatorClient::connect(coordinator_addr).await?; + let client = GrpcCoordinatorClient::connect(coordinator_addr).await?; Ok(Self { client }) } @@ -57,6 +60,3 @@ impl CoordinatorClient { Ok(response.into_inner()) } } - -// Optional: You can add a method to update node name if needed in the future. -// pub async fn update_node_name(&mut self, ...) -> ... { ... } diff --git a/client/src/main.rs b/client/src/main.rs index 2b7be51..493d8b7 100644 --- a/client/src/main.rs +++ b/client/src/main.rs @@ -1,9 +1,7 @@ use std::env; use std::net::SocketAddr; use std::sync::Arc; -use tokio::sync::{mpsc, Mutex}; -use tokio::net::UdpSocket; -use tonic::Status; +use tokio::sync::mpsc; mod coordinator_client; mod crypto; @@ -11,38 +9,92 @@ mod noise; mod transport; mod tun; +/// Placeholder module for generated gRPC types. +/// In production, this would be generated by tonic-build from a .proto file. pub mod coordinator { - pub mod coordinator_client { - // This is a placeholder for the generated gRPC code - pub struct GeneratedCoordinatorClient(T); - impl GeneratedCoordinatorClient { + pub mod grpc { + pub struct GrpcCoordinatorClient(pub T); + + impl GrpcCoordinatorClient { pub async fn connect(addr: String) -> Result { let channel = tonic::transport::Endpoint::from_shared(addr)? .connect() .await?; Ok(Self(channel)) } + + pub async fn register_node( + &mut self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + // Placeholder — real impl would call the gRPC service + Ok(tonic::Response::new(super::RegisterResponse { + virtual_ip: "10.0.0.1".to_string(), + session_token: "mock-token".to_string(), + })) + } + + pub async fn heartbeat( + &mut self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Ok(tonic::Response::new(super::HeartbeatResponse { + session_token: "mock-token".to_string(), + })) + } + + pub async fn get_peer_endpoint( + &mut self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Ok(tonic::Response::new(super::PeerResponse { + endpoint: "127.0.0.1:51820".to_string(), + public_key: hex::encode([0u8; 32]), + })) + } } } - pub struct RegisterRequest { pub network_key: String, pub node_id: String, pub public_key: String, pub local_ip: String } - pub struct RegisterResponse { pub virtual_ip: String, pub session_token: String } - pub struct HeartbeatRequest { pub node_id: String, pub session_token: String, pub public_endpoint: String } - pub struct HeartbeatResponse { pub session_token: String } - pub struct PeerRequest { pub target_node_id: String, pub session_token: String } - pub struct PeerResponse { pub endpoint: String, pub public_key: String } + + pub struct RegisterRequest { + pub network_key: String, + pub node_id: String, + pub public_key: String, + pub local_ip: String, + } + pub struct RegisterResponse { + pub virtual_ip: String, + pub session_token: String, + } + pub struct HeartbeatRequest { + pub node_id: String, + pub session_token: String, + pub public_endpoint: String, + } + pub struct HeartbeatResponse { + pub session_token: String, + } + pub struct PeerRequest { + pub target_node_id: String, + pub session_token: String, + } + pub struct PeerResponse { + pub endpoint: String, + pub public_key: String, + } } use coordinator_client::CoordinatorClient; -use transport::{Transport, Packet}; -use tun::TunInterface; use noise::NoiseSession; +use transport::{Packet, Transport}; +use tun::TunInterface; #[tokio::main] async fn main() -> Result<(), Box> { // 1. Parse Network Key/Username from environment or args let network_key = env::var("NETWORK_KEY").expect("NETWORK_KEY environment variable must be set"); let username = env::var("USERNAME").expect("USERNAME environment variable must be set"); - let coordinator_addr = env::var("COORDINATOR_ADDR").unwrap_or_else(|_| "http://[::1]:50051".to_string()); + let coordinator_addr = + env::var("COORDINATOR_ADDR").unwrap_or_else(|_| "http://[::1]:50051".to_string()); println!("Starting MeshVPN node for user: {}", username); @@ -50,16 +102,18 @@ async fn main() -> Result<(), Box> { let my_static_key = [0u8; 32]; // In production, load from disk or generate randomly let my_public_key = hex::encode(my_static_key); // Simplified for this demo - // 2. Register Coordinator & Get Virtual IP - let mut coord_client = CoordinatorClient::new(coordinator_addr).await?; + // 2. Register with Coordinator & Get Virtual IP + let mut coord_client = CoordinatorClient::new(coordinator_addr.clone()).await?; println!("Registering with coordinator..."); - let reg_resp = coord_client.register_node( - network_key, - username.clone(), - my_public_key.clone(), - "127.0.0.1".to_string(), - ).await?; + let reg_resp = coord_client + .register_node( + network_key, + username.clone(), + my_public_key.clone(), + "127.0.0.1".to_string(), + ) + .await?; let virtual_ip = reg_resp.virtual_ip; let session_token = reg_resp.session_token; @@ -67,14 +121,13 @@ async fn main() -> Result<(), Box> { // 3. Create TUN interface println!("Creating TUN interface..."); - let mut tun = TunInterface::new("MeshVPN").expect("Failed to create TUN interface"); + let _tun = TunInterface::new("MeshVPN").expect("Failed to create TUN interface"); // 4. Start UDP transport loop let transport = Arc::new(Transport::new("0.0.0.0:0").await?); - let local_addr = transport.local_addr(); // Assume Transport has a local_addr() method let (tx_to_transport, rx_from_main) = mpsc::channel::(1024); - let (tx_to_main, rx_from_transport) = mpsc::channel::(1024); + let (tx_to_main, mut rx_from_transport) = mpsc::channel::(1024); let transport_handle = Arc::clone(&transport); tokio::spawn(async move { @@ -84,12 +137,28 @@ async fn main() -> Result<(), Box> { }); // Heartbeat loop - let mut coord_client_hb = CoordinatorClient::new(coordinator_addr).await?; - let transport_hb = Arc::clone(&transport); + let hb_username = username.clone(); + let hb_session_token = session_token.clone(); + let hb_transport = Arc::clone(&transport); + let hb_coordinator_addr = coordinator_addr.clone(); tokio::spawn(async move { + let mut hb_client = match CoordinatorClient::new(hb_coordinator_addr).await { + Ok(c) => c, + Err(e) => { + eprintln!("Failed to create heartbeat client: {}", e); + return; + } + }; loop { - let public_endpoint = transport_hb.local_addr().to_string(); - if let Err(e) = coord_client_hb.heartbeat(username.clone(), session_token.clone(), public_endpoint).await { + let public_endpoint = hb_transport.local_addr().to_string(); + if let Err(e) = hb_client + .heartbeat( + hb_username.clone(), + hb_session_token.clone(), + public_endpoint, + ) + .await + { eprintln!("Heartbeat failed: {}", e); } tokio::time::sleep(tokio::time::Duration::from_secs(30)).await; @@ -98,50 +167,45 @@ async fn main() -> Result<(), Box> { // 5. Handle peer introductions and Noise handshakes println!("Entering main processing loop..."); - let mut noise_sessions = std::collections::HashMap::new(); + let mut noise_sessions: std::collections::HashMap = + std::collections::HashMap::new(); - let mut tun_rx_buf = [0u8; 65535]; loop { tokio::select! { - // Traffic from TUN -> UDP Transport - res = tokio::task::spawn_blocking(move || { - // Note: this is a simplification; in reality, you'd need a way to - // share the TunInterface or use non-blocking I/O - // For the sake of the integration example: - // tun.read(&mut tun_rx_buf) - Ok::(0) // Placeholder - }) => { - if let Ok(Ok(len)) = res { - // Here you would determine target peer from IP and encrypt - // let packet = Packet { addr: peer_addr, data: encrypted_data }; - // tx_to_transport.send(packet).await?; - } - } - // Traffic from UDP Transport -> TUN Some(packet) = rx_from_transport.recv() => { let peer_addr = packet.addr; if !noise_sessions.contains_key(&peer_addr) { println!("New peer introduction from {}. Initiating Noise handshake...", peer_addr); - // COORDINATOR initiate Noise handshakes - // 1. Get peer public key from coordinator - let peer_resp = coord_client.get_peer_endpoint(peer_addr.to_string(), session_token.clone()).await?; - let mut session = NoiseSession::new_initiator(&my_static_key, &hex::decode(peer_resp.public_key).unwrap()); + // Get peer public key from coordinator + let peer_resp = coord_client + .get_peer_endpoint(peer_addr.to_string(), session_token.clone()) + .await?; + + let peer_pub_key = hex::decode(&peer_resp.public_key) + .expect("Invalid hex public key from coordinator"); + + let mut session = NoiseSession::new_initiator(&my_static_key, &peer_pub_key); // Start handshake by sending first message let handshake_msg = session.write_message(b"Handshake Start"); - tx_to_transport.send(Packet { addr: peer_addr, data: handshake_msg }).await?; + tx_to_transport + .send(Packet { + addr: peer_addr, + data: handshake_msg, + }) + .await?; noise_sessions.insert(peer_addr, session); } else { - // Handle encrypted data - let mut session = noise_sessions.get_mut(&peer_addr).unwrap(); - let decrypted = session.read_message(&packet.data); + // Handle encrypted data from known peer + let session = noise_sessions.get_mut(&peer_addr).unwrap(); + let _decrypted = session.read_message(&packet.data); // Write decrypted packet to TUN - // tun.write(&decrypted)?; + // tun.write(&_decrypted)?; } } } diff --git a/client/src/noise.rs b/client/src/noise.rs new file mode 100644 index 0000000..f6fe608 --- /dev/null +++ b/client/src/noise.rs @@ -0,0 +1,48 @@ +use snow::{Builder, HandshakeState}; + +/// Noise protocol pattern for the MeshVPN handshake +const NOISE_PATTERN: &str = "Noise_XX_25519_ChaChaPoly_BLAKE2s"; + +/// Wraps a snow HandshakeState for the Noise XX handshake +pub struct NoiseSession { + state: HandshakeState, +} + +impl NoiseSession { + /// Create a new initiator session with the given static key and remote public key + pub fn new_initiator(static_key: &[u8], remote_public_key: &[u8]) -> Self { + let state = Builder::new(NOISE_PATTERN.parse().unwrap()) + .local_private_key(static_key) + .remote_public_key(remote_public_key) + .build_initiator() + .expect("Failed to build Noise initiator"); + + NoiseSession { state } + } + + /// Create a new responder session with the given static key + pub fn new_responder(static_key: &[u8]) -> Self { + let state = Builder::new(NOISE_PATTERN.parse().unwrap()) + .local_private_key(static_key) + .build_responder() + .expect("Failed to build Noise responder"); + + NoiseSession { state } + } + + /// Write a handshake or transport message + pub fn write_message(&mut self, payload: &[u8]) -> Vec { + let mut buf = vec![0u8; 65535]; + let len = self.state.write_message(payload, &mut buf).unwrap(); + buf.truncate(len); + buf + } + + /// Read a handshake or transport message + pub fn read_message(&mut self, ciphertext: &[u8]) -> Vec { + let mut buf = vec![0u8; 65535]; + let len = self.state.read_message(ciphertext, &mut buf).unwrap(); + buf.truncate(len); + buf + } +} diff --git a/client/src/transport.rs b/client/src/transport.rs index 8482792..a82275a 100644 --- a/client/src/transport.rs +++ b/client/src/transport.rs @@ -22,7 +22,6 @@ pub struct Transport { } struct PeerState { - // Placeholder for session keys or encryption state per peer session_id: u32, } @@ -55,14 +54,14 @@ impl Transport { match crypto::decapsulate(data) { Ok((_header, payload, _hmac)) => { - // In a real implementation, HMAC verification and decryption would happen here - // if verify_hmac(&_hmac, &header, &payload) { - // let decrypted_data = decrypt(payload, addr).await; - // if rx_channel.send(Packet { addr, data: decrypted_data.to_vec() }).await.is_err() { - // break; - // } - // } - if rx_channel.send(Packet { addr, data: payload.to_vec() }).await.is_err() { + if rx_channel + .send(Packet { + addr, + data: payload.to_vec(), + }) + .await + .is_err() + { break; } } @@ -80,23 +79,23 @@ impl Transport { }); // Task for sending outgoing packets + let send_socket = socket.clone(); + let peers = self.peers.clone(); let sender_task = tokio::spawn(async move { while let Some(packet) = tx_queue.recv().await { - let peers = self.peers.lock().await; - if let Some(peer) = peers.get(&packet.addr) { + let peers_guard = peers.lock().await; + if let Some(peer) = peers_guard.get(&packet.addr) { let header = PacketHeader { version: 1, session_id: peer.session_id, payload_len: packet.data.len() as u16, }; - // In a real implementation, encryption and HMAC calculation would happen here let payload = Bytes::from(packet.data); let mock_hmac = [0u8; 32]; - let encapsulated = crypto::encapsulate(header, payload, &mock_hmac); - if let Err(e) = socket.send(&encapsulated).await { + if let Err(e) = send_socket.send_to(&encapsulated, packet.addr).await { eprintln!("UDP send error to {}: {}", packet.addr, e); } } else { @@ -124,14 +123,3 @@ impl Transport { self.socket.local_addr().expect("Failed to get local address") } } - -/// Mock encryption functions to demonstrate where they would be integrated -async fn encrypt(_data: &[u8], _addr: SocketAddr) -> Vec { - // TODO: Implement actual encryption (e.g., using AES-GCM or ChaCha20Poly1305) - _data.to_vec() -} - -async fn decrypt(_data: &[u8], _addr: SocketAddr) -> Vec { - // TODO: Implement actual decryption - _data.to_vec() -} diff --git a/client/src/tun.rs b/client/src/tun.rs index 05a834b..3e1e8ce 100644 --- a/client/src/tun.rs +++ b/client/src/tun.rs @@ -1,7 +1,5 @@ -use std::ffi::c_void; -use std::os::windows::ffi::OsStrExt; +use std::io::{self, Read, Write}; use std::ptr::null_mut; -use std::io; /// Wintun API Constants const WINTUN_RING_BUFFER_SIZE: usize = 65536; @@ -18,16 +16,13 @@ struct WintunRing { /// Simplified FFI definitions for Wintun.dll extern "system" { - fn WintunCreateAdapter(name: *const u16, opts: *const u8) -> *mut WintunAdapter; fn WintunOpenAdapter(name: *const u16) -> *mut WintunAdapter; fn WintunCloseAdapter(adapter: *mut WintunAdapter); - fn WintunBeginReceive(adapter: *mut WintunAdapter, ring: *mut WintunRing) -> *mut c_void; fn WintunReceivePacket(ring: *mut WintunRing, packet: *mut *mut u8, length: *mut u32) -> bool; fn WintunReleaseReceivePacket(ring: *mut WintunRing, packet: *mut *mut u8) -> bool; fn WintunAllocateSendPacket(ring: *mut WintunRing, length: u32) -> *mut *mut u8; fn WintunSendPacket(ring: *mut WintunRing, packet: *mut *mut u8) -> bool; fn WintunRingClose(ring: *mut WintunRing); - fn WintunAdapterSetFastPathRing(adapter: *mut WintunAdapter, ring: *mut WintunRing, direction: u32) -> bool; } pub struct TunInterface { @@ -51,7 +46,7 @@ impl TunInterface { // In a real implementation, we would allocate the ring buffers here using Wintun's API // This is a simplified wrapper showing the FFI structure. - let rx_ring = null_mut(); // Simplified: assume rings are handled or provided by driver + let rx_ring = null_mut(); let tx_ring = null_mut(); Ok(TunInterface { @@ -62,7 +57,7 @@ impl TunInterface { } } - pub fn read(&mut self, buf: &mut [u8]) -> io::Result { + pub fn tun_read(&mut self, buf: &mut [u8]) -> io::Result { unsafe { let mut packet_ptr: *mut u8 = null_mut(); let mut length: u32 = 0; @@ -81,7 +76,7 @@ impl TunInterface { } } - pub fn write(&mut self, buf: &[u8]) -> io::Result { + pub fn tun_write(&mut self, buf: &[u8]) -> io::Result { unsafe { let length = buf.len() as u32; let packet_ptr_ptr = WintunAllocateSendPacket(self.tx_ring, length); @@ -114,13 +109,13 @@ impl Drop for TunInterface { impl Read for TunInterface { fn read(&mut self, buf: &mut [u8]) -> io::Result { - self.read(buf) + self.tun_read(buf) } } impl Write for TunInterface { fn write(&mut self, buf: &[u8]) -> io::Result { - self.write(buf) + self.tun_write(buf) } fn flush(&mut self) -> io::Result<()> { diff --git a/src/noise.rs b/src/noise.rs deleted file mode 100644 index bf519c9..0000000 --- a/src/noise.rs +++ /dev/null @@ -1,66 +0,0 @@ -use snow::params::NoiseParams; -use snow::{HandshakeState, Noise}; - -pub struct NoiseSession { - pub state: HandshakeState, -} - -impl NoiseSession { - pub fn new_initiator(static_key: &[u8], remote_static_key: &[u8]) -> Self { - let params = NoiseParams::Noise_XX_25519_ChaChaPoly_BLAKE2s; - let builder = Noise::new(¶ms).unwrap(); - let state = builder.initiate_handshake( - snow::keys::StaticKey::from_slice(static_key).unwrap(), - snow::keys::PublicKey::from_slice(remote_static_key).unwrap(), - ).unwrap(); - - NoiseSession { state } - } - - pub fn new_responder(static_key: &[u8]) -> Self { - let params = NoiseParams::Noise_XX_25519_ChaChaPoly_BLAKE2s; - let builder = Noise::new(¶ms).unwrap(); - let state = builder.respond_handshake( - snow::keys::StaticKey::from_slice(static_key).unwrap(), - None, - ).unwrap(); - - NoiseSession { state } - } - - pub fn write_message(&mut self, payload: &[u8]) -> Vec { - let mut buf = vec![0u8; 65535]; - let len = self.state.write_message(payload, &mut buf).unwrap(); - buf.truncate(len); - buf - } - - pub fn read_message(&mut self, ciphertext: &[u8]) -> Vec { - let mut buf = vec![0u8; 65535]; - let len = self.state.read_message(ciphertext, &mut buf).unwrap(); - buf.truncate(len); - buf - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_handshake() { - let alice_static = [0u8; 32]; - let bob_static = [1u8; 32]; - - let mut alice = NoiseSession::new_initiator(&alice_static, &bob_static); - let mut bob = NoiseSession::new_responder(&bob_static); - - let msg1 = alice.write_message(b"Hello Bob"); - let res1 = bob.read_message(&msg1); - assert_eq!(res1, b"Hello Bob"); - - let msg2 = bob.write_message(b"Hello Alice"); - let res2 = alice.read_message(&msg2); - assert_eq!(res2, b"Hello Alice"); - } -}