#[cfg(not(feature = "std"))]
extern crate alloc;
#[cfg(not(feature = "std"))]
use alloc::vec::Vec;
#[cfg(feature = "std")]
use std::sync::Arc;
use core::sync::atomic::{AtomicU16, AtomicU64, AtomicU8, Ordering};
use crate::colony::common::current_timestamp_ms;
use crate::policy::{GatePolicy, TransitStatus};
use crate::utils::BasisPoints;
use crate::Frame;
#[cfg(feature = "x509")]
mod x509 {
pub use std::collections::HashMap;
pub use std::sync::Mutex;
pub use crate::colony::common::ClusterCommand;
pub use crate::crypto::x509::store::CertificateTrust;
pub use crate::der::Encode;
}
#[cfg(feature = "x509")]
use x509::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum CircuitState {
Closed = 0,
Open = 1,
HalfOpen = 2,
}
pub struct ClusterCircuitBreaker {
state: AtomicU8,
failures: AtomicU8,
opened_at: AtomicU64,
failure_threshold: u8,
cooldown_ms: u64,
}
impl ClusterCircuitBreaker {
pub fn new(failure_threshold: u8, cooldown_ms: u64) -> Self {
Self {
state: AtomicU8::new(CircuitState::Closed as u8),
failures: AtomicU8::new(0),
opened_at: AtomicU64::new(0),
failure_threshold,
cooldown_ms,
}
}
pub fn allow_request(&self) -> bool {
match self.state() {
CircuitState::Closed => true,
CircuitState::Open => {
let now = current_timestamp_ms();
let opened = self.opened_at.load(Ordering::Relaxed);
if now.saturating_sub(opened) < self.cooldown_ms {
return false;
}
self.state
.compare_exchange(
CircuitState::Open as u8,
CircuitState::HalfOpen as u8,
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
}
CircuitState::HalfOpen => true,
}
}
pub fn record_success(&self) {
self.failures.store(0, Ordering::Relaxed);
self.state.store(CircuitState::Closed as u8, Ordering::Release);
}
pub fn record_auth_failure(&self) {
if self.state() == CircuitState::HalfOpen {
self.trip();
return;
}
let previous = self
.failures
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |count| Some(count.saturating_add(1)))
.unwrap_or(u8::MAX);
if previous.saturating_add(1) >= self.failure_threshold {
self.trip();
}
}
fn trip(&self) {
self.opened_at.store(current_timestamp_ms(), Ordering::Relaxed);
self.state.store(CircuitState::Open as u8, Ordering::Release);
}
pub fn state(&self) -> CircuitState {
match self.state.load(Ordering::Acquire) {
0 => CircuitState::Closed,
1 => CircuitState::Open,
_ => CircuitState::HalfOpen,
}
}
pub fn is_open(&self) -> bool {
self.state() == CircuitState::Open
}
pub fn reset(&self) {
self.failures.store(0, Ordering::Relaxed);
self.state.store(CircuitState::Closed as u8, Ordering::Release);
}
}
#[cfg(feature = "x509")]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TrustVerification {
MissingSignature,
UnknownSigner,
Invalid,
Verified,
}
#[cfg(feature = "x509")]
pub fn verify_frame_signature(trust_store: &dyn CertificateTrust, frame: &Frame) -> TrustVerification {
let Some(signer_info) = frame.nonrepudiation.as_ref() else {
return TrustVerification::MissingSignature;
};
let Some(cert) = trust_store.find_by_signer_info(signer_info) else {
return TrustVerification::UnknownSigner;
};
let algorithm_oid = signer_info.signature_algorithm.oid;
let signature = signer_info.signature.as_bytes();
let Ok(public_key_der) = cert.tbs_certificate.subject_public_key_info.to_der() else {
return TrustVerification::Invalid;
};
let Ok(message) = frame.to_tbs() else {
return TrustVerification::Invalid;
};
match trust_store
.to_policy_ref()
.verify_signature(&algorithm_oid, &public_key_der, &message, signature)
{
Ok(()) => TrustVerification::Verified,
Err(_) => TrustVerification::Invalid,
}
}
#[cfg(feature = "x509")]
pub const REPLAY_GUARD_CAPACITY: usize = 1024;
#[cfg(feature = "x509")]
type SignerPartitions = HashMap<Vec<u8>, HashMap<Vec<u8>, u64>>;
#[cfg(feature = "x509")]
pub struct ReplayGuard {
seen: Mutex<SignerPartitions>,
window_ms: u64,
}
#[cfg(feature = "x509")]
impl ReplayGuard {
pub fn new(window_ms: u64) -> Self {
Self { seen: Mutex::new(HashMap::new()), window_ms }
}
pub fn is_fresh(&self, issued_at_ms: u64, now_ms: u64) -> bool {
now_ms.abs_diff(issued_at_ms) <= self.window_ms
}
pub fn check_and_insert(&self, signer: &[u8], signature: &[u8], now_ms: u64) -> bool {
let Ok(mut seen) = self.seen.lock() else {
return false;
};
seen.retain(|_, sigs| {
sigs.retain(|_, ts| now_ms.abs_diff(*ts) <= self.window_ms);
!sigs.is_empty()
});
if seen.values().any(|sigs| sigs.contains_key(signature)) {
return false;
}
let sigs = seen.entry(signer.to_vec()).or_default();
if sigs.len() >= REPLAY_GUARD_CAPACITY {
return false;
}
sigs.insert(signature.to_vec(), now_ms);
true
}
pub fn forget(&self, signature: &[u8]) {
let Ok(mut seen) = self.seen.lock() else {
return;
};
seen.retain(|_, sigs| {
sigs.remove(signature);
!sigs.is_empty()
});
}
}
#[cfg(feature = "x509")]
pub struct ClusterSecurityGate {
circuit_breaker: Arc<ClusterCircuitBreaker>,
trust_store: Arc<dyn CertificateTrust>,
replay_guard: Arc<ReplayGuard>,
}
#[cfg(feature = "x509")]
impl ClusterSecurityGate {
pub fn new(
circuit_breaker: Arc<ClusterCircuitBreaker>,
trust_store: Arc<dyn CertificateTrust>,
replay_guard: Arc<ReplayGuard>,
) -> Self {
Self { circuit_breaker, trust_store, replay_guard }
}
}
#[cfg(feature = "x509")]
impl GatePolicy for ClusterSecurityGate {
fn evaluate(&self, frame: &Frame) -> TransitStatus {
if !self.circuit_breaker.allow_request() {
return TransitStatus::Forbidden;
}
let Some(signer_info) = frame.nonrepudiation.as_ref() else {
return TransitStatus::Unauthorized;
};
if frame.integrity.is_none() {
return TransitStatus::Unauthorized;
}
match verify_frame_signature(self.trust_store.as_ref(), frame) {
TrustVerification::MissingSignature => return TransitStatus::Unauthorized,
TrustVerification::UnknownSigner => return TransitStatus::Forbidden,
TrustVerification::Invalid => {
self.circuit_breaker.record_auth_failure();
return TransitStatus::Forbidden;
}
TrustVerification::Verified => {}
}
let Ok(command) = crate::decode::<ClusterCommand>(&frame.message) else {
return TransitStatus::Forbidden;
};
let now = current_timestamp_ms();
if !self.replay_guard.is_fresh(command.issued_at_ms, now) {
return TransitStatus::Forbidden;
}
let Ok(signer_id) = signer_info.sid.to_der() else {
return TransitStatus::Forbidden;
};
if !self
.replay_guard
.check_and_insert(&signer_id, signer_info.signature.as_bytes(), now)
{
return TransitStatus::Forbidden;
}
self.circuit_breaker.record_success();
TransitStatus::Accepted
}
}
pub struct BackpressureGate {
utilization: Arc<AtomicU16>,
threshold: BasisPoints,
}
impl BackpressureGate {
pub fn new(utilization: Arc<AtomicU16>, threshold: BasisPoints) -> Self {
Self { utilization, threshold }
}
pub fn current_utilization(&self) -> BasisPoints {
BasisPoints::new_saturating(self.utilization.load(Ordering::Relaxed))
}
}
impl GatePolicy for BackpressureGate {
fn evaluate(&self, _frame: &Frame) -> TransitStatus {
let current = self.utilization.load(Ordering::Relaxed);
if current >= self.threshold.get() {
TransitStatus::Busy
} else {
TransitStatus::Accepted
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn breaker_trips_after_threshold() {
let breaker = ClusterCircuitBreaker::new(3, 60_000);
breaker.record_auth_failure();
breaker.record_auth_failure();
assert_eq!(breaker.state(), CircuitState::Closed);
breaker.record_auth_failure();
assert_eq!(breaker.state(), CircuitState::Open);
assert!(!breaker.allow_request());
}
#[test]
fn breaker_probe_success_closes() {
let breaker = ClusterCircuitBreaker::new(1, 0);
breaker.record_auth_failure();
assert!(breaker.allow_request());
assert_eq!(breaker.state(), CircuitState::HalfOpen);
breaker.record_success();
assert_eq!(breaker.state(), CircuitState::Closed);
}
#[test]
fn breaker_probe_failure_reopens() {
let breaker = ClusterCircuitBreaker::new(1, 0);
breaker.record_auth_failure();
assert!(breaker.allow_request());
assert_eq!(breaker.state(), CircuitState::HalfOpen);
breaker.record_auth_failure();
assert_eq!(breaker.state(), CircuitState::Open);
}
#[test]
fn breaker_reset_clears_state() {
let breaker = ClusterCircuitBreaker::new(1, 60_000);
breaker.record_auth_failure();
assert!(breaker.is_open());
breaker.reset();
assert_eq!(breaker.state(), CircuitState::Closed);
assert!(breaker.allow_request());
}
#[cfg(feature = "x509")]
#[test]
fn replay_guard_accepts_first_rejects_second() {
let guard = ReplayGuard::new(30_000);
assert!(guard.check_and_insert(b"signer-1", b"sig-a", 1_000));
assert!(!guard.check_and_insert(b"signer-1", b"sig-a", 2_000));
assert!(guard.check_and_insert(b"signer-1", b"sig-b", 2_000));
}
#[cfg(feature = "x509")]
#[test]
fn replay_guard_prunes_expired_entries() {
let guard = ReplayGuard::new(1_000);
assert!(guard.check_and_insert(b"signer-1", b"sig-a", 1_000));
assert!(guard.check_and_insert(b"signer-1", b"sig-a", 3_000));
}
#[cfg(feature = "x509")]
#[test]
fn replay_guard_prunes_future_dated_entries_after_clock_regression() {
let guard = ReplayGuard::new(1_000);
assert!(guard.check_and_insert(b"signer-1", b"sig-a", 10_000));
assert!(guard.check_and_insert(b"signer-1", b"sig-a", 5_000));
}
#[cfg(feature = "x509")]
#[test]
fn replay_guard_saturated_signer_does_not_block_others() {
let guard = ReplayGuard::new(30_000);
for i in 0..REPLAY_GUARD_CAPACITY {
assert!(guard.check_and_insert(b"signer-1", &i.to_be_bytes(), 1_000));
}
assert!(!guard.check_and_insert(b"signer-1", b"sig-overflow", 1_000));
assert!(guard.check_and_insert(b"signer-2", b"sig-a", 1_000));
}
#[cfg(feature = "x509")]
#[test]
fn replay_guard_rejects_replay_across_signer_partitions() {
let guard = ReplayGuard::new(30_000);
assert!(guard.check_and_insert(b"signer-1", b"sig-a", 1_000));
assert!(!guard.check_and_insert(b"signer-2", b"sig-a", 1_000));
}
#[cfg(feature = "x509")]
#[test]
fn replay_guard_forget_permits_retry() {
let guard = ReplayGuard::new(30_000);
assert!(guard.check_and_insert(b"signer-1", b"sig-a", 1_000));
guard.forget(b"sig-a");
assert!(guard.check_and_insert(b"signer-1", b"sig-a", 2_000));
}
#[cfg(feature = "x509")]
#[test]
fn replay_guard_freshness_window_is_bidirectional() {
let guard = ReplayGuard::new(1_000);
assert!(guard.is_fresh(9_500, 10_000));
assert!(guard.is_fresh(10_500, 10_000));
assert!(!guard.is_fresh(8_999, 10_000));
assert!(!guard.is_fresh(11_001, 10_000));
}
fn work_frame(priority: Option<crate::MessagePriority>) -> Result<Frame, crate::TightBeamError> {
use crate::builder::TypeBuilder;
let mut builder = crate::utils::compose(crate::Version::V2)
.with_id(b"work")
.with_order(0)
.with_message(crate::testing::TestMessage { content: "payload".into() });
if let Some(priority) = priority {
builder = builder.with_priority(priority);
}
builder.build()
}
#[test]
fn backpressure_gate_ignores_priority() -> Result<(), crate::TightBeamError> {
let utilization = Arc::new(AtomicU16::new(9_500));
let gate = BackpressureGate::new(utilization, BasisPoints::new_saturating(9_000));
let frame = work_frame(Some(crate::MessagePriority::NetworkControl))?;
assert_eq!(GatePolicy::evaluate(&gate, &frame), TransitStatus::Busy);
Ok(())
}
#[test]
fn backpressure_gate_accepts_below_threshold() -> Result<(), crate::TightBeamError> {
let utilization = Arc::new(AtomicU16::new(1_000));
let gate = BackpressureGate::new(utilization, BasisPoints::new_saturating(9_000));
let frame = work_frame(None)?;
assert_eq!(GatePolicy::evaluate(&gate, &frame), TransitStatus::Accepted);
Ok(())
}
}