use std::collections::HashMap;
use std::fmt;
use std::sync::{Arc, Mutex, Weak};
use sha2::{Digest, Sha384};
use tokio::sync::Notify;
use crate::binding::{connect_binding, status_statement, BindingError, SignedTbs};
use crate::node_key::{KeyError, NodeKey, Purpose};
pub const STATEMENT_EVERY_MS: i64 = 15 * 60 * 1000;
pub const STATEMENT_VALID_MS: i64 = 60 * 60 * 1000;
pub const CONNECT_BINDING_VALID_MS: i64 = 7 * 24 * 60 * 60 * 1000;
pub const CONNECT_ROTATE_EVERY_MS: i64 = 5 * 24 * 60 * 60 * 1000;
pub const ROTATION_MARGIN_MS: i64 = 24 * 60 * 60 * 1000;
const TOLERANCE_MS: i64 = 5 * 60 * 1000;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum IssuerError {
NotAnIdentityKey,
UnknownBinding,
NoConnectMaterial(String),
RotationOverdue { failures: u64, left_ms: i64 },
Key(KeyError),
Binding(BindingError),
}
impl fmt::Display for IssuerError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
IssuerError::NotAnIdentityKey => f.write_str("the issuer needs an identity key"),
IssuerError::UnknownBinding => f.write_str("no binding in force has that hash"),
IssuerError::NoConnectMaterial(why) => write!(f, "no CONNECT binding and status statement in force: {why}"),
IssuerError::RotationOverdue { failures, left_ms } => write!(
f,
"the CONNECT key has not rotated: {failures} failed rotations, {left_ms} ms left on its binding"
),
IssuerError::Key(e) => write!(f, "{e}"),
IssuerError::Binding(e) => write!(f, "{e}"),
}
}
}
impl std::error::Error for IssuerError {}
impl From<KeyError> for IssuerError {
fn from(e: KeyError) -> Self {
IssuerError::Key(e)
}
}
impl From<BindingError> for IssuerError {
fn from(e: BindingError) -> Self {
IssuerError::Binding(e)
}
}
#[derive(Debug, Clone)]
pub struct ConnectMaterial {
pub key: Arc<NodeKey>,
pub binding: SignedTbs,
pub status: SignedTbs,
}
pub type Clock = Box<dyn Fn() -> i64 + Send + Sync>;
struct StatedBinding {
key: Option<Arc<NodeKey>>,
binding: SignedTbs,
statement: SignedTbs,
bound_at: i64,
stated_at: i64,
not_after: i64,
}
impl StatedBinding {
fn rotation_due(&self, now: i64) -> bool {
now < self.bound_at - TOLERANCE_MS || now >= self.bound_at + CONNECT_ROTATE_EVERY_MS
}
fn restatement_due(&self, now: i64) -> bool {
now < self.stated_at - TOLERANCE_MS || now >= self.stated_at + STATEMENT_EVERY_MS
}
fn in_force(&self, now: i64) -> bool {
self.bound_at - TOLERANCE_MS <= now
&& now <= self.not_after
&& self.stated_at - TOLERANCE_MS <= now
&& now < self.stated_at + STATEMENT_VALID_MS
}
}
struct Slot {
newest: Mutex<(Option<SignedTbs>, bool)>,
notify: Notify,
}
struct State {
identity: Arc<NodeKey>,
clock: Clock,
current: [u8; 48],
bindings: HashMap<[u8; 48], StatedBinding>,
subscribers: HashMap<[u8; 48], Vec<Arc<Slot>>>,
rotation_failures: u64,
}
#[derive(Clone)]
pub struct StatementIssuer {
state: Arc<Mutex<State>>,
}
impl StatementIssuer {
pub fn new(identity: Arc<NodeKey>, clock: Clock) -> Result<StatementIssuer, IssuerError> {
if identity.purpose() != Purpose::Identity {
return Err(IssuerError::NotAnIdentityKey);
}
let now = clock();
let mut state = State {
identity,
clock,
current: [0; 48],
bindings: HashMap::new(),
subscribers: HashMap::new(),
rotation_failures: 0,
};
state.rotate_connect(now)?;
Ok(StatementIssuer {
state: Arc::new(Mutex::new(state)),
})
}
pub fn with_wall_clock(identity: Arc<NodeKey>) -> Result<StatementIssuer, IssuerError> {
StatementIssuer::new(identity, Box::new(|| crate::uuid_v7::now_ms() as i64))
}
pub fn connect_material(&self) -> Result<ConnectMaterial, IssuerError> {
let mut state = self.lock();
let now = (state.clock)();
let mut work = Ok(());
let due = state
.bindings
.get(&state.current)
.is_some_and(|b| b.rotation_due(now) || b.restatement_due(now));
if due {
work = state.tick(now);
}
let current = &state.bindings[&state.current];
match ¤t.key {
Some(key) if current.in_force(now) => Ok(ConnectMaterial {
key: key.clone(),
binding: current.binding.clone(),
status: current.statement.clone(),
}),
_ => Err(IssuerError::NoConnectMaterial(match work {
Err(e) => e.to_string(),
Ok(()) => "the current binding is out of force".into(),
})),
}
}
pub fn rotation_failures(&self) -> u64 {
self.lock().rotation_failures
}
pub fn subscribe(&self, binding: &SignedTbs) -> Result<StatementSubscription, IssuerError> {
let hash: [u8; 48] = Sha384::digest(&binding.tbs).into();
let mut state = self.lock();
let now = (state.clock)();
match state.bindings.get(&hash) {
Some(held) if held.not_after >= now => {}
_ => return Err(IssuerError::UnknownBinding),
}
let slot = Arc::new(Slot {
newest: Mutex::new((None, false)),
notify: Notify::new(),
});
state
.subscribers
.entry(hash)
.or_default()
.push(slot.clone());
Ok(StatementSubscription {
slot,
hash,
issuer: Arc::downgrade(&self.state),
})
}
pub fn tick(&self) -> Result<(), IssuerError> {
let mut state = self.lock();
let now = (state.clock)();
state.tick(now)
}
pub fn spawn_ticks(
&self,
on_error: impl Fn(IssuerError) + Send + 'static,
) -> tokio::task::JoinHandle<()> {
let weak = Arc::downgrade(&self.state);
tokio::spawn(async move {
let mut ticks =
tokio::time::interval(std::time::Duration::from_millis(STATEMENT_EVERY_MS as u64));
ticks.tick().await;
loop {
ticks.tick().await;
let Some(state) = weak.upgrade() else { return };
let issuer = StatementIssuer { state };
if let Err(e) = issuer.tick() {
on_error(e);
}
}
})
}
fn lock(&self) -> std::sync::MutexGuard<'_, State> {
self.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
}
impl State {
fn tick(&mut self, now: i64) -> Result<(), IssuerError> {
let reissued = self.reissue(now);
let mut rotated = Ok(());
if self
.bindings
.get(&self.current)
.is_some_and(|b| b.rotation_due(now))
{
rotated = self.rotate(now);
}
self.drop_expired(now);
reissued.and(rotated)
}
fn rotate(&mut self, now: i64) -> Result<(), IssuerError> {
match self.rotate_connect(now) {
Ok(()) => {
self.rotation_failures = 0;
Ok(())
}
Err(e) => {
self.rotation_failures += 1;
let left = self
.bindings
.get(&self.current)
.map_or(0, |b| b.not_after - now);
if left < ROTATION_MARGIN_MS {
return Err(IssuerError::RotationOverdue {
failures: self.rotation_failures,
left_ms: left,
});
}
Err(e)
}
}
}
fn reissue(&mut self, now: i64) -> Result<(), IssuerError> {
let mut first_error = Ok(());
for (hash, held) in self.bindings.iter_mut() {
if held.not_after < now {
continue;
}
match status_statement(&self.identity, &held.binding, now, now + STATEMENT_VALID_MS) {
Ok(statement) => {
held.statement = statement.clone();
held.stated_at = now;
for slot in self.subscribers.get(hash).into_iter().flatten() {
deliver(slot, statement.clone());
}
}
Err(e) => {
if first_error.is_ok() {
first_error = Err(e.into());
}
}
}
}
first_error
}
fn drop_expired(&mut self, now: i64) {
let expired: Vec<[u8; 48]> = self
.bindings
.iter()
.filter(|(_, b)| b.not_after < now)
.map(|(h, _)| *h)
.collect();
for hash in expired {
for slot in self.subscribers.remove(&hash).into_iter().flatten() {
close(&slot);
}
if hash != self.current {
self.bindings.remove(&hash);
}
}
}
fn rotate_connect(&mut self, now: i64) -> Result<(), IssuerError> {
let key = NodeKey::generate(Purpose::Connect, self.identity.profile())?;
let not_after = now + CONNECT_BINDING_VALID_MS;
let binding = connect_binding(&self.identity, &key.public_key(), now, not_after)?;
let statement = status_statement(&self.identity, &binding, now, now + STATEMENT_VALID_MS)?;
if let Some(previous) = self.bindings.get_mut(&self.current) {
previous.key = None;
}
let hash: [u8; 48] = Sha384::digest(&binding.tbs).into();
self.bindings.insert(
hash,
StatedBinding {
key: Some(Arc::new(key)),
binding,
statement,
bound_at: now,
stated_at: now,
not_after,
},
);
self.current = hash;
Ok(())
}
}
fn deliver(slot: &Slot, statement: SignedTbs) {
let mut newest = slot.newest.lock().unwrap_or_else(|p| p.into_inner());
newest.0 = Some(statement);
drop(newest);
slot.notify.notify_one();
}
fn close(slot: &Slot) {
slot.newest.lock().unwrap_or_else(|p| p.into_inner()).1 = true;
slot.notify.notify_one();
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SubscriptionEmpty {
Empty,
Closed,
}
pub struct StatementSubscription {
slot: Arc<Slot>,
hash: [u8; 48],
issuer: Weak<Mutex<State>>,
}
impl StatementSubscription {
pub fn try_recv(&mut self) -> Result<SignedTbs, SubscriptionEmpty> {
let mut newest = self.slot.newest.lock().unwrap_or_else(|p| p.into_inner());
match newest.0.take() {
Some(statement) => Ok(statement),
None if newest.1 => Err(SubscriptionEmpty::Closed),
None => Err(SubscriptionEmpty::Empty),
}
}
pub async fn recv(&mut self) -> Option<SignedTbs> {
let slot = self.slot.clone();
loop {
let notified = slot.notify.notified();
match self.try_recv() {
Ok(statement) => return Some(statement),
Err(SubscriptionEmpty::Closed) => return None,
Err(SubscriptionEmpty::Empty) => notified.await,
}
}
}
}
impl Drop for StatementSubscription {
fn drop(&mut self) {
let Some(state) = self.issuer.upgrade() else {
return;
};
let mut state = state.lock().unwrap_or_else(|p| p.into_inner());
if let Some(slots) = state.subscribers.get_mut(&self.hash) {
slots.retain(|s| !Arc::ptr_eq(s, &self.slot));
}
}
}