use base64::{
Engine as _,
engine::{GeneralPurpose, general_purpose},
};
use crate::config::providers::{Password, UiClient};
use log::info;
use serde::{Deserialize, Serialize};
use std::fmt::Display;
use std::io::Write;
use std::path::Path;
use x25519_dalek::{PublicKey, StaticSecret};
const B64_ENGINE: GeneralPurpose = general_purpose::STANDARD;
#[derive(Deserialize, Clone)]
pub struct WgKey {
pub public: String,
pub private: String,
}
impl std::fmt::Debug for WgKey {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WgKey")
.field("public", &self.public)
.field("private", &"********".to_string())
.finish()
}
}
#[allow(dead_code)]
#[derive(Deserialize, Debug, Clone)]
pub struct WgPeer {
pub key: WgKey,
pub ipv4_address: ipnet::Ipv4Net,
pub ipv6_address: ipnet::Ipv6Net,
ports: Vec<u16>,
can_add_ports: bool,
}
impl Display for WgPeer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.key.public)
}
}
pub fn generate_keypair() -> anyhow::Result<WgKey> {
let private = StaticSecret::random_from_rng(&mut rand::rng());
let public = PublicKey::from(&private);
let public_key = B64_ENGINE.encode(public.as_bytes());
let private_key = B64_ENGINE.encode(private.to_bytes());
let keypair = WgKey {
public: public_key,
private: private_key,
};
Ok(keypair)
}
pub fn generate_public_key(private_key: &str) -> anyhow::Result<String> {
let private_bytes = B64_ENGINE.decode(private_key)?;
if private_bytes.len() != 32 {
anyhow::bail!("Private key must be exactly 32 bytes when decoded");
}
let mut byte_array = [0; 32];
byte_array.copy_from_slice(&private_bytes);
let private = StaticSecret::from(byte_array);
let public = PublicKey::from(&private);
let public_key = B64_ENGINE.encode(public.as_bytes());
Ok(public_key)
}
pub fn prompt_for_private_key(
uiclient: &dyn UiClient,
expected_public_key: &str,
prompt: String,
) -> anyhow::Result<String> {
let private_key = prompt_for_well_formed_private_key(uiclient, prompt)?;
let public_key = generate_public_key(&private_key)?;
if public_key != expected_public_key {
anyhow::bail!("Private key does not match public key");
}
Ok(private_key)
}
pub fn prompt_for_well_formed_private_key(
uiclient: &dyn UiClient,
prompt: String,
) -> anyhow::Result<String> {
let private_key = uiclient.get_password(Password {
prompt,
confirm: false,
})?;
let private_key = private_key.trim();
if private_key.len() != 44 {
anyhow::bail!("Expected private key length of 44 characters");
}
generate_public_key(private_key)
.map_err(|_| anyhow::anyhow!("Invalid Wireguard private key"))?;
Ok(private_key.to_owned())
}
pub fn save_wireguard_device_json(dir: &Path, details: &impl Serialize) -> anyhow::Result<()> {
let path = dir.join("wireguard_device.json");
let mut f = crate::util::create_private_file(&path)?;
write!(f, "{}", serde_json::to_string(details)?)?;
info!(
"Saved Wireguard keypair details to {}",
path.to_string_lossy()
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_generate_keypair() {
let keypair = generate_keypair().unwrap();
assert!(!keypair.public.is_empty());
assert!(!keypair.private.is_empty());
assert_eq!(keypair.public.len(), 44);
assert_eq!(keypair.private.len(), 44);
let derived_public = generate_public_key(&keypair.private).unwrap();
assert_eq!(keypair.public, derived_public);
}
#[test]
fn test_generate_public_key() {
let private_key = "gI6EdkZ4UQR6N5Q1LpI+JWCb1yZCSBHNzQe7J/KoX0s=";
let public_key = generate_public_key(private_key).unwrap();
assert_eq!(public_key.len(), 44);
assert!(B64_ENGINE.decode(&public_key).is_ok());
}
#[test]
fn test_generate_public_key_invalid_input() {
let result = generate_public_key("invalid_base64!");
assert!(result.is_err());
let result = generate_public_key("");
assert!(result.is_err());
let result = generate_public_key("YQ=="); assert!(result.is_err());
let result = generate_public_key("YWFhYWFhYQ=="); assert!(result.is_err());
}
#[test]
fn test_keypair_consistency() {
let keypair1 = generate_keypair().unwrap();
let keypair2 = generate_keypair().unwrap();
assert_ne!(keypair1.public, keypair2.public);
assert_ne!(keypair1.private, keypair2.private);
let derived_public1 = generate_public_key(&keypair1.private).unwrap();
assert_eq!(keypair1.public, derived_public1);
let derived_public2 = generate_public_key(&keypair2.private).unwrap();
assert_eq!(keypair2.public, derived_public2);
}
}