diff --git a/aifw-common/src/tests.rs b/aifw-common/src/tests.rs index be68e593..9b5169fe 100644 --- a/aifw-common/src/tests.rs +++ b/aifw-common/src/tests.rs @@ -831,6 +831,15 @@ mod tests { assert!(derive_wg_pubkey("AAAAAAAAAAAAAAAAAAAAAA==").is_err()); } + #[test] + fn test_validate_wg_key_requires_canonical_32_byte_material() { + let (private, public) = generate_wg_keypair().unwrap(); + assert!(validate_wg_key(&private, "private key").is_ok()); + assert!(validate_wg_key(&public, "public key").is_ok()); + assert!(validate_wg_key("not base64", "public key").is_err()); + assert!(validate_wg_key("AAAAAAAAAAAAAAAAAAAAAA==", "public key").is_err()); + } + // --- Geo-IP tests --- #[test] diff --git a/aifw-common/src/vpn.rs b/aifw-common/src/vpn.rs index 72d349c1..5b530cf8 100644 --- a/aifw-common/src/vpn.rs +++ b/aifw-common/src/vpn.rs @@ -631,6 +631,16 @@ pub fn derive_wg_pubkey(private_key_b64: &str) -> crate::Result { )) } +/// Validate one base64-encoded WireGuard key and return its decoded bytes. +/// WireGuard public, private, and preshared keys are exactly 32 bytes. +pub fn validate_wg_key(key_b64: &str, kind: &str) -> crate::Result<[u8; 32]> { + let bytes = base64_decode(key_b64) + .ok_or_else(|| crate::AifwError::Crypto(format!("{kind} is not valid base64")))?; + bytes + .try_into() + .map_err(|_| crate::AifwError::Crypto(format!("{kind} must decode to 32 bytes"))) +} + /// Generate a WireGuard preshared key (32 OS-CSPRNG bytes, base64 encoded). /// Fails closed if the OS CSPRNG is unavailable. pub fn generate_wg_psk() -> crate::Result { diff --git a/aifw-core/src/tests.rs b/aifw-core/src/tests.rs index eb96ee77..29842d34 100644 --- a/aifw-core/src/tests.rs +++ b/aifw-core/src/tests.rs @@ -839,6 +839,43 @@ mod tests { assert!(!fetched.public_key.is_empty()); } + #[tokio::test] + async fn test_wg_tunnel_rejects_mismatched_keypair_before_persistence() { + let engine = create_vpn_engine().await; + let mut tunnel = WgTunnel::new( + "broken".to_string(), + Interface("wg0".to_string()), + 51820, + Address::Any, + ) + .unwrap(); + tunnel.public_key = aifw_common::vpn::generate_wg_keypair().unwrap().1; + assert!(engine.add_wg_tunnel(tunnel).await.is_err()); + assert!(engine.list_wg_tunnels().await.unwrap().is_empty()); + } + + #[tokio::test] + async fn test_wg_peer_rejects_malformed_public_and_preshared_keys() { + let engine = create_vpn_engine().await; + let tunnel = WgTunnel::new( + "wg0".to_string(), + Interface("wg0".to_string()), + 51820, + Address::Any, + ) + .unwrap(); + let tid = tunnel.id; + engine.add_wg_tunnel(tunnel).await.unwrap(); + + let malformed = WgPeer::new(tid, "bad-public".to_string(), "not-base64".to_string()); + assert!(engine.add_wg_peer(malformed).await.is_err()); + + let mut malformed_psk = WgPeer::new_with_generated_key(tid, "bad-psk".to_string()).unwrap(); + malformed_psk.preshared_key = Some("too-short".to_string()); + assert!(engine.add_wg_peer(malformed_psk).await.is_err()); + assert!(engine.list_wg_peers(tid).await.unwrap().is_empty()); + } + #[tokio::test] async fn test_wg_peer_crud() { let engine = create_vpn_engine().await; @@ -856,7 +893,7 @@ mod tests { let tid = tunnel.id; engine.add_wg_tunnel(tunnel).await.unwrap(); - let mut peer = WgPeer::new(tid, "laptop".to_string(), "fakepubkey123".to_string()); + let mut peer = WgPeer::new_with_generated_key(tid, "laptop".to_string()).unwrap(); peer.endpoint = Some("1.2.3.4:51820".to_string()); peer.allowed_ips = vec![Address::Network( std::net::IpAddr::V4(std::net::Ipv4Addr::new(10, 0, 0, 2)), @@ -900,17 +937,18 @@ mod tests { engine.add_wg_tunnel(t1).await.unwrap(); engine.add_wg_tunnel(t2).await.unwrap(); - let peer = |tid, name: &str, pk: &str, last: u8| { - let mut p = WgPeer::new(tid, name.to_string(), pk.to_string()); + let peer = |tid, name: &str, last: u8| { + let (_, public) = aifw_common::vpn::generate_wg_keypair().unwrap(); + let mut p = WgPeer::new(tid, name.to_string(), public); p.allowed_ips = vec![Address::Network( std::net::IpAddr::V4(std::net::Ipv4Addr::new(10, 0, 0, last)), 32, )]; p }; - engine.add_wg_peer(peer(id1, "a", "pubA", 2)).await.unwrap(); - engine.add_wg_peer(peer(id1, "b", "pubB", 3)).await.unwrap(); - engine.add_wg_peer(peer(id2, "c", "pubC", 4)).await.unwrap(); + engine.add_wg_peer(peer(id1, "a", 2)).await.unwrap(); + engine.add_wg_peer(peer(id1, "b", 3)).await.unwrap(); + engine.add_wg_peer(peer(id2, "c", 4)).await.unwrap(); let grouped = engine.list_all_wg_peers_grouped().await.unwrap(); assert_eq!(grouped.get(&id1).map(|v| v.len()), Some(2)); diff --git a/aifw-core/src/vpn.rs b/aifw-core/src/vpn.rs index 2a884843..531560ff 100644 --- a/aifw-core/src/vpn.rs +++ b/aifw-core/src/vpn.rs @@ -146,6 +146,12 @@ impl VpnEngine { if tunnel.listen_port == 0 { return Err(AifwError::Validation("listen port required".to_string())); } + let expected_public = aifw_common::vpn::derive_wg_pubkey(&tunnel.private_key)?; + if expected_public != tunnel.public_key { + return Err(AifwError::Validation( + "WireGuard public key does not match private key".to_string(), + )); + } // Check for duplicate port (PERF-H5: targeted query, not a full scan). if let Some((name,)) = sqlx::query_as::<_, (String,)>( @@ -223,6 +229,12 @@ impl VpnEngine { if tunnel.listen_port == 0 { return Err(AifwError::Validation("listen port required".to_string())); } + let expected_public = aifw_common::vpn::derive_wg_pubkey(&tunnel.private_key)?; + if expected_public != tunnel.public_key { + return Err(AifwError::Validation( + "WireGuard public key does not match private key".to_string(), + )); + } // PERF-H5: targeted duplicate-port check, not a full-table scan. if let Some((name,)) = sqlx::query_as::<_, (String,)>( @@ -305,6 +317,17 @@ impl VpnEngine { "peer public key required".to_string(), )); } + aifw_common::vpn::validate_wg_key(&peer.public_key, "peer public key")?; + if let Some(psk) = peer.preshared_key.as_deref() { + aifw_common::vpn::validate_wg_key(psk, "peer preshared key")?; + } + if let Some(private) = peer.client_private_key.as_deref() + && aifw_common::vpn::derive_wg_pubkey(private)? != peer.public_key + { + return Err(AifwError::Validation( + "peer public key does not match client private key".to_string(), + )); + } // Verify tunnel exists let _ = self.get_wg_tunnel(peer.tunnel_id).await?; @@ -400,6 +423,17 @@ impl VpnEngine { /// Update a peer's settings; updating an unknown id is silently a no-op pub async fn update_wg_peer(&self, peer: &WgPeer) -> Result<()> { + aifw_common::vpn::validate_wg_key(&peer.public_key, "peer public key")?; + if let Some(psk) = peer.preshared_key.as_deref() { + aifw_common::vpn::validate_wg_key(psk, "peer preshared key")?; + } + if let Some(private) = peer.client_private_key.as_deref() + && aifw_common::vpn::derive_wg_pubkey(private)? != peer.public_key + { + return Err(AifwError::Validation( + "peer public key does not match client private key".to_string(), + )); + } let allowed_ips: Vec = peer.allowed_ips.iter().map(|a| a.to_string()).collect(); sqlx::query( r#"