use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher};
use std::sync::Arc;
use tokio::sync::{Mutex, MutexGuard};
const SHARDS: usize = 64;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum Chain {
Memory,
Event,
}
#[derive(Debug)]
pub struct ChainLocks {
shards: Vec<Arc<Mutex<()>>>,
}
impl Default for ChainLocks {
fn default() -> Self {
Self::new()
}
}
impl ChainLocks {
pub fn new() -> Self {
Self {
shards: (0..SHARDS).map(|_| Arc::new(Mutex::new(()))).collect(),
}
}
fn shard_for(&self, chain: Chain, agent_id: &str, thread_id: Option<&str>) -> usize {
let mut h = DefaultHasher::new();
chain.hash(&mut h);
agent_id.hash(&mut h);
if chain == Chain::Memory {
thread_id.is_some().hash(&mut h);
thread_id.unwrap_or("").hash(&mut h);
}
(h.finish() % SHARDS as u64) as usize
}
pub async fn lock(
&self,
chain: Chain,
agent_id: &str,
thread_id: Option<&str>,
) -> MutexGuard<'_, ()> {
self.shards[self.shard_for(chain, agent_id, thread_id)]
.lock()
.await
}
}
pub fn thread_key(chain: Chain, thread_id: Option<&str>) -> String {
match (chain, thread_id) {
(Chain::Event, _) | (_, None) => "-".to_string(),
(Chain::Memory, Some(t)) => format!("t:{t}"),
}
}
pub fn chain_name(chain: Chain) -> &'static str {
match chain {
Chain::Memory => "memory",
Chain::Event => "event",
}
}
pub fn advisory_key(chain: Chain, agent_id: &str, thread_id: Option<&str>) -> i64 {
use sha2::{Digest, Sha256};
let mut h = Sha256::new();
h.update(match chain {
Chain::Memory => b"mnemo:chain:memory\x00".as_slice(),
Chain::Event => b"mnemo:chain:event\x00".as_slice(),
});
h.update(agent_id.as_bytes());
h.update([0u8]);
if chain == Chain::Memory {
match thread_id {
Some(t) => {
h.update([1u8]);
h.update(t.as_bytes());
}
None => h.update([0u8]),
}
}
let d = h.finalize();
i64::from_be_bytes(d[..8].try_into().expect("sha256 yields 32 bytes"))
}
#[cfg(test)]
mod tests {
use super::*;
const PINNED_EVENT: i64 = -516657226202830256;
const PINNED_MEMORY: i64 = -7096874915296593842;
#[test]
fn the_two_chains_do_not_share_a_lock_for_the_same_key() {
let locks = ChainLocks::new();
assert_ne!(
locks.shard_for(Chain::Memory, "a", None),
locks.shard_for(Chain::Event, "a", None),
);
}
#[test]
fn a_thread_named_none_is_not_the_absent_memory_thread() {
let locks = ChainLocks::new();
assert_ne!(
locks.shard_for(Chain::Memory, "a", None),
locks.shard_for(Chain::Memory, "a", Some("None")),
);
}
#[test]
fn event_appends_for_one_agent_share_a_lock_across_threads() {
let locks = ChainLocks::new();
let base = locks.shard_for(Chain::Event, "a", None);
assert_eq!(base, locks.shard_for(Chain::Event, "a", Some("t1")));
assert_eq!(base, locks.shard_for(Chain::Event, "a", Some("t2")));
}
#[test]
fn the_advisory_key_is_pinned_to_a_literal() {
assert_eq!(
advisory_key(Chain::Event, "agent-a", None),
advisory_key(Chain::Event, "agent-a", Some("ignored")),
);
assert_ne!(
advisory_key(Chain::Memory, "agent-a", None),
advisory_key(Chain::Memory, "agent-a", Some("t")),
);
assert_ne!(
advisory_key(Chain::Memory, "agent-a", None),
advisory_key(Chain::Event, "agent-a", None),
);
assert_eq!(advisory_key(Chain::Event, "agent-a", None), PINNED_EVENT);
assert_eq!(advisory_key(Chain::Memory, "agent-a", None), PINNED_MEMORY);
}
#[test]
fn thread_keys_are_unambiguous() {
assert_eq!(thread_key(Chain::Memory, None), "-");
assert_eq!(thread_key(Chain::Memory, Some("a")), "t:a");
assert_ne!(
thread_key(Chain::Memory, Some("-")),
thread_key(Chain::Memory, None)
);
assert_eq!(thread_key(Chain::Event, Some("anything")), "-");
assert_eq!(thread_key(Chain::Event, None), "-");
}
#[tokio::test]
async fn the_same_key_serialises_and_releasing_frees_it() {
let locks = ChainLocks::new();
let held = locks.lock(Chain::Event, "agent-a", None).await;
let shard = locks.shard_for(Chain::Event, "agent-a", None);
assert!(locks.shards[shard].try_lock().is_err());
drop(held);
assert!(locks.shards[shard].try_lock().is_ok());
}
}