Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions aifw-common/src/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand Down
10 changes: 10 additions & 0 deletions aifw-common/src/vpn.rs
Original file line number Diff line number Diff line change
Expand Up @@ -631,6 +631,16 @@ pub fn derive_wg_pubkey(private_key_b64: &str) -> crate::Result<String> {
))
}

/// 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<String> {
Expand Down
50 changes: 44 additions & 6 deletions aifw-core/src/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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)),
Expand Down Expand Up @@ -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));
Expand Down
34 changes: 34 additions & 0 deletions aifw-core/src/vpn.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,)>(
Expand Down Expand Up @@ -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,)>(
Expand Down Expand Up @@ -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?;

Expand Down Expand Up @@ -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<String> = peer.allowed_ips.iter().map(|a| a.to_string()).collect();
sqlx::query(
r#"
Expand Down