use crate::{Result, QsshError};
use crate::crypto::{PqKeyExchange, SessionKeyDerivation, SymmetricCrypto};
use crate::transport::{Transport, Message, RekeyMessage};
#[cfg(feature = "qkd")]
use crate::qkd::QkdClient;
use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tokio::sync::{Mutex, RwLock};
use tokio::time::{interval, Instant};
use serde::{Serialize, Deserialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct KeyRotationStats {
pub rotations_completed: u64,
pub rotations_failed: u64,
pub last_rotation_secs: Option<u64>, pub next_rotation_secs: Option<u64>, pub qkd_rotations: u64,
pub pqc_only_rotations: u64,
}
impl Default for KeyRotationStats {
fn default() -> Self {
Self {
rotations_completed: 0,
rotations_failed: 0,
last_rotation_secs: None,
next_rotation_secs: None,
qkd_rotations: 0,
pqc_only_rotations: 0,
}
}
}
pub struct KeyRotationManager {
transport: Arc<Transport>,
interval_seconds: u64,
stats: Arc<RwLock<KeyRotationStats>>,
#[cfg(feature = "qkd")]
qkd_client: Option<Arc<QkdClient>>,
is_server: bool,
current_keys: Arc<Mutex<SessionKeys>>,
}
struct SessionKeys {
client_write_key: [u8; 32],
server_write_key: [u8; 32],
generation: u64,
created_at: Instant,
}
impl KeyRotationManager {
pub fn new(
transport: Arc<Transport>,
interval_seconds: u64,
is_server: bool,
#[cfg(feature = "qkd")]
qkd_client: Option<Arc<QkdClient>>,
) -> Self {
let initial_keys = SessionKeys {
client_write_key: [0u8; 32], server_write_key: [0u8; 32],
generation: 0,
created_at: Instant::now(),
};
Self {
transport,
interval_seconds,
stats: Arc::new(RwLock::new(KeyRotationStats::default())),
#[cfg(feature = "qkd")]
qkd_client,
is_server,
current_keys: Arc::new(Mutex::new(initial_keys)),
}
}
pub async fn start_rotation_task(self: Arc<Self>) {
let mut rotation_interval = interval(Duration::from_secs(self.interval_seconds));
loop {
rotation_interval.tick().await;
{
let mut stats = self.stats.write().await;
let next_time = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs() + self.interval_seconds;
stats.next_rotation_secs = Some(next_time);
}
match self.rotate_keys().await {
Ok(_) => {
log::info!("Key rotation completed successfully");
let mut stats = self.stats.write().await;
stats.rotations_completed += 1;
stats.last_rotation_secs = Some(
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
);
}
Err(e) => {
log::error!("Key rotation failed: {}", e);
let mut stats = self.stats.write().await;
stats.rotations_failed += 1;
}
}
}
}
pub async fn rotate_keys(&self) -> Result<()> {
log::info!("Starting key rotation (generation {})",
self.current_keys.lock().await.generation);
if self.is_server {
self.rotate_as_server().await
} else {
self.rotate_as_client().await
}
}
async fn rotate_as_server(&self) -> Result<()> {
let new_kex = PqKeyExchange::new()?;
let (server_share, server_signature) = new_kex.create_key_share()?;
#[cfg(feature = "qkd")]
let qkd_key = if let Some(qkd_client) = &self.qkd_client {
match qkd_client.get_key(256).await {
Ok(key) => {
log::info!("Got QKD key for rotation: {} bytes", key.len());
Some(key)
}
Err(e) => {
log::warn!("QKD unavailable for rotation: {}", e);
None
}
}
} else {
None
};
#[cfg(not(feature = "qkd"))]
let qkd_key: Option<Vec<u8>> = None;
let rekey_msg = RekeyMessage {
sequence_number: self.current_keys.lock().await.generation + 1,
falcon_public_key: new_kex.public_bytes(),
key_share: server_share.clone(),
key_share_signature: server_signature,
qkd_proof: qkd_key.as_ref().map(|k| {
use sha3::{Sha3_256, Digest};
let mut hasher = Sha3_256::new();
hasher.update(b"QKD_REKEY_PROOF");
hasher.update(k);
hasher.finalize().to_vec()
}),
};
self.transport.send_message(&Message::Rekey(rekey_msg)).await?;
let client_rekey = self.wait_for_rekey_response().await?;
let client_share = new_kex.process_key_share(
&client_rekey.falcon_public_key,
&client_rekey.key_share,
&client_rekey.key_share_signature,
)?;
let mut shared_secret = new_kex.compute_shared_secret(
&server_share,
&client_share,
&[0u8; 32], &[0u8; 32],
);
#[cfg(feature = "qkd")]
if let Some(qkd_key) = qkd_key {
for (i, byte) in shared_secret.iter_mut().enumerate() {
if i < qkd_key.len() {
*byte ^= qkd_key[i];
}
}
let mut stats = self.stats.write().await;
stats.qkd_rotations += 1;
} else {
let mut stats = self.stats.write().await;
stats.pqc_only_rotations += 1;
}
#[cfg(not(feature = "qkd"))]
{
let mut stats = self.stats.write().await;
stats.pqc_only_rotations += 1;
}
let new_keys = SessionKeyDerivation::derive_keys(
&shared_secret,
&[0u8; 32],
&[0u8; 32],
)?;
self.update_transport_keys(&new_keys).await?;
let mut keys = self.current_keys.lock().await;
keys.client_write_key = new_keys.client_write_key;
keys.server_write_key = new_keys.server_write_key;
keys.generation += 1;
keys.created_at = Instant::now();
log::info!("Server completed key rotation to generation {}", keys.generation);
Ok(())
}
async fn rotate_as_client(&self) -> Result<()> {
let server_rekey = self.wait_for_rekey_request().await?;
let new_kex = PqKeyExchange::new()?;
let (client_share, client_signature) = new_kex.create_key_share()?;
#[cfg(feature = "qkd")]
let qkd_key = if let Some(qkd_proof) = &server_rekey.qkd_proof {
if let Some(qkd_client) = &self.qkd_client {
match qkd_client.verify_and_get_key(qkd_proof).await {
Ok(key) => {
log::info!("Client verified QKD proof for rotation");
Some(key)
}
Err(e) => {
log::warn!("Failed to verify QKD proof: {}", e);
None
}
}
} else {
None
}
} else {
None
};
#[cfg(not(feature = "qkd"))]
let qkd_key: Option<Vec<u8>> = None;
let server_share = new_kex.process_key_share(
&server_rekey.falcon_public_key,
&server_rekey.key_share,
&server_rekey.key_share_signature,
)?;
let client_rekey = RekeyMessage {
sequence_number: server_rekey.sequence_number,
falcon_public_key: new_kex.public_bytes(),
key_share: client_share.clone(),
key_share_signature: client_signature,
qkd_proof: None, };
self.transport.send_message(&Message::Rekey(client_rekey)).await?;
let mut shared_secret = new_kex.compute_shared_secret(
&client_share,
&server_share,
&[0u8; 32],
&[0u8; 32],
);
#[cfg(feature = "qkd")]
if let Some(qkd_key) = qkd_key {
for (i, byte) in shared_secret.iter_mut().enumerate() {
if i < qkd_key.len() {
*byte ^= qkd_key[i];
}
}
let mut stats = self.stats.write().await;
stats.qkd_rotations += 1;
} else {
let mut stats = self.stats.write().await;
stats.pqc_only_rotations += 1;
}
#[cfg(not(feature = "qkd"))]
{
let mut stats = self.stats.write().await;
stats.pqc_only_rotations += 1;
}
let new_keys = SessionKeyDerivation::derive_keys(
&shared_secret,
&[0u8; 32],
&[0u8; 32],
)?;
self.update_transport_keys(&new_keys).await?;
let mut keys = self.current_keys.lock().await;
keys.client_write_key = new_keys.client_write_key;
keys.server_write_key = new_keys.server_write_key;
keys.generation = server_rekey.sequence_number;
keys.created_at = Instant::now();
log::info!("Client completed key rotation to generation {}", keys.generation);
Ok(())
}
async fn wait_for_rekey_request(&self) -> Result<RekeyMessage> {
let timeout_duration = Duration::from_secs(30);
let start = Instant::now();
while start.elapsed() < timeout_duration {
match self.transport.receive_message::<Message>().await {
Ok(Message::Rekey(rekey)) => {
return Ok(rekey);
}
Ok(_) => {
tokio::time::sleep(Duration::from_millis(100)).await;
}
Err(e) => {
return Err(QsshError::Protocol(format!("Error waiting for rekey: {}", e)));
}
}
}
Err(QsshError::Protocol("Timeout waiting for rekey request".into()))
}
async fn wait_for_rekey_response(&self) -> Result<RekeyMessage> {
let timeout_duration = Duration::from_secs(30);
let start = Instant::now();
while start.elapsed() < timeout_duration {
match self.transport.receive_message::<Message>().await {
Ok(Message::Rekey(rekey)) => {
return Ok(rekey);
}
Ok(_) => {
tokio::time::sleep(Duration::from_millis(100)).await;
}
Err(e) => {
return Err(QsshError::Protocol(format!("Error waiting for rekey response: {}", e)));
}
}
}
Err(QsshError::Protocol("Timeout waiting for rekey response".into()))
}
async fn update_transport_keys(&self, keys: &crate::crypto::kdf::SessionKeys) -> Result<()> {
let send_crypto = if self.is_server {
SymmetricCrypto::from_shared_secret(&keys.server_write_key)?
} else {
SymmetricCrypto::from_shared_secret(&keys.client_write_key)?
};
let recv_crypto = if self.is_server {
SymmetricCrypto::from_shared_secret(&keys.client_write_key)?
} else {
SymmetricCrypto::from_shared_secret(&keys.server_write_key)?
};
self.transport.update_encryption(send_crypto, recv_crypto).await?;
Ok(())
}
pub async fn get_stats(&self) -> KeyRotationStats {
self.stats.read().await.clone()
}
pub async fn force_rotation(&self) -> Result<()> {
log::info!("Forcing immediate key rotation");
self.rotate_keys().await
}
pub async fn is_rotation_due(&self) -> bool {
let keys = self.current_keys.lock().await;
keys.created_at.elapsed() >= Duration::from_secs(self.interval_seconds)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rotation_stats() {
let stats = KeyRotationStats::default();
assert_eq!(stats.rotations_completed, 0);
assert_eq!(stats.rotations_failed, 0);
assert_eq!(stats.qkd_rotations, 0);
assert_eq!(stats.pqc_only_rotations, 0);
}
}