use std::cell::RefCell;
use std::collections::{HashMap, HashSet};
#[derive(Default)]
struct XmssIndexState {
committed: HashMap<Vec<u8>, u32>,
observed: Vec<(Vec<u8>, u32)>,
lamport_used: HashSet<Vec<u8>>,
lamport_observed: Vec<Vec<u8>>,
}
impl XmssIndexState {
fn check_and_observe(&mut self, pubkey: &[u8], index: u32) -> Result<(), String> {
if let Some(&max) = self.committed.get(pubkey) {
if index <= max {
return Err(format!(
"XMSS leaf index {index} reused or rolled back (last committed index {max})"
));
}
}
self.observed.push((pubkey.to_vec(), index));
Ok(())
}
fn check_and_observe_lamport(&mut self, pubkey: &[u8]) -> Result<(), String> {
if self.lamport_used.contains(pubkey) {
return Err(
"Lamport one-time key reused: this public key already signed an earlier entry"
.to_string(),
);
}
self.lamport_observed.push(pubkey.to_vec());
Ok(())
}
fn commit(&mut self) {
for (pubkey, index) in self.observed.drain(..) {
let slot = self.committed.entry(pubkey).or_insert(index);
if index > *slot {
*slot = index;
}
}
for pubkey in self.lamport_observed.drain(..) {
self.lamport_used.insert(pubkey);
}
}
fn discard(&mut self) {
self.observed.clear();
self.lamport_observed.clear();
}
}
thread_local! {
static XMSS_STATE: RefCell<Option<XmssIndexState>> = const { RefCell::new(None) };
}
pub struct XmssEnforcement {
prev: Option<XmssIndexState>,
}
impl XmssEnforcement {
#[must_use]
pub fn install() -> Self {
let prev = XMSS_STATE.with(|s| s.borrow_mut().replace(XmssIndexState::default()));
Self { prev }
}
pub fn commit_entry(&self) {
XMSS_STATE.with(|s| {
if let Some(state) = s.borrow_mut().as_mut() {
state.commit();
}
});
}
pub fn discard_entry(&self) {
XMSS_STATE.with(|s| {
if let Some(state) = s.borrow_mut().as_mut() {
state.discard();
}
});
}
}
impl Drop for XmssEnforcement {
fn drop(&mut self) {
XMSS_STATE.with(|s| *s.borrow_mut() = self.prev.take());
}
}
pub(crate) fn enforce_xmss_index(pubkey: &[u8], index: u32) -> Result<(), String> {
XMSS_STATE.with(|s| match s.borrow_mut().as_mut() {
Some(state) => state.check_and_observe(pubkey, index),
None => Ok(()),
})
}
pub(crate) fn enforce_lamport_once(pubkey: &[u8]) -> Result<(), String> {
XMSS_STATE.with(|s| match s.borrow_mut().as_mut() {
Some(state) => state.check_and_observe_lamport(pubkey),
None => Ok(()),
})
}
#[cfg(test)]
mod tests {
use super::*;
const K1: &[u8] = b"xmss-public-key-one";
const K2: &[u8] = b"xmss-public-key-two";
#[test]
fn monotonic_index_per_key() {
let mut s = XmssIndexState::default();
assert!(s.check_and_observe(K1, 0).is_ok());
s.commit();
assert!(s.check_and_observe(K1, 0).is_err());
s.discard();
assert!(s.check_and_observe(K1, 1).is_ok());
s.commit();
assert!(s.check_and_observe(K1, 0).is_err());
assert!(s.check_and_observe(K1, 1).is_err());
assert!(s.check_and_observe(K2, 0).is_ok());
}
#[test]
fn lamport_one_time_per_key() {
let mut s = XmssIndexState::default();
assert!(s.check_and_observe_lamport(K1).is_ok());
s.commit();
assert!(s.check_and_observe_lamport(K1).is_err());
assert!(s.check_and_observe_lamport(K2).is_ok());
s.commit();
assert!(s.check_and_observe_lamport(K2).is_err());
}
#[test]
fn lamport_reuse_within_entry_tolerated_until_commit() {
let mut s = XmssIndexState::default();
assert!(s.check_and_observe_lamport(K1).is_ok());
assert!(s.check_and_observe_lamport(K1).is_ok());
s.discard();
assert!(s.check_and_observe_lamport(K1).is_ok());
s.commit();
assert!(s.check_and_observe_lamport(K1).is_err());
}
#[test]
fn observations_only_count_after_commit() {
let mut s = XmssIndexState::default();
assert!(s.check_and_observe(K1, 5).is_ok());
s.discard();
assert!(s.check_and_observe(K1, 0).is_ok());
}
#[test]
fn same_index_within_one_entry_is_tolerated_until_commit() {
let mut s = XmssIndexState::default();
assert!(s.check_and_observe(K1, 3).is_ok());
assert!(s.check_and_observe(K1, 3).is_ok()); s.commit();
assert!(s.check_and_observe(K1, 3).is_err());
}
#[test]
fn enforce_is_noop_without_installed_guard() {
assert!(enforce_xmss_index(K1, 0).is_ok());
assert!(enforce_xmss_index(K1, 0).is_ok());
}
#[test]
fn raii_guard_installs_and_restores() {
assert!(enforce_xmss_index(K1, 7).is_ok()); {
let guard = XmssEnforcement::install();
assert!(enforce_xmss_index(K1, 0).is_ok());
guard.commit_entry();
assert!(enforce_xmss_index(K1, 0).is_err());
}
assert!(enforce_xmss_index(K1, 0).is_ok());
}
}