use std::collections::{HashMap, HashSet};
use std::str::FromStr;
use std::sync::{Arc, Condvar, Mutex, MutexGuard};
use base64::Engine;
use base64::engine::general_purpose::URL_SAFE_NO_PAD;
use hmac::{Hmac, Mac};
use kcode_k1_peering::K1Peering;
use kcode_k1_transaction::SubsystemId;
use kcode_k1_txn_ordering::{K1TxnOrdering, Subsystem, TxId};
use sha2::Sha256;
use zeroize::{Zeroize, Zeroizing};
const SUBSYSTEM_BYTES: [u8; 20] = *b"k1-invites-subsystem";
const CODE_BYTES: usize = 6;
const COMMITMENT_BYTES: usize = 32;
const ISSUE_BYTES: usize = 34;
const CONSUME_BYTES: usize = 46;
const VERSION: u8 = 1;
const ISSUE: u8 = 1;
const CONSUME: u8 = 2;
const DOMAIN: &[u8] = b"k1-invite-v1";
type HmacSha256 = Hmac<Sha256>;
pub struct InviteCode {
bytes: [u8; CODE_BYTES],
}
impl InviteCode {
fn from_bytes(bytes: [u8; CODE_BYTES]) -> Self {
Self { bytes }
}
#[must_use]
pub fn expose(&self) -> String {
URL_SAFE_NO_PAD.encode(self.bytes)
}
}
impl FromStr for InviteCode {
type Err = String;
fn from_str(text: &str) -> Result<Self, Self::Err> {
if text.len() != 8 || !text.is_ascii() {
return Err("invite code must be exactly eight URL-safe characters".to_owned());
}
let mut decoded = URL_SAFE_NO_PAD
.decode(text.as_bytes())
.map_err(|_| "invite code is not canonical URL-safe base64".to_owned())?;
if decoded.len() != CODE_BYTES {
decoded.zeroize();
return Err("invite code must decode to exactly six bytes".to_owned());
}
let mut bytes = [0; CODE_BYTES];
bytes.copy_from_slice(&decoded);
decoded.zeroize();
if URL_SAFE_NO_PAD.encode(bytes) != text {
bytes.zeroize();
return Err("invite code is not canonical URL-safe base64".to_owned());
}
Ok(Self::from_bytes(bytes))
}
}
impl Drop for InviteCode {
fn drop(&mut self) {
self.bytes.zeroize();
}
}
pub struct InviteVerifierKey {
bytes: Zeroizing<[u8; 32]>,
}
impl InviteVerifierKey {
#[must_use]
pub fn from_bytes(bytes: [u8; 32]) -> Self {
Self {
bytes: Zeroizing::new(bytes),
}
}
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub struct UserId(TxId);
impl UserId {
#[must_use]
pub fn as_tx_id(self) -> TxId {
self.0
}
}
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
struct Commitment([u8; COMMITMENT_BYTES]);
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum Record {
Issue(Commitment),
Consume {
issue_id: TxId,
commitment: Commitment,
},
}
#[derive(Clone, Copy, Eq, PartialEq)]
enum Phase {
Replaying,
Ready,
Unavailable,
}
struct Projection {
issue_id: TxId,
user_id: Option<UserId>,
}
struct State {
phase: Phase,
issued: HashMap<Commitment, Projection>,
pending_issues: HashSet<Commitment>,
pending_consumes: HashSet<Commitment>,
consume_failures: HashMap<Commitment, String>,
}
impl State {
fn new() -> Self {
Self {
phase: Phase::Replaying,
issued: HashMap::new(),
pending_issues: HashSet::new(),
pending_consumes: HashSet::new(),
consume_failures: HashMap::new(),
}
}
fn invalidate(&mut self) {
self.phase = Phase::Unavailable;
self.issued.clear();
self.pending_issues.clear();
self.pending_consumes.clear();
self.consume_failures.clear();
}
}
trait CodeGenerator: Send + Sync {
fn generate(&self) -> Result<[u8; CODE_BYTES], String>;
}
struct OsGenerator;
impl CodeGenerator for OsGenerator {
fn generate(&self) -> Result<[u8; CODE_BYTES], String> {
let mut bytes = [0; CODE_BYTES];
getrandom::fill(&mut bytes)
.map_err(|error| format!("invite randomness unavailable: {error}"))?;
Ok(bytes)
}
}
trait Submitter: Send + Sync {
fn submit(&self, subsystem: SubsystemId, payload: &[u8]) -> Result<TxId, String>;
}
impl Submitter for K1Peering {
fn submit(&self, subsystem: SubsystemId, payload: &[u8]) -> Result<TxId, String> {
self.submit_txn(subsystem, payload)
}
}
struct InviteInner {
key: InviteVerifierKey,
submitter: Arc<dyn Submitter>,
generator: Arc<dyn CodeGenerator>,
state: Mutex<State>,
changed: Condvar,
}
impl InviteInner {
fn new(
key: InviteVerifierKey,
submitter: Arc<dyn Submitter>,
generator: Arc<dyn CodeGenerator>,
) -> Self {
Self {
key,
submitter,
generator,
state: Mutex::new(State::new()),
changed: Condvar::new(),
}
}
fn lock(&self) -> Result<MutexGuard<'_, State>, String> {
self.state
.lock()
.map_err(|_| "invite subsystem synchronization unavailable".to_owned())
}
fn require_ready(state: &State) -> Result<(), String> {
(state.phase == Phase::Ready)
.then_some(())
.ok_or_else(|| "invite subsystem unavailable".to_owned())
}
fn finish_replay(&self) -> Result<(), String> {
let mut state = self.lock()?;
if state.phase != Phase::Replaying {
return Err("invite subsystem unavailable after replay".to_owned());
}
state.phase = Phase::Ready;
self.changed.notify_all();
Ok(())
}
fn invalidate(&self) {
let mut state = self
.state
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
state.invalidate();
self.changed.notify_all();
}
fn commitment(&self, code: &InviteCode) -> Commitment {
let mut mac = HmacSha256::new_from_slice(&self.key.bytes[..])
.expect("HMAC-SHA256 accepts a 32-byte key");
mac.update(DOMAIN);
mac.update(&code.bytes);
let output = mac.finalize().into_bytes();
let mut commitment = [0; COMMITMENT_BYTES];
commitment.copy_from_slice(&output);
Commitment(commitment)
}
fn create(&self) -> Result<(TxId, InviteCode), String> {
loop {
let code = InviteCode::from_bytes(self.generator.generate()?);
let commitment = self.commitment(&code);
{
let mut state = self.lock()?;
Self::require_ready(&state)?;
if state.issued.contains_key(&commitment)
|| !state.pending_issues.insert(commitment)
{
continue;
}
}
let result = self
.submitter
.submit(subsystem(), &encode_issue(commitment));
let mut state = self.lock()?;
state.pending_issues.remove(&commitment);
self.changed.notify_all();
Self::require_ready(&state)?;
let projected = state.issued.get(&commitment).map(|entry| entry.issue_id);
match (result, projected) {
(_, Some(issue_id)) => return Ok((issue_id, code)),
(Err(error), None) => return Err(error),
(Ok(_), None) => {
return Err("invite issue acknowledgement was not projected".to_owned());
}
}
}
}
fn consume(&self, code: &InviteCode) -> Result<UserId, String> {
let commitment = self.commitment(code);
let issue_id = loop {
let mut state = self.lock()?;
Self::require_ready(&state)?;
let Some(entry) = state.issued.get(&commitment) else {
return Err("invite is unavailable".to_owned());
};
if let Some(user_id) = entry.user_id {
return Ok(user_id);
}
let issue_id = entry.issue_id;
if let Some(error) = state.consume_failures.get(&commitment) {
return Err(error.clone());
}
if state.pending_consumes.insert(commitment) {
break issue_id;
}
state = self
.changed
.wait(state)
.map_err(|_| "invite subsystem synchronization unavailable".to_owned())?;
drop(state);
};
let result = self
.submitter
.submit(subsystem(), &encode_consume(issue_id, commitment));
let mut state = self.lock()?;
state.pending_consumes.remove(&commitment);
self.changed.notify_all();
Self::require_ready(&state)?;
let projected = state
.issued
.get(&commitment)
.filter(|entry| entry.issue_id == issue_id)
.and_then(|entry| entry.user_id);
if let Some(user_id) = projected {
return Ok(user_id);
}
let error = match result {
Ok(_) => "invite consume acknowledgement was not projected".to_owned(),
Err(error) => error,
};
state.consume_failures.insert(commitment, error.clone());
Err(error)
}
fn apply(&self, id: TxId, payload: &[u8]) -> Result<(), String> {
let record = match decode(payload) {
Ok(record) => record,
Err(error) => {
self.invalidate();
return Err(error);
}
};
let mut state = self.lock()?;
if state.phase == Phase::Unavailable {
return Err("invite subsystem unavailable".to_owned());
}
match record {
Record::Issue(commitment) => {
state.issued.entry(commitment).or_insert(Projection {
issue_id: id,
user_id: None,
});
}
Record::Consume {
issue_id,
commitment,
} => {
if let Some(entry) = state.issued.get_mut(&commitment)
&& entry.issue_id == issue_id
&& entry.user_id.is_none()
{
entry.user_id = Some(UserId(id));
state.consume_failures.remove(&commitment);
}
}
}
self.changed.notify_all();
Ok(())
}
}
impl Subsystem for InviteInner {
fn submit_txn(&self, id: TxId, payload: &[u8]) -> Result<(), String> {
self.apply(id, payload)
}
fn reorg(&self) -> Result<(), String> {
self.invalidate();
Ok(())
}
}
pub struct K1Invites {
inner: Arc<InviteInner>,
}
impl K1Invites {
pub fn open(
ordering: Arc<K1TxnOrdering>,
peering: Arc<K1Peering>,
verifier_key: InviteVerifierKey,
) -> Result<Self, String> {
let submitter: Arc<dyn Submitter> = peering;
let inner = Arc::new(InviteInner::new(
verifier_key,
submitter,
Arc::new(OsGenerator),
));
let handler: Arc<dyn Subsystem> = inner.clone();
if let Err(error) = ordering.register_subsystem(subsystem(), None, handler) {
inner.invalidate();
return Err(error);
}
inner.finish_replay()?;
Ok(Self { inner })
}
pub fn create(&self) -> Result<(TxId, InviteCode), String> {
self.inner.create()
}
pub fn consume(&self, code: &InviteCode) -> Result<UserId, String> {
self.inner.consume(code)
}
}
fn subsystem() -> SubsystemId {
SubsystemId::from_bytes(SUBSYSTEM_BYTES)
.expect("invite subsystem ID is valid fixed-width UTF-8")
}
fn encode_issue(commitment: Commitment) -> [u8; ISSUE_BYTES] {
let mut payload = [0; ISSUE_BYTES];
payload[0] = VERSION;
payload[1] = ISSUE;
payload[2..].copy_from_slice(&commitment.0);
payload
}
fn encode_consume(issue_id: TxId, commitment: Commitment) -> [u8; CONSUME_BYTES] {
let mut payload = [0; CONSUME_BYTES];
payload[0] = VERSION;
payload[1] = CONSUME;
payload[2..14].copy_from_slice(issue_id.as_bytes());
payload[14..].copy_from_slice(&commitment.0);
payload
}
fn decode(payload: &[u8]) -> Result<Record, String> {
if payload.len() < 2 {
return Err("malformed invite transaction header".to_owned());
}
if payload[0] != VERSION {
return Err("unsupported invite transaction version".to_owned());
}
match payload[1] {
ISSUE if payload.len() == ISSUE_BYTES => {
let mut commitment = [0; COMMITMENT_BYTES];
commitment.copy_from_slice(&payload[2..]);
Ok(Record::Issue(Commitment(commitment)))
}
ISSUE => Err("malformed invite issue transaction".to_owned()),
CONSUME if payload.len() == CONSUME_BYTES => {
let mut issue_id = [0; 12];
issue_id.copy_from_slice(&payload[2..14]);
let mut commitment = [0; COMMITMENT_BYTES];
commitment.copy_from_slice(&payload[14..]);
Ok(Record::Consume {
issue_id: TxId::from_bytes(issue_id),
commitment: Commitment(commitment),
})
}
CONSUME => Err("malformed invite consume transaction".to_owned()),
_ => Err("unknown invite transaction kind".to_owned()),
}
}
#[cfg(test)]
mod tests {
use super::*;
use kcode_k1_transaction::Transaction;
use std::collections::VecDeque;
use std::fs;
use std::path::PathBuf;
use std::sync::Barrier;
use std::sync::Weak;
use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::thread;
const KEY: u8 = 0x42;
static NEXT_ROOT: AtomicU64 = AtomicU64::new(0);
struct SequenceGenerator(Mutex<VecDeque<[u8; CODE_BYTES]>>);
impl SequenceGenerator {
fn new(values: Vec<[u8; CODE_BYTES]>) -> Self {
Self(Mutex::new(values.into()))
}
}
impl CodeGenerator for SequenceGenerator {
fn generate(&self) -> Result<[u8; CODE_BYTES], String> {
self.0
.lock()
.map_err(|_| "test generator unavailable".to_owned())?
.pop_front()
.ok_or_else(|| "test generator exhausted".to_owned())
}
}
#[derive(Clone)]
struct Committed {
id: TxId,
payload: Vec<u8>,
}
struct FakePeer {
next: AtomicU64,
attempts: AtomicUsize,
committed: Mutex<Vec<Committed>>,
handler: Mutex<Weak<InviteInner>>,
fail_before: AtomicBool,
fail_after: AtomicBool,
}
impl FakePeer {
fn new() -> Self {
Self {
next: AtomicU64::new(1),
attempts: AtomicUsize::new(0),
committed: Mutex::new(Vec::new()),
handler: Mutex::new(Weak::new()),
fail_before: AtomicBool::new(false),
fail_after: AtomicBool::new(false),
}
}
fn attach(&self, inner: &Arc<InviteInner>) {
*self.handler.lock().expect("handler lock") = Arc::downgrade(inner);
}
fn committed(&self) -> Vec<Committed> {
self.committed.lock().expect("commit lock").clone()
}
}
impl Submitter for FakePeer {
fn submit(&self, id: SubsystemId, payload: &[u8]) -> Result<TxId, String> {
assert_eq!(id, subsystem());
self.attempts.fetch_add(1, Ordering::SeqCst);
if self.fail_before.swap(false, Ordering::SeqCst) {
return Err("injected peering error".to_owned());
}
let transaction_id = tx_id(self.next.fetch_add(1, Ordering::SeqCst));
self.committed.lock().expect("commit lock").push(Committed {
id: transaction_id,
payload: payload.to_vec(),
});
self.handler
.lock()
.expect("handler lock")
.upgrade()
.ok_or_else(|| "test handler unavailable".to_owned())?
.apply(transaction_id, payload)?;
if self.fail_after.swap(false, Ordering::SeqCst) {
return Err("injected post-commit error".to_owned());
}
Ok(transaction_id)
}
}
fn tx_id(value: u64) -> TxId {
let mut bytes = [0; 12];
bytes[4..].copy_from_slice(&value.to_be_bytes());
TxId::from_bytes(bytes)
}
fn test_invites(peer: Arc<FakePeer>, key: u8, codes: Vec<[u8; CODE_BYTES]>) -> K1Invites {
let submitter: Arc<dyn Submitter> = peer.clone();
let inner = Arc::new(InviteInner::new(
InviteVerifierKey::from_bytes([key; 32]),
submitter,
Arc::new(SequenceGenerator::new(codes)),
));
peer.attach(&inner);
inner.finish_replay().unwrap();
K1Invites { inner }
}
fn replay(peer: Arc<FakePeer>, key: u8) -> K1Invites {
let invites = test_invites(peer.clone(), key, Vec::new());
{
let mut state = invites.inner.lock().unwrap();
state.phase = Phase::Replaying;
}
for record in peer.committed() {
invites.inner.apply(record.id, &record.payload).unwrap();
}
invites.inner.finish_replay().unwrap();
invites
}
#[test]
fn strict_code_and_wire_codecs() {
let code = InviteCode::from_bytes([0, 1, 2, 3, 4, 5]);
let text = code.expose();
assert_eq!(text, "AAECAwQF");
assert_eq!(text.parse::<InviteCode>().unwrap().bytes, code.bytes);
for bad in [
"",
"AAAAAAA",
"AAAAAAAAA",
"AAAAAAA=",
"AAAAAAA!",
"////////",
] {
assert!(bad.parse::<InviteCode>().is_err());
}
let commitment = Commitment([7; COMMITMENT_BYTES]);
let issue_id = tx_id(9);
assert_eq!(
decode(&encode_issue(commitment)),
Ok(Record::Issue(commitment))
);
assert_eq!(
decode(&encode_consume(issue_id, commitment)),
Ok(Record::Consume {
issue_id,
commitment,
})
);
let mut malformed = vec![vec![], vec![VERSION], vec![2, ISSUE], vec![VERSION, 9]];
malformed.push(encode_issue(commitment)[..ISSUE_BYTES - 1].to_vec());
malformed.push([encode_issue(commitment).as_slice(), &[0]].concat());
malformed.push(encode_consume(issue_id, commitment)[..CONSUME_BYTES - 1].to_vec());
malformed.push([encode_consume(issue_id, commitment).as_slice(), &[0]].concat());
assert!(malformed.iter().all(|payload| decode(payload).is_err()));
}
#[test]
fn creation_collision_consumption_and_wrong_key_are_deterministic() {
let peer = Arc::new(FakePeer::new());
let first = [1, 2, 3, 4, 5, 6];
let replacement = [7, 8, 9, 10, 11, 12];
let invites = test_invites(peer.clone(), KEY, vec![first, first, replacement]);
let (first_id, first_code) = invites.create().unwrap();
let (second_id, second_code) = invites.create().unwrap();
assert_ne!(first_id, second_id);
assert_eq!(first_code.bytes, first);
assert_eq!(second_code.bytes, replacement);
assert_eq!(peer.committed().len(), 2);
let first_payload = &peer.committed()[0].payload;
assert_eq!(
first_payload,
&encode_issue(invites.inner.commitment(&first_code))
);
assert!(!first_payload.windows(CODE_BYTES).any(|part| part == first));
let wrong = InviteCode::from_bytes([99; CODE_BYTES]);
assert!(invites.consume(&wrong).is_err());
assert_eq!(peer.committed().len(), 2);
let user = invites.consume(&first_code).unwrap();
assert_eq!(user.as_tx_id(), peer.committed()[2].id);
let attempts = peer.attempts.load(Ordering::SeqCst);
assert_eq!(invites.consume(&first_code).unwrap(), user);
assert_eq!(peer.attempts.load(Ordering::SeqCst), attempts);
let wrong_key = replay(peer.clone(), KEY + 1);
assert!(wrong_key.consume(&second_code).is_err());
}
#[test]
fn concurrent_consumers_share_one_commit_and_errors_are_not_retried() {
let peer = Arc::new(FakePeer::new());
let invites = Arc::new(test_invites(peer.clone(), KEY, vec![[3; CODE_BYTES]]));
peer.fail_after.store(true, Ordering::SeqCst);
let (_, code) = invites.create().unwrap();
let text = code.expose();
let barrier = Arc::new(Barrier::new(12));
let handles: Vec<_> = (0..12)
.map(|_| {
let invites = invites.clone();
let barrier = barrier.clone();
let text = text.clone();
thread::spawn(move || {
let code: InviteCode = text.parse().unwrap();
barrier.wait();
invites.consume(&code)
})
})
.collect();
let users: Vec<_> = handles
.into_iter()
.map(|handle| handle.join().unwrap().unwrap())
.collect();
assert!(users.iter().all(|user| *user == users[0]));
assert_eq!(peer.committed().len(), 2);
let failed_peer = Arc::new(FakePeer::new());
let failed = test_invites(failed_peer.clone(), KEY, vec![[4; CODE_BYTES]]);
let (_, failed_code) = failed.create().unwrap();
failed_peer.fail_before.store(true, Ordering::SeqCst);
let before = failed_peer.attempts.load(Ordering::SeqCst);
let error = failed.consume(&failed_code).unwrap_err();
assert_eq!(failed.consume(&failed_code).unwrap_err(), error);
assert_eq!(failed_peer.attempts.load(Ordering::SeqCst), before + 1);
}
#[test]
fn replay_and_reorg_are_fail_closed() {
let peer = Arc::new(FakePeer::new());
let invites = test_invites(peer.clone(), KEY, vec![[5; CODE_BYTES]]);
let (_, code) = invites.create().unwrap();
let user = invites.consume(&code).unwrap();
let restarted = replay(peer.clone(), KEY);
assert_eq!(restarted.consume(&code).unwrap(), user);
let before = peer.committed().len();
restarted.inner.reorg().unwrap();
assert!(restarted.consume(&code).is_err());
assert!(restarted.create().is_err());
assert_eq!(peer.committed().len(), before);
let reopened = replay(peer, KEY);
assert_eq!(reopened.consume(&code).unwrap(), user);
}
struct TempRoots(PathBuf);
impl TempRoots {
fn new(label: &str) -> Self {
let number = NEXT_ROOT.fetch_add(1, Ordering::Relaxed);
let root = std::env::temp_dir().join(format!(
"kcode-k1-invites-{}-{number}-{label}",
std::process::id()
));
let _ = fs::remove_dir_all(&root);
Self(root)
}
fn ordering(&self) -> PathBuf {
self.0.join("ordering")
}
fn peering(&self) -> PathBuf {
self.0.join("peering")
}
}
impl Drop for TempRoots {
fn drop(&mut self) {
let _ = fs::remove_dir_all(&self.0);
}
}
fn open_real(roots: &TempRoots) -> (Arc<K1TxnOrdering>, Arc<K1Peering>, K1Invites) {
let ordering = Arc::new(K1TxnOrdering::open(&roots.ordering()).unwrap());
let peering = Arc::new(K1Peering::open(&roots.peering(), ordering.clone()).unwrap());
let invites = K1Invites::open(
ordering.clone(),
peering.clone(),
InviteVerifierKey::from_bytes([KEY; 32]),
)
.unwrap();
(ordering, peering, invites)
}
#[test]
fn real_stack_receipts_restart_and_persistence_boundary() {
let roots = TempRoots::new("restart");
let (ordering, peering, invites) = open_real(&roots);
let (issue_id, code) = invites.create().unwrap();
assert_eq!(ordering.tip(), Some(issue_id));
let bytes = ordering.get_txn(issue_id).unwrap().unwrap();
let transaction = Transaction::parse(&bytes).unwrap();
assert_eq!(transaction.subsystem(), subsystem());
assert_eq!(
decode(transaction.payload()),
Ok(Record::Issue(invites.inner.commitment(&code)))
);
let user = invites.consume(&code).unwrap();
assert_eq!(ordering.tip(), Some(user.as_tx_id()));
let tip = ordering.tip();
assert_eq!(invites.consume(&code).unwrap(), user);
assert_eq!(ordering.tip(), tip);
let text = code.expose();
drop(invites);
drop(peering);
drop(ordering);
let (ordering, peering, invites) = open_real(&roots);
let code: InviteCode = text.parse().unwrap();
assert_eq!(invites.consume(&code).unwrap(), user);
assert_eq!(ordering.tip(), tip);
let mut entries: Vec<_> = fs::read_dir(&roots.0)
.unwrap()
.map(|entry| entry.unwrap().file_name())
.collect();
entries.sort();
assert_eq!(entries, vec!["ordering", "peering"]);
drop(invites);
drop(peering);
drop(ordering);
}
struct Noop;
impl Subsystem for Noop {
fn submit_txn(&self, _id: TxId, _payload: &[u8]) -> Result<(), String> {
Ok(())
}
fn reorg(&self) -> Result<(), String> {
Ok(())
}
}
#[test]
fn malformed_canonical_payload_faults_only_invites() {
let roots = TempRoots::new("isolation");
let (ordering, peering, invites) = open_real(&roots);
let other = SubsystemId::from_bytes([b'o'; 20]).unwrap();
ordering
.register_subsystem(other, None, Arc::new(Noop))
.unwrap();
let result = ordering.submit_local_txn(
1,
[9; 32],
subsystem(),
&[VERSION],
|_| Ok([9; 64]),
|_| Ok(()),
);
assert!(result.is_err());
assert!(invites.create().is_err());
let other_id = peering.submit_txn(other, b"still available").unwrap();
assert_eq!(ordering.tip(), Some(other_id));
}
}