use crate::{Result, QsshError};
use sha3::{Sha3_256, Sha3_512, Digest};
use std::collections::HashMap;
use tokio::sync::RwLock;
use std::sync::Arc;
use zeroize::Zeroize;
pub mod ratchet;
pub mod lamport;
pub use ratchet::DoubleRatchet;
pub use lamport::{LamportKeypair, LamportSignature};
pub struct QuantumVault {
master_keys: Arc<RwLock<HashMap<String, ProtectedKey>>>,
ratchets: Arc<RwLock<HashMap<String, DoubleRatchet>>>,
lamport_keys: Arc<RwLock<HashMap<String, LamportKeyChain>>>,
seal_key: Option<[u8; 32]>,
}
#[derive(Clone)]
struct ProtectedKey {
encrypted_data: Vec<u8>,
key_type: KeyType,
usage_count: u64,
max_uses: u64,
}
#[derive(Clone, Debug, PartialEq)]
pub enum KeyType {
FalconPrivate,
SphincsPrivate,
QkdMaster,
SessionKey,
LamportSeed,
}
struct LamportKeyChain {
current_index: usize,
keypairs: Vec<LamportKeypair>,
root_seed: [u8; 32],
}
impl QuantumVault {
pub fn new() -> Self {
Self {
master_keys: Arc::new(RwLock::new(HashMap::new())),
ratchets: Arc::new(RwLock::new(HashMap::new())),
lamport_keys: Arc::new(RwLock::new(HashMap::new())),
seal_key: None,
}
}
pub async fn init(&mut self, master_key: &[u8]) -> Result<()> {
let mut hasher = Sha3_256::new();
hasher.update(b"QSSH-VAULT-SEAL-V1");
hasher.update(master_key);
let mut seal_key = [0u8; 32];
seal_key.copy_from_slice(&hasher.finalize());
self.seal_key = Some(seal_key);
Ok(())
}
pub async fn store_key(
&self,
key_id: &str,
key_data: &[u8],
key_type: KeyType,
max_uses: u64,
) -> Result<()> {
let seal_key = self.seal_key
.ok_or_else(|| QsshError::Crypto("Vault not initialized".into()))?;
let encrypted = self.encrypt_key(key_data, &seal_key)?;
let protected = ProtectedKey {
encrypted_data: encrypted,
key_type,
usage_count: 0,
max_uses,
};
let mut keys = self.master_keys.write().await;
keys.insert(key_id.to_string(), protected);
Ok(())
}
pub async fn get_key(&self, key_id: &str) -> Result<Vec<u8>> {
let seal_key = self.seal_key
.ok_or_else(|| QsshError::Crypto("Vault not initialized".into()))?;
let mut keys = self.master_keys.write().await;
let protected = keys.get_mut(key_id)
.ok_or_else(|| QsshError::Crypto("Key not found".into()))?;
if protected.max_uses > 0 {
if protected.usage_count >= protected.max_uses {
return Err(QsshError::Crypto("Key usage limit exceeded".into()));
}
protected.usage_count += 1;
}
self.decrypt_key(&protected.encrypted_data, &seal_key)
}
pub async fn create_ratchet(&self, session_id: &str, root_key: &[u8]) -> Result<()> {
let ratchet = DoubleRatchet::new(root_key);
let mut ratchets = self.ratchets.write().await;
ratchets.insert(session_id.to_string(), ratchet);
Ok(())
}
pub async fn ratchet_forward(&self, session_id: &str) -> Result<[u8; 32]> {
let mut ratchets = self.ratchets.write().await;
let ratchet = ratchets.get_mut(session_id)
.ok_or_else(|| QsshError::Crypto("Ratchet session not found".into()))?;
Ok(ratchet.next_key())
}
pub async fn init_lamport_chain(&self, chain_id: &str, seed: &[u8], count: usize) -> Result<()> {
let mut hasher = Sha3_256::new();
hasher.update(b"QSSH-LAMPORT-SEED");
hasher.update(seed);
let mut root_seed = [0u8; 32];
root_seed.copy_from_slice(&hasher.finalize());
let mut keypairs = Vec::with_capacity(count);
for i in 0..count {
let mut key_seed = [0u8; 32];
let mut hasher = Sha3_256::new();
hasher.update(&root_seed);
hasher.update(&i.to_le_bytes());
key_seed.copy_from_slice(&hasher.finalize());
keypairs.push(LamportKeypair::generate(&key_seed));
}
let chain = LamportKeyChain {
current_index: 0,
keypairs,
root_seed,
};
let mut chains = self.lamport_keys.write().await;
chains.insert(chain_id.to_string(), chain);
Ok(())
}
pub async fn get_lamport_keypair(&self, chain_id: &str) -> Result<LamportKeypair> {
let mut chains = self.lamport_keys.write().await;
let chain = chains.get_mut(chain_id)
.ok_or_else(|| QsshError::Crypto("Lamport chain not found".into()))?;
if chain.current_index >= chain.keypairs.len() {
return Err(QsshError::Crypto("Lamport chain exhausted".into()));
}
let keypair = chain.keypairs[chain.current_index].clone();
chain.current_index += 1;
Ok(keypair)
}
fn encrypt_key(&self, key_data: &[u8], seal_key: &[u8; 32]) -> Result<Vec<u8>> {
use aes_gcm::{
aead::{Aead, KeyInit, OsRng},
Aes256Gcm, Nonce,
};
use rand::RngCore;
let cipher = Aes256Gcm::new_from_slice(seal_key)
.map_err(|e| QsshError::Crypto(format!("Cipher init failed: {}", e)))?;
let mut nonce_bytes = [0u8; 12];
OsRng.fill_bytes(&mut nonce_bytes);
let nonce = Nonce::from_slice(&nonce_bytes);
let ciphertext = cipher
.encrypt(nonce, key_data)
.map_err(|e| QsshError::Crypto(format!("Encryption failed: {}", e)))?;
let mut result = nonce_bytes.to_vec();
result.extend_from_slice(&ciphertext);
Ok(result)
}
fn decrypt_key(&self, encrypted: &[u8], seal_key: &[u8; 32]) -> Result<Vec<u8>> {
use aes_gcm::{
aead::{Aead, KeyInit},
Aes256Gcm, Nonce,
};
if encrypted.len() < 12 {
return Err(QsshError::Crypto("Invalid encrypted data".into()));
}
let (nonce_bytes, ciphertext) = encrypted.split_at(12);
let nonce = Nonce::from_slice(nonce_bytes);
let cipher = Aes256Gcm::new_from_slice(seal_key)
.map_err(|e| QsshError::Crypto(format!("Cipher init failed: {}", e)))?;
cipher
.decrypt(nonce, ciphertext)
.map_err(|e| QsshError::Crypto(format!("Decryption failed: {}", e)))
}
}
impl Drop for QuantumVault {
fn drop(&mut self) {
if let Some(ref mut key) = self.seal_key {
key.zeroize();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_vault_basic_operations() {
let mut vault = QuantumVault::new();
vault.init(b"test-master-key").await.expect("Failed to initialize vault");
let test_key = b"super-secret-key";
vault.store_key("test", test_key, KeyType::SessionKey, 0).await.expect("Failed to store key");
let retrieved = vault.get_key("test").await.expect("Failed to retrieve key");
assert_eq!(test_key.to_vec(), retrieved);
}
#[tokio::test]
async fn test_one_time_key() {
let mut vault = QuantumVault::new();
vault.init(b"test-master-key").await.expect("Failed to initialize vault");
let test_key = b"one-time-secret";
vault.store_key("once", test_key, KeyType::LamportSeed, 1).await.expect("Failed to store one-time key");
let retrieved = vault.get_key("once").await.expect("Failed to retrieve one-time key");
assert_eq!(test_key.to_vec(), retrieved);
assert!(vault.get_key("once").await.is_err());
}
}