use crate::{QsshError, Result};
use crate::crypto::quantum_kem::QuantumKem;
use tokio::net::TcpStream;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
use std::sync::Arc;
use tokio::sync::{Mutex, RwLock};
use rand::RngCore;
use hmac::{Hmac, Mac};
use sha2::Sha256;
pub const QUANTUM_FRAME_SIZE: usize = 768;
#[repr(u8)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum QuantumFrameType {
Noise = 0x00, Handshake = 0x01, Data = 0x02, Control = 0x03, }
pub struct QuantumTransport {
reader: Arc<Mutex<OwnedReadHalf>>,
writer: Arc<Mutex<OwnedWriteHalf>>,
send_key: Arc<RwLock<Vec<u8>>>,
recv_key: Arc<RwLock<Vec<u8>>>,
send_sequence: Arc<Mutex<u64>>,
recv_sequence: Arc<Mutex<u64>>,
auth_key: Arc<[u8; 32]>,
stealth_config: StealthConfig,
}
#[derive(Debug, Clone)]
pub struct StealthConfig {
pub dummy_traffic: bool,
pub timing_obfuscation: bool,
pub min_frame_rate: u32,
pub max_padding: bool,
}
impl Default for StealthConfig {
fn default() -> Self {
Self {
dummy_traffic: true,
timing_obfuscation: true,
min_frame_rate: 10,
max_padding: true,
}
}
}
#[repr(C)]
pub struct QuantumFrame {
header: EncryptedHeader,
payload: [u8; 719],
mac: [u8; 32],
}
const _: () = assert!(std::mem::size_of::<QuantumFrame>() == QUANTUM_FRAME_SIZE);
const _: () = assert!(std::mem::size_of::<EncryptedHeader>() == 17);
#[repr(C)]
struct EncryptedHeader {
sequence: [u8; 8], timestamp: [u8; 8], frame_type: u8, }
const KEM_CIPHERTEXT_SIZE: usize = 897 + 64 + 16976;
const HANDSHAKE_HEADER_SIZE: usize = 4;
impl QuantumTransport {
pub async fn new(
stream: TcpStream,
kem: &QuantumKem,
peer_pk: &[u8],
is_client: bool,
) -> Result<Self> {
log::info!("Creating quantum-native transport (768-byte indistinguishable frames)");
let (mut reader, mut writer) = stream.into_split();
let (shared_secret, auth_key) = if is_client {
Self::client_kem_handshake(kem, peer_pk, &mut writer).await?
} else {
Self::server_kem_handshake(kem, peer_pk, &mut reader).await?
};
let mut send_key = [0u8; 32];
let mut recv_key = [0u8; 32];
if is_client {
send_key.copy_from_slice(&shared_secret[0..32]);
recv_key.copy_from_slice(&shared_secret[32..64]);
} else {
recv_key.copy_from_slice(&shared_secret[0..32]);
send_key.copy_from_slice(&shared_secret[32..64]);
}
Ok(Self {
reader: Arc::new(Mutex::new(reader)),
writer: Arc::new(Mutex::new(writer)),
send_key: Arc::new(RwLock::new(send_key.to_vec())),
recv_key: Arc::new(RwLock::new(recv_key.to_vec())),
send_sequence: Arc::new(Mutex::new(0)),
recv_sequence: Arc::new(Mutex::new(0)),
auth_key: Arc::new(auth_key),
stealth_config: StealthConfig::default(),
})
}
async fn client_kem_handshake(
kem: &QuantumKem,
server_pk: &[u8],
writer: &mut OwnedWriteHalf,
) -> Result<(Vec<u8>, [u8; 32])> {
log::debug!("Client: performing SPHINCS+/Falcon KEM handshake");
let (ciphertext, shared_secret) = kem.encapsulate(server_pk)?;
log::debug!(
"Client: KEM encapsulation successful, ciphertext: {} bytes, shared secret: {} bytes",
ciphertext.len(),
shared_secret.len()
);
let ciphertext_len = ciphertext.len() as u32;
writer.write_all(&ciphertext_len.to_be_bytes()).await
.map_err(QsshError::Io)?;
writer.write_all(&ciphertext).await
.map_err(QsshError::Io)?;
writer.flush().await
.map_err(QsshError::Io)?;
log::debug!("Client: KEM ciphertext sent to server");
let mut expanded = vec![0u8; 96];
Self::kdf(&shared_secret, b"QSSH-quantum-session", &mut expanded)?;
let mut auth_key = [0u8; 32];
auth_key.copy_from_slice(&expanded[64..96]);
Ok((expanded[0..64].to_vec(), auth_key))
}
async fn server_kem_handshake(
kem: &QuantumKem,
client_pk: &[u8],
reader: &mut OwnedReadHalf,
) -> Result<(Vec<u8>, [u8; 32])> {
log::debug!("Server: performing SPHINCS+/Falcon KEM handshake");
let mut len_bytes = [0u8; HANDSHAKE_HEADER_SIZE];
reader.read_exact(&mut len_bytes).await
.map_err(QsshError::Io)?;
let ciphertext_len = u32::from_be_bytes(len_bytes) as usize;
if ciphertext_len == 0 || ciphertext_len > KEM_CIPHERTEXT_SIZE + 1024 {
return Err(QsshError::Protocol(format!(
"Invalid KEM ciphertext size: {} (expected around {})",
ciphertext_len,
KEM_CIPHERTEXT_SIZE
)));
}
let mut ciphertext = vec![0u8; ciphertext_len];
reader.read_exact(&mut ciphertext).await
.map_err(QsshError::Io)?;
log::debug!(
"Server: received KEM ciphertext: {} bytes",
ciphertext.len()
);
let shared_secret = kem.decapsulate(&ciphertext, client_pk)?;
log::debug!(
"Server: KEM decapsulation successful, shared secret: {} bytes",
shared_secret.len()
);
let mut expanded = vec![0u8; 96];
Self::kdf(&shared_secret, b"QSSH-quantum-session", &mut expanded)?;
let mut auth_key = [0u8; 32];
auth_key.copy_from_slice(&expanded[64..96]);
log::debug!("Server: KEM handshake complete");
Ok((expanded[0..64].to_vec(), auth_key))
}
fn kdf(secret: &[u8], info: &[u8], output: &mut [u8]) -> Result<()> {
use hkdf::Hkdf;
use sha2::Sha256;
let hkdf = Hkdf::<Sha256>::new(None, secret);
hkdf.expand(info, output)
.map_err(|e| QsshError::Crypto(format!("KDF failed: {}", e)))?;
Ok(())
}
pub async fn send_frame(&self, frame_type: QuantumFrameType, data: &[u8]) -> Result<()> {
let sequence = {
let mut seq = self.send_sequence.lock().await;
let current = *seq;
*seq += 1;
current
};
let frame = self.build_frame(sequence, frame_type, data).await?;
let mut writer = self.writer.lock().await;
let frame_bytes = unsafe {
std::slice::from_raw_parts(
&frame as *const QuantumFrame as *const u8,
QUANTUM_FRAME_SIZE
)
};
writer.write_all(frame_bytes).await
.map_err(QsshError::Io)?;
writer.flush().await
.map_err(QsshError::Io)?;
log::trace!("Quantum frame sent: type={:?}, sequence={}", frame_type, sequence);
if self.stealth_config.dummy_traffic && rand::random::<f32>() < 0.3 {
self.send_dummy_frame_impl().await?;
}
Ok(())
}
pub async fn receive_frame(&self) -> Result<(QuantumFrameType, Vec<u8>)> {
let mut reader = self.reader.lock().await;
let mut frame_bytes = vec![0u8; QUANTUM_FRAME_SIZE];
reader.read_exact(&mut frame_bytes).await
.map_err(QsshError::Io)?;
let frame = unsafe {
std::ptr::read(frame_bytes.as_ptr() as *const QuantumFrame)
};
let (frame_type, payload) = self.verify_and_decrypt_frame(frame).await?;
if frame_type == QuantumFrameType::Noise {
log::trace!("Received noise frame, skipping");
return Box::pin(self.receive_frame()).await;
}
log::trace!("Quantum frame received: type={:?}, payload={} bytes", frame_type, payload.len());
Ok((frame_type, payload))
}
async fn build_frame(&self, sequence: u64, frame_type: QuantumFrameType, data: &[u8]) -> Result<QuantumFrame> {
let mut frame = QuantumFrame {
header: EncryptedHeader {
sequence: sequence.to_be_bytes(),
timestamp: Self::get_timestamp(),
frame_type: frame_type as u8,
},
payload: [0; 719],
mac: [0; 32],
};
let data_len = data.len().min(717);
frame.payload[0..2].copy_from_slice(&(data_len as u16).to_be_bytes());
frame.payload[2..2 + data_len].copy_from_slice(&data[..data_len]);
if self.stealth_config.max_padding {
rand::thread_rng().fill_bytes(&mut frame.payload[2 + data_len..]);
}
self.encrypt_header(&mut frame.header).await?;
frame.mac = self.calculate_mac(&frame).await?;
Ok(frame)
}
async fn verify_and_decrypt_frame(&self, mut frame: QuantumFrame) -> Result<(QuantumFrameType, Vec<u8>)> {
let calculated_mac = self.calculate_mac_for_verify(&frame).await?;
if calculated_mac != frame.mac {
return Err(QsshError::Crypto("Frame MAC verification failed".to_string()));
}
self.decrypt_header(&mut frame.header).await?;
let received_seq = u64::from_be_bytes(frame.header.sequence);
let expected_seq = {
let mut seq = self.recv_sequence.lock().await;
let current = *seq;
*seq += 1;
current
};
if received_seq != expected_seq {
return Err(QsshError::Protocol(format!(
"Invalid sequence: expected {}, got {}", expected_seq, received_seq
)));
}
let frame_type = match frame.header.frame_type {
0x00 => QuantumFrameType::Noise,
0x01 => QuantumFrameType::Handshake,
0x02 => QuantumFrameType::Data,
0x03 => QuantumFrameType::Control,
_ => return Err(QsshError::Protocol("Invalid frame type".to_string())),
};
let payload_len = u16::from_be_bytes([frame.payload[0], frame.payload[1]]) as usize;
if payload_len > 717 {
return Err(QsshError::Protocol("Invalid payload length".to_string()));
}
let payload = frame.payload[2..2 + payload_len].to_vec();
Ok((frame_type, payload))
}
async fn send_dummy_frame_impl(&self) -> Result<()> {
let sequence = {
let mut seq = self.send_sequence.lock().await;
let current = *seq;
*seq += 1;
current
};
let mut dummy_data = vec![0u8; rand::random::<usize>() % 500];
rand::thread_rng().fill_bytes(&mut dummy_data);
let frame = self.build_frame(sequence, QuantumFrameType::Noise, &dummy_data).await?;
let mut writer = self.writer.lock().await;
let frame_bytes = unsafe {
std::slice::from_raw_parts(
&frame as *const QuantumFrame as *const u8,
QUANTUM_FRAME_SIZE
)
};
writer.write_all(frame_bytes).await
.map_err(QsshError::Io)?;
writer.flush().await
.map_err(QsshError::Io)?;
log::trace!("Dummy frame sent: sequence={}", sequence);
Ok(())
}
async fn encrypt_header(&self, header: &mut EncryptedHeader) -> Result<()> {
let key = self.send_key.read().await;
for (i, byte) in header.sequence.iter_mut().enumerate() {
*byte ^= key[i % 32];
}
for (i, byte) in header.timestamp.iter_mut().enumerate() {
*byte ^= key[(i + 8) % 32];
}
header.frame_type ^= key[16];
Ok(())
}
async fn decrypt_header(&self, header: &mut EncryptedHeader) -> Result<()> {
let key = self.recv_key.read().await;
for (i, byte) in header.sequence.iter_mut().enumerate() {
*byte ^= key[i % 32];
}
for (i, byte) in header.timestamp.iter_mut().enumerate() {
*byte ^= key[(i + 8) % 32];
}
header.frame_type ^= key[16];
Ok(())
}
async fn calculate_mac(&self, frame: &QuantumFrame) -> Result<[u8; 32]> {
let mut mac = Hmac::<Sha256>::new_from_slice(&*self.auth_key)
.map_err(|e| QsshError::Crypto(format!("MAC creation failed: {}", e)))?;
mac.update(&frame.header.sequence);
mac.update(&frame.header.timestamp);
mac.update(&[frame.header.frame_type]);
mac.update(&frame.payload);
let result = mac.finalize();
let mut mac_bytes = [0u8; 32];
mac_bytes.copy_from_slice(result.into_bytes().as_slice());
Ok(mac_bytes)
}
async fn calculate_mac_for_verify(&self, frame: &QuantumFrame) -> Result<[u8; 32]> {
self.calculate_mac(frame).await
}
fn get_timestamp() -> [u8; 8] {
use std::time::{SystemTime, UNIX_EPOCH};
let micros = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_micros() as u64;
micros.to_be_bytes()
}
pub async fn close(&self) -> Result<()> {
let mut writer = self.writer.lock().await;
writer.shutdown().await
.map_err(QsshError::Io)?;
Ok(())
}
}
impl QuantumTransport {
pub async fn send_message<T: serde::Serialize>(&self, message: &T) -> Result<()> {
let data = bincode::serialize(message)
.map_err(|e| QsshError::Protocol(format!("Serialization failed: {}", e)))?;
self.send_frame(QuantumFrameType::Data, &data).await
}
pub async fn receive_message<T: for<'de> serde::Deserialize<'de>>(&self) -> Result<T> {
let (frame_type, data) = self.receive_frame().await?;
if frame_type != QuantumFrameType::Data {
return Err(QsshError::Protocol("Expected data frame".to_string()));
}
let message = bincode::deserialize(&data)
.map_err(|e| QsshError::Protocol(format!("Deserialization failed: {}", e)))?;
Ok(message)
}
}
#[async_trait::async_trait]
impl crate::transport::QsshTransport for QuantumTransport {
async fn send_message<T: serde::Serialize + Send + Sync>(&self, message: &T) -> Result<()> {
self.send_message(message).await
}
async fn receive_message<T: for<'de> serde::Deserialize<'de>>(&self) -> Result<T> {
self.receive_message().await
}
async fn close(&self) -> Result<()> {
self.close().await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_frame_size() {
assert_eq!(std::mem::size_of::<QuantumFrame>(), QUANTUM_FRAME_SIZE);
println!("✅ Quantum frame is exactly {} bytes", QUANTUM_FRAME_SIZE);
}
#[test]
fn test_frame_indistinguishability() {
let mut frame1 = QuantumFrame {
header: EncryptedHeader {
sequence: [1, 2, 3, 4, 5, 6, 7, 8],
timestamp: [9, 10, 11, 12, 13, 14, 15, 16],
frame_type: QuantumFrameType::Data as u8,
},
payload: [0; 719],
mac: [0; 32],
};
let mut frame2 = QuantumFrame {
header: EncryptedHeader {
sequence: [8, 7, 6, 5, 4, 3, 2, 1],
timestamp: [16, 15, 14, 13, 12, 11, 10, 9],
frame_type: QuantumFrameType::Noise as u8,
},
payload: [0; 719],
mac: [0; 32],
};
rand::thread_rng().fill_bytes(&mut frame1.payload);
rand::thread_rng().fill_bytes(&mut frame2.payload);
rand::thread_rng().fill_bytes(&mut frame1.mac);
rand::thread_rng().fill_bytes(&mut frame2.mac);
assert_eq!(std::mem::size_of_val(&frame1), std::mem::size_of_val(&frame2));
assert_eq!(std::mem::size_of_val(&frame1), QUANTUM_FRAME_SIZE);
println!("✅ All frames are indistinguishable (same size)");
}
}
#[cfg(kani)]
mod kani_proofs {
use super::*;
#[kani::proof]
fn proof_quantum_frame_size() {
assert_eq!(std::mem::size_of::<QuantumFrame>(), QUANTUM_FRAME_SIZE);
assert_eq!(QUANTUM_FRAME_SIZE, 768);
}
#[kani::proof]
fn proof_quantum_frame_alignment() {
assert_eq!(std::mem::align_of::<QuantumFrame>(), 1);
assert_eq!(std::mem::align_of::<EncryptedHeader>(), 1);
}
#[kani::proof]
fn proof_frame_roundtrip() {
let sequence: [u8; 8] = kani::any();
let timestamp: [u8; 8] = kani::any();
let frame_type: u8 = kani::any();
let mac: [u8; 32] = kani::any();
let frame = QuantumFrame {
header: EncryptedHeader {
sequence,
timestamp,
frame_type,
},
payload: [0u8; 719], mac,
};
let frame_bytes: &[u8] = unsafe {
std::slice::from_raw_parts(
&frame as *const QuantumFrame as *const u8,
QUANTUM_FRAME_SIZE,
)
};
assert_eq!(frame_bytes.len(), QUANTUM_FRAME_SIZE);
let frame_back: QuantumFrame = unsafe {
std::ptr::read(frame_bytes.as_ptr() as *const QuantumFrame)
};
assert_eq!(frame_back.header.sequence, sequence);
assert_eq!(frame_back.header.timestamp, timestamp);
assert_eq!(frame_back.header.frame_type, frame_type);
assert_eq!(frame_back.mac, mac);
}
#[kani::proof]
fn proof_build_frame_no_panic() {
let data_len: usize = kani::any();
kani::assume(data_len <= 719);
let mut payload = [0u8; 719];
let actual_len = data_len.min(717);
let len_bytes = (actual_len as u16).to_be_bytes();
payload[0] = len_bytes[0];
payload[1] = len_bytes[1];
assert!(2 + actual_len <= 719);
}
#[kani::proof]
fn proof_payload_bounds() {
let payload = [0u8; 719];
let byte0: u8 = kani::any();
let byte1: u8 = kani::any();
let payload_len = u16::from_be_bytes([byte0, byte1]) as usize;
if payload_len <= 717 {
assert!(2 + payload_len <= 719);
let _extracted = &payload[2..2 + payload_len];
}
}
#[kani::proof]
fn proof_frame_type_exhaustive() {
let frame_type_byte: u8 = kani::any();
let result = match frame_type_byte {
0x00 => Ok(QuantumFrameType::Noise),
0x01 => Ok(QuantumFrameType::Handshake),
0x02 => Ok(QuantumFrameType::Data),
0x03 => Ok(QuantumFrameType::Control),
_ => Err("Invalid frame type"),
};
if frame_type_byte <= 3 {
assert!(result.is_ok());
} else {
assert!(result.is_err());
}
}
#[kani::proof]
fn proof_payload_len_cast() {
let data_len: usize = kani::any();
kani::assume(data_len <= 719);
let actual_len = data_len.min(717);
let cast_result = actual_len as u16;
assert_eq!(cast_result as usize, actual_len);
}
#[kani::proof]
fn proof_sequence_no_overflow() {
let seq: u64 = kani::any();
kani::assume(seq < (1u64 << 48));
let next = seq + 1;
assert!(next > seq); }
#[kani::proof]
fn proof_sequence_from_be_bytes() {
let bytes: [u8; 8] = kani::any();
let seq = u64::from_be_bytes(bytes);
let roundtrip = seq.to_be_bytes();
assert_eq!(roundtrip, bytes);
}
}