use std::collections::HashMap;
use std::sync::{Arc, Mutex as StdMutex, Weak};
use tokio::sync::{Mutex as AsyncMutex, OwnedMutexGuard};
use vta_sdk::webvh::WebvhDidRecord;
#[derive(Clone, Default)]
pub struct DidUpdateLocks {
inner: Arc<StdMutex<HashMap<String, Weak<AsyncMutex<()>>>>>,
}
impl DidUpdateLocks {
pub async fn acquire(&self, did: &str) -> OwnedMutexGuard<()> {
let lock = {
let mut locks = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
match locks.get(did).and_then(Weak::upgrade) {
Some(live) => live,
None => {
locks.retain(|_, weak| weak.strong_count() > 0);
let fresh = Arc::new(AsyncMutex::new(()));
locks.insert(did.to_string(), Arc::downgrade(&fresh));
fresh
}
}
};
lock.lock_owned().await
}
}
pub static DID_UPDATE_LOCKS: std::sync::LazyLock<DidUpdateLocks> =
std::sync::LazyLock::new(DidUpdateLocks::default);
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RecordSnapshot {
did: String,
log_entry_count: u32,
updated_at: chrono::DateTime<chrono::Utc>,
server_id: String,
}
impl RecordSnapshot {
pub fn capture(record: &WebvhDidRecord) -> Self {
Self {
did: record.did.clone(),
log_entry_count: record.log_entry_count,
updated_at: record.updated_at,
server_id: record.server_id.clone(),
}
}
pub fn assert_unchanged(&self, current: &WebvhDidRecord) -> Result<(), RaceDetected> {
debug_assert_eq!(
self.did, current.did,
"RecordSnapshot::assert_unchanged comparing different DIDs — caller bug"
);
if self.log_entry_count != current.log_entry_count {
return Err(RaceDetected::LogEntryCountChanged {
did: self.did.clone(),
expected: self.log_entry_count,
current: current.log_entry_count,
});
}
if self.updated_at != current.updated_at {
return Err(RaceDetected::UpdatedAtChanged {
did: self.did.clone(),
expected: self.updated_at,
current: current.updated_at,
});
}
if self.server_id != current.server_id {
return Err(RaceDetected::ServerIdChanged {
did: self.did.clone(),
expected: self.server_id.clone(),
current: current.server_id.clone(),
});
}
Ok(())
}
}
#[allow(clippy::enum_variant_names)]
#[derive(Debug, thiserror::Error)]
pub enum RaceDetected {
#[error(
"DID `{did}` log_entry_count changed concurrently \
(expected {expected}, got {current}) — another caller appended a log entry"
)]
LogEntryCountChanged {
did: String,
expected: u32,
current: u32,
},
#[error(
"DID `{did}` was modified concurrently \
(record updated_at moved from {expected} to {current})"
)]
UpdatedAtChanged {
did: String,
expected: chrono::DateTime<chrono::Utc>,
current: chrono::DateTime<chrono::Utc>,
},
#[error(
"DID `{did}` server_id changed concurrently \
(`{expected}` → `{current}`) — another caller registered or moved this DID"
)]
ServerIdChanged {
did: String,
expected: String,
current: String,
},
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::sync::oneshot;
fn record(did: &str, count: u32, ts: i64, server: &str) -> WebvhDidRecord {
WebvhDidRecord {
did: did.into(),
server_id: server.into(),
mnemonic: "irrelevant".into(),
scid: "scid".into(),
context_id: "vta".into(),
portable: true,
log_entry_count: count,
pre_rotation_count: 0,
next_fragment_id: 0,
created_at: chrono::Utc::now(),
updated_at: chrono::DateTime::<chrono::Utc>::from_timestamp(ts, 0).unwrap(),
}
}
#[test]
fn unchanged_record_passes() {
let r = record("did:webvh:foo", 1, 1_000_000, "serverless");
let snap = RecordSnapshot::capture(&r);
snap.assert_unchanged(&r).expect("identity case must pass");
}
#[test]
fn log_entry_count_change_detected() {
let before = record("did:webvh:foo", 1, 1_000_000, "serverless");
let after = record("did:webvh:foo", 2, 1_000_000, "serverless");
let snap = RecordSnapshot::capture(&before);
let err = snap.assert_unchanged(&after).unwrap_err();
assert!(
matches!(
err,
RaceDetected::LogEntryCountChanged {
expected: 1,
current: 2,
..
}
),
"got {err:?}"
);
}
#[test]
fn updated_at_change_detected() {
let before = record("did:webvh:foo", 1, 1_000_000, "serverless");
let after = record("did:webvh:foo", 1, 1_000_001, "serverless");
let snap = RecordSnapshot::capture(&before);
let err = snap.assert_unchanged(&after).unwrap_err();
assert!(
matches!(err, RaceDetected::UpdatedAtChanged { .. }),
"got {err:?}"
);
}
#[test]
fn server_id_change_detected() {
let before = record("did:webvh:foo", 1, 1_000_000, "serverless");
let after = record("did:webvh:foo", 1, 1_000_000, "webvh-prod");
let snap = RecordSnapshot::capture(&before);
let err = snap.assert_unchanged(&after).unwrap_err();
assert!(
matches!(err, RaceDetected::ServerIdChanged { ref expected, ref current, .. }
if expected == "serverless" && current == "webvh-prod"),
"got {err:?}"
);
}
#[test]
fn unrelated_field_changes_do_not_trip_assertion() {
let before = record("did:webvh:foo", 1, 1_000_000, "serverless");
let mut after = before.clone();
after.mnemonic = "rotated-by-this-op".into();
after.next_fragment_id = 42;
after.pre_rotation_count = 3;
let snap = RecordSnapshot::capture(&before);
snap.assert_unchanged(&after)
.expect("only log_entry_count, updated_at, server_id are version-vector fields");
}
#[tokio::test]
async fn did_update_lock_serializes_same_did() {
let locks = DidUpdateLocks::default();
let first = locks.acquire("did:webvh:scid:example.com:agent").await;
let (started_tx, started_rx) = oneshot::channel();
let (acquired_tx, mut acquired_rx) = oneshot::channel();
let waiting_locks = locks.clone();
let waiting = tokio::spawn(async move {
let _ = started_tx.send(());
let _second = waiting_locks
.acquire("did:webvh:scid:example.com:agent")
.await;
let _ = acquired_tx.send(());
});
started_rx.await.expect("waiting task started");
assert!(
tokio::time::timeout(std::time::Duration::from_millis(20), &mut acquired_rx)
.await
.is_err(),
"the second update must wait for the first holder"
);
drop(first);
acquired_rx
.await
.expect("second update acquired after release");
waiting.await.expect("waiting task joined");
}
#[tokio::test]
async fn did_update_locks_are_reclaimed_when_idle() {
let locks = DidUpdateLocks::default();
for i in 0..10 {
let _guard = locks
.acquire(&format!("did:webvh:scid:example.com:a{i}"))
.await;
}
let live = {
let map = locks
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
map.values().filter(|w| w.strong_count() > 0).count()
};
assert_eq!(live, 0, "no lock is held, so none should still be live");
let retained = {
let map = locks
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
map.len()
};
assert!(
retained <= 1,
"dead entries must be pruned, map still holds {retained}"
);
}
}