use crate::transport::handshake::error::HandshakeError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Phase {
Hello,
KeyExchange,
Finished,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AbortReason {
LocalPolicy,
Timeout(Phase),
PeerAbort,
Shutdown,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FailureKind {
ProtocolViolation,
ReplayDetected,
DowngradeAttempt,
CertificateInvalid,
SignatureInvalid,
IntegrityMismatch,
DerDecodeError,
KeyDerivationError,
UnsupportedAlgorithm,
InternalError,
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub enum ClientHandshakeState {
#[default]
Init,
HelloSent,
ServerHelloReceived,
KeyExchangeSent,
ServerFinishedReceived,
ClientFinishedSent,
Completed,
Aborted(AbortReason),
Failed(FailureKind),
}
impl ClientHandshakeState {
pub fn is_completed(&self) -> bool {
matches!(self, Self::Completed)
}
pub fn is_failed(&self) -> bool {
matches!(self, Self::Failed(_))
}
pub fn is_aborted(&self) -> bool {
matches!(self, Self::Aborted(_))
}
pub fn is_terminal(&self) -> bool {
self.is_completed() || self.is_failed() || self.is_aborted()
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub enum ServerHandshakeState {
#[default]
Init,
ClientHelloReceived,
ServerHelloSent,
KeyExchangeReceived,
ServerFinishedSent,
ClientFinishedReceived,
Completed,
Aborted(AbortReason),
Failed(FailureKind),
}
impl ServerHandshakeState {
pub fn is_completed(&self) -> bool {
matches!(self, Self::Completed)
}
pub fn is_failed(&self) -> bool {
matches!(self, Self::Failed(_))
}
pub fn is_aborted(&self) -> bool {
matches!(self, Self::Aborted(_))
}
pub fn is_terminal(&self) -> bool {
self.is_completed() || self.is_failed() || self.is_aborted()
}
}
#[cfg(not(feature = "std"))]
use alloc::collections::{BTreeMap as HashMap, VecDeque};
#[cfg(feature = "std")]
use std::collections::{HashMap, VecDeque};
#[derive(Debug)]
pub struct NonceReplaySet<const N: usize> {
seen: HashMap<[u8; N], usize>,
order: VecDeque<[u8; N]>,
cap: usize,
counter: usize,
}
impl<const N: usize> NonceReplaySet<N> {
pub fn new(cap: usize) -> Self {
Self {
#[cfg(feature = "std")]
seen: HashMap::with_capacity(cap),
#[cfg(not(feature = "std"))]
seen: HashMap::new(),
order: VecDeque::with_capacity(cap),
cap,
counter: 0,
}
}
pub fn insert_or_replay(&mut self, n: [u8; N]) -> bool {
if self.seen.contains_key(&n) {
return true; }
if self.seen.len() >= self.cap {
if let Some(oldest) = self.order.pop_front() {
self.seen.remove(&oldest);
}
}
self.seen.insert(n, self.counter);
self.order.push_back(n);
self.counter = self.counter.wrapping_add(1);
false }
pub fn clear(&mut self) {
self.seen.clear();
self.order.clear();
self.counter = 0;
}
pub fn len(&self) -> usize {
self.seen.len()
}
pub fn is_empty(&self) -> bool {
self.seen.is_empty()
}
}
#[derive(Debug, Default)]
pub struct HandshakeInvariant {
pub transcript_locked: bool,
pub aead_key_derived: bool,
pub finished_sent: bool,
}
impl HandshakeInvariant {
pub fn lock_transcript(&mut self) -> Result<bool, HandshakeError> {
if self.transcript_locked {
return Err(HandshakeError::TranscriptAlreadyLocked);
}
self.transcript_locked = true;
Ok(true)
}
pub fn derive_aead_once(&mut self) -> Result<bool, HandshakeError> {
if !self.transcript_locked {
return Err(HandshakeError::TranscriptNotLocked);
}
if self.aead_key_derived {
return Err(HandshakeError::AeadAlreadyDerived);
}
self.aead_key_derived = true;
Ok(true)
}
pub fn mark_finished_sent(&mut self) -> Result<bool, HandshakeError> {
if !self.transcript_locked {
return Err(HandshakeError::FinishedBeforeTranscriptLock);
}
if self.finished_sent {
return Err(HandshakeError::FinishedAlreadySent);
}
self.finished_sent = true;
Ok(true)
}
}
#[derive(Debug, Default)]
pub struct ClientStateMachine {
state: ClientHandshakeState,
}
impl ClientStateMachine {
pub fn state(&self) -> ClientHandshakeState {
self.state
}
fn can_transition(&self, to: ClientHandshakeState) -> bool {
use ClientHandshakeState::*;
match (self.state, to) {
(Init, HelloSent)
| (Init, KeyExchangeSent)
| (HelloSent, ServerHelloReceived)
| (ServerHelloReceived, KeyExchangeSent)
| (KeyExchangeSent, ServerFinishedReceived)
| (ServerFinishedReceived, ClientFinishedSent)
| (ClientFinishedSent, Completed)
| (KeyExchangeSent, Completed)
| (_, Aborted(_))
| (_, Failed(_)) => true,
_ => false,
}
}
pub fn transition(&mut self, to: ClientHandshakeState) -> Result<(), HandshakeError> {
if self.state.is_terminal() {
return Err(HandshakeError::InvalidState);
}
if self.can_transition(to) {
self.state = to;
Ok(())
} else {
Err(HandshakeError::InvalidState)
}
}
}
#[derive(Debug, Default)]
pub struct ServerStateMachine {
state: ServerHandshakeState,
}
impl ServerStateMachine {
pub fn state(&self) -> ServerHandshakeState {
self.state
}
fn can_transition(&self, to: ServerHandshakeState) -> bool {
use ServerHandshakeState::*;
match (self.state, to) {
(Init, ClientHelloReceived)
| (Init, KeyExchangeReceived)
| (ClientHelloReceived, ServerHelloSent)
| (ServerHelloSent, KeyExchangeReceived)
| (KeyExchangeReceived, ServerFinishedSent)
| (ServerFinishedSent, ClientFinishedReceived)
| (ClientFinishedReceived, Completed)
| (KeyExchangeReceived, Completed)
| (_, Aborted(_))
| (_, Failed(_)) => true,
_ => false,
}
}
pub fn transition(&mut self, to: ServerHandshakeState) -> Result<(), HandshakeError> {
if self.state.is_terminal() {
return Err(HandshakeError::InvalidState);
}
if self.can_transition(to) {
self.state = to;
Ok(())
} else {
Err(HandshakeError::InvalidState)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn client_linear_flow_ecies_short() {
let mut sm = ClientStateMachine::default();
assert_eq!(sm.state(), ClientHandshakeState::Init);
assert!(sm.transition(ClientHandshakeState::HelloSent).is_ok());
assert!(sm.transition(ClientHandshakeState::ServerHelloReceived).is_ok());
assert!(sm.transition(ClientHandshakeState::KeyExchangeSent).is_ok());
assert!(sm.transition(ClientHandshakeState::Completed).is_ok());
assert!(sm.state().is_completed());
}
#[test]
fn client_full_flow_cms() {
let mut sm = ClientStateMachine::default();
assert!(sm.transition(ClientHandshakeState::HelloSent).is_ok());
assert!(sm.transition(ClientHandshakeState::ServerHelloReceived).is_ok());
assert!(sm.transition(ClientHandshakeState::KeyExchangeSent).is_ok());
assert!(sm.transition(ClientHandshakeState::ServerFinishedReceived).is_ok());
assert!(sm.transition(ClientHandshakeState::ClientFinishedSent).is_ok());
assert!(sm.transition(ClientHandshakeState::Completed).is_ok());
}
#[test]
fn server_linear_flow_ecies_short() {
let mut sm = ServerStateMachine::default();
assert_eq!(sm.state(), ServerHandshakeState::Init);
assert!(sm.transition(ServerHandshakeState::ClientHelloReceived).is_ok());
assert!(sm.transition(ServerHandshakeState::ServerHelloSent).is_ok());
assert!(sm.transition(ServerHandshakeState::KeyExchangeReceived).is_ok());
assert!(sm.transition(ServerHandshakeState::Completed).is_ok());
assert!(sm.state().is_completed());
}
#[test]
fn server_full_flow_cms() {
let mut sm = ServerStateMachine::default();
assert!(sm.transition(ServerHandshakeState::KeyExchangeReceived).is_ok());
assert!(sm.transition(ServerHandshakeState::ServerFinishedSent).is_ok());
assert!(sm.transition(ServerHandshakeState::ClientFinishedReceived).is_ok());
assert!(sm.transition(ServerHandshakeState::Completed).is_ok());
}
#[test]
fn abort_and_failure_are_terminal() {
let mut sm = ClientStateMachine::default();
assert!(sm.transition(ClientHandshakeState::HelloSent).is_ok());
assert!(sm.transition(ClientHandshakeState::Aborted(AbortReason::PeerAbort)).is_ok());
assert!(sm.state().is_aborted());
assert!(sm.transition(ClientHandshakeState::ServerHelloReceived).is_err());
let mut sm2 = ServerStateMachine::default();
assert!(sm2
.transition(ServerHandshakeState::Failed(FailureKind::ProtocolViolation))
.is_ok());
assert!(sm2.state().is_failed());
}
#[test]
fn test_nonce_replay_detection() {
let mut set = NonceReplaySet::<32>::new(3);
let nonce1 = [1u8; 32];
let nonce2 = [2u8; 32];
assert!(!set.insert_or_replay(nonce1));
assert_eq!(set.len(), 1);
assert!(set.insert_or_replay(nonce1));
assert_eq!(set.len(), 1);
assert!(!set.insert_or_replay(nonce2));
assert_eq!(set.len(), 2);
}
#[test]
fn test_nonce_lru_eviction() {
let mut set = NonceReplaySet::<32>::new(3);
let nonce1 = [1u8; 32];
let nonce2 = [2u8; 32];
let nonce3 = [3u8; 32];
let nonce4 = [4u8; 32];
assert!(!set.insert_or_replay(nonce1));
assert!(!set.insert_or_replay(nonce2));
assert!(!set.insert_or_replay(nonce3));
assert_eq!(set.len(), 3);
assert!(!set.insert_or_replay(nonce4));
assert_eq!(set.len(), 3);
assert!(!set.insert_or_replay(nonce1)); assert_eq!(set.len(), 3);
assert!(set.insert_or_replay(nonce3)); assert!(set.insert_or_replay(nonce4)); assert!(set.insert_or_replay(nonce1)); }
#[test]
fn test_nonce_clear() {
let mut set = NonceReplaySet::<32>::new(10);
let nonce = [42u8; 32];
set.insert_or_replay(nonce);
assert_eq!(set.len(), 1);
set.clear();
assert_eq!(set.len(), 0);
assert!(set.is_empty());
assert!(!set.insert_or_replay(nonce));
}
#[test]
fn test_nonce_capacity_boundary() {
let mut set = NonceReplaySet::<32>::new(1);
let nonce1 = [1u8; 32];
let nonce2 = [2u8; 32];
assert!(!set.insert_or_replay(nonce1));
assert_eq!(set.len(), 1);
assert!(!set.insert_or_replay(nonce2));
assert_eq!(set.len(), 1);
assert!(!set.insert_or_replay(nonce1));
assert_eq!(set.len(), 1);
}
#[test]
fn test_nonce_64byte_ukm() {
let mut set = NonceReplaySet::<64>::new(3);
let ukm1 = [1u8; 64];
let ukm2 = [2u8; 64];
assert!(!set.insert_or_replay(ukm1));
assert_eq!(set.len(), 1);
assert!(set.insert_or_replay(ukm1));
assert_eq!(set.len(), 1);
assert!(!set.insert_or_replay(ukm2));
assert_eq!(set.len(), 2);
}
}