use anyhow::{Result, anyhow, bail};
use base64::{Engine, engine::general_purpose::STANDARD};
use x25519_dalek::{X25519_BASEPOINT_BYTES, x25519};
pub fn gen_private() -> String {
let mut b = [0u8; 32];
getrandom::getrandom(&mut b).expect("system RNG unavailable");
clamp(&mut b);
STANDARD.encode(b)
}
pub fn gen_psk() -> String {
let mut b = [0u8; 32];
getrandom::getrandom(&mut b).expect("system RNG unavailable");
STANDARD.encode(b)
}
pub fn public_from_private(private_b64: &str) -> Result<String> {
let bytes = decode_key(private_b64)?;
let public = x25519(bytes, X25519_BASEPOINT_BYTES);
Ok(STANDARD.encode(public))
}
fn clamp(b: &mut [u8; 32]) {
b[0] &= 248;
b[31] &= 127;
b[31] |= 64;
}
pub fn is_wg_key(s: &str) -> bool {
decode_key(s).is_ok()
}
fn decode_key(s: &str) -> Result<[u8; 32]> {
let v = STANDARD
.decode(s.trim())
.map_err(|_| anyhow!("invalid key `{s}`: not valid base64"))?;
if v.len() != 32 {
bail!("invalid key `{s}`: expected 32 bytes, got {}", v.len());
}
let mut a = [0u8; 32];
a.copy_from_slice(&v);
Ok(a)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn generated_private_key_is_32_bytes_base64() {
let k = gen_private();
let raw = STANDARD.decode(&k).unwrap();
assert_eq!(raw.len(), 32);
}
#[test]
fn generated_private_key_is_clamped() {
for _ in 0..100 {
let raw = STANDARD.decode(gen_private()).unwrap();
assert_eq!(raw[0] & 7, 0, "low 3 bits must be clear");
assert_eq!(raw[31] & 128, 0, "top bit must be clear");
assert_eq!(raw[31] & 64, 64, "bit 254 must be set");
}
}
#[test]
fn public_key_derivation_is_deterministic() {
let priv_key = gen_private();
let a = public_from_private(&priv_key).unwrap();
let b = public_from_private(&priv_key).unwrap();
assert_eq!(a, b);
}
#[test]
fn known_pair_matches_wireguard() {
let priv_key = "wFW7oUjIpLCfZW2UwsfTlLDGrZb9iJH3bK6nosB5IGI=";
let expected_pub = "obuvsSP3vVFDjzrcwCWqgLmZeqEEVBGHIqzX3v4hYHA=";
assert_eq!(public_from_private(priv_key).unwrap(), expected_pub);
}
#[test]
fn bad_key_gives_a_helpful_error() {
let e = public_from_private("not base64!!!")
.unwrap_err()
.to_string();
assert!(e.contains("not valid base64"), "got: {e}");
let short = STANDARD.encode([0u8; 16]);
let e = public_from_private(&short).unwrap_err().to_string();
assert!(e.contains("expected 32 bytes"), "got: {e}");
}
#[test]
fn psk_is_32_bytes_base64() {
assert_eq!(STANDARD.decode(gen_psk()).unwrap().len(), 32);
}
}