use std::collections::BTreeMap;
use proto_blue_lex_cbor::encode;
use proto_blue_lex_data::LexValue;
use sqlx::sqlite::SqliteConnection;
use sqlx::{Pool, Sqlite};
use crate::audit::append::{
AuditRowForAppend, append_in_tx, read_latest_chain_hash, read_latest_chain_timestamp_ms,
};
use crate::audit::hash::compute_chain_hash;
use crate::error::{Error, Result};
use crate::writer::epoch_ms_now;
#[derive(Debug, Clone)]
pub struct MembershipRow {
pub did: String,
pub note: Option<String>,
pub added_by_moderator: String,
pub added_at: i64,
pub revoked_at: Option<i64>,
pub revoked_by_moderator: Option<String>,
}
pub async fn is_known_caller(pool: &Pool<Sqlite>, did: &str) -> Result<bool> {
let count: i64 = sqlx::query_scalar!(
"SELECT COUNT(*) FROM xrpc_known_callers
WHERE did = ?1 AND revoked_at IS NULL",
did
)
.fetch_one(pool)
.await?;
Ok(count > 0)
}
pub async fn is_trusted_pds(pool: &Pool<Sqlite>, did: &str) -> Result<bool> {
let count: i64 = sqlx::query_scalar!(
"SELECT COUNT(*) FROM xrpc_trusted_pdses
WHERE did = ?1 AND revoked_at IS NULL",
did
)
.fetch_one(pool)
.await?;
Ok(count > 0)
}
pub async fn list_known_callers(
pool: &Pool<Sqlite>,
include_revoked: bool,
) -> Result<Vec<MembershipRow>> {
let rows = if include_revoked {
sqlx::query!(
"SELECT did, note, added_by_moderator, added_at, revoked_at, revoked_by_moderator
FROM xrpc_known_callers ORDER BY added_at ASC"
)
.fetch_all(pool)
.await?
.into_iter()
.map(|r| MembershipRow {
did: r.did,
note: r.note,
added_by_moderator: r.added_by_moderator,
added_at: r.added_at,
revoked_at: r.revoked_at,
revoked_by_moderator: r.revoked_by_moderator,
})
.collect()
} else {
sqlx::query!(
"SELECT did, note, added_by_moderator, added_at, revoked_at, revoked_by_moderator
FROM xrpc_known_callers
WHERE revoked_at IS NULL
ORDER BY added_at ASC"
)
.fetch_all(pool)
.await?
.into_iter()
.map(|r| MembershipRow {
did: r.did,
note: r.note,
added_by_moderator: r.added_by_moderator,
added_at: r.added_at,
revoked_at: r.revoked_at,
revoked_by_moderator: r.revoked_by_moderator,
})
.collect()
};
Ok(rows)
}
pub async fn list_trusted_pdses(
pool: &Pool<Sqlite>,
include_revoked: bool,
) -> Result<Vec<MembershipRow>> {
let rows = if include_revoked {
sqlx::query!(
"SELECT did, note, added_by_moderator, added_at, revoked_at, revoked_by_moderator
FROM xrpc_trusted_pdses ORDER BY added_at ASC"
)
.fetch_all(pool)
.await?
.into_iter()
.map(|r| MembershipRow {
did: r.did,
note: r.note,
added_by_moderator: r.added_by_moderator,
added_at: r.added_at,
revoked_at: r.revoked_at,
revoked_by_moderator: r.revoked_by_moderator,
})
.collect()
} else {
sqlx::query!(
"SELECT did, note, added_by_moderator, added_at, revoked_at, revoked_by_moderator
FROM xrpc_trusted_pdses
WHERE revoked_at IS NULL
ORDER BY added_at ASC"
)
.fetch_all(pool)
.await?
.into_iter()
.map(|r| MembershipRow {
did: r.did,
note: r.note,
added_by_moderator: r.added_by_moderator,
added_at: r.added_at,
revoked_at: r.revoked_at,
revoked_by_moderator: r.revoked_by_moderator,
})
.collect()
};
Ok(rows)
}
#[derive(Debug, Clone, Copy)]
enum MembershipTable {
KnownCallers,
TrustedPdses,
}
impl MembershipTable {
fn name(self) -> &'static str {
match self {
Self::KnownCallers => "xrpc_known_callers",
Self::TrustedPdses => "xrpc_trusted_pdses",
}
}
}
pub async fn add_known_caller(
pool: &Pool<Sqlite>,
did: &str,
note: Option<&str>,
added_by_moderator: &str,
) -> Result<MembershipRow> {
add_membership_row(
pool,
MembershipTable::KnownCallers,
did,
note,
added_by_moderator,
)
.await
}
pub async fn add_trusted_pds(
pool: &Pool<Sqlite>,
did: &str,
note: Option<&str>,
added_by_moderator: &str,
) -> Result<MembershipRow> {
add_membership_row(
pool,
MembershipTable::TrustedPdses,
did,
note,
added_by_moderator,
)
.await
}
pub async fn revoke_known_caller(
pool: &Pool<Sqlite>,
did: &str,
revoked_by_moderator: &str,
) -> Result<()> {
revoke_membership_row(
pool,
MembershipTable::KnownCallers,
did,
revoked_by_moderator,
)
.await
}
pub async fn revoke_trusted_pds(
pool: &Pool<Sqlite>,
did: &str,
revoked_by_moderator: &str,
) -> Result<()> {
revoke_membership_row(
pool,
MembershipTable::TrustedPdses,
did,
revoked_by_moderator,
)
.await
}
async fn add_membership_row(
pool: &Pool<Sqlite>,
table: MembershipTable,
did: &str,
note: Option<&str>,
added_by_moderator: &str,
) -> Result<MembershipRow> {
let mut conn = pool
.acquire()
.await
.map_err(|e| Error::Signing(format!("membership acquire: {e}")))?;
sqlx::query("BEGIN IMMEDIATE")
.execute(&mut *conn)
.await
.map_err(|e| Error::Signing(format!("membership begin: {e}")))?;
let result = perform_add(&mut conn, table, did, note, added_by_moderator).await;
match result {
Ok(row) => {
sqlx::query("COMMIT")
.execute(&mut *conn)
.await
.map_err(|e| Error::Signing(format!("membership commit: {e}")))?;
Ok(row)
}
Err(e) => {
let _ = sqlx::query("ROLLBACK").execute(&mut *conn).await;
Err(e)
}
}
}
async fn perform_add(
conn: &mut SqliteConnection,
table: MembershipTable,
did: &str,
note: Option<&str>,
added_by_moderator: &str,
) -> Result<MembershipRow> {
let now_ms = epoch_ms_now();
let max_existing = read_latest_chain_timestamp_ms(&mut *conn)
.await?
.unwrap_or(0);
let added_at = now_ms.max(max_existing.saturating_add(1));
let prev_hash = read_latest_chain_hash(&mut *conn).await?;
let row_hash = compute_membership_row_hash(
&prev_hash,
&MembershipRowForHashing {
did,
note,
added_by_moderator,
added_at,
},
)?;
let prev_hash_slice: &[u8] = &prev_hash;
let row_hash_slice: &[u8] = &row_hash;
let sql = match table {
MembershipTable::KnownCallers => {
"INSERT INTO xrpc_known_callers
(did, note, added_by_moderator, added_at, prev_hash, row_hash)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)"
}
MembershipTable::TrustedPdses => {
"INSERT INTO xrpc_trusted_pdses
(did, note, added_by_moderator, added_at, prev_hash, row_hash)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)"
}
};
sqlx::query(sql)
.bind(did)
.bind(note)
.bind(added_by_moderator)
.bind(added_at)
.bind(prev_hash_slice)
.bind(row_hash_slice)
.execute(&mut *conn)
.await
.map_err(|e| Error::Signing(format!("{} insert: {e}", table.name())))?;
Ok(MembershipRow {
did: did.to_string(),
note: note.map(str::to_string),
added_by_moderator: added_by_moderator.to_string(),
added_at,
revoked_at: None,
revoked_by_moderator: None,
})
}
async fn revoke_membership_row(
pool: &Pool<Sqlite>,
table: MembershipTable,
did: &str,
revoked_by_moderator: &str,
) -> Result<()> {
let mut tx = pool.begin().await?;
let existing = match table {
MembershipTable::KnownCallers => sqlx::query!(
"SELECT added_by_moderator, added_at, revoked_at FROM xrpc_known_callers
WHERE did = ?1",
did
)
.fetch_optional(&mut *tx)
.await?
.map(|r| (r.added_by_moderator, r.added_at, r.revoked_at)),
MembershipTable::TrustedPdses => sqlx::query!(
"SELECT added_by_moderator, added_at, revoked_at FROM xrpc_trusted_pdses
WHERE did = ?1",
did
)
.fetch_optional(&mut *tx)
.await?
.map(|r| (r.added_by_moderator, r.added_at, r.revoked_at)),
};
let Some((_added_by, _added_at, revoked_at)) = existing else {
return Err(Error::Signing(format!(
"{} row not found for did {did:?}",
table.name()
)));
};
if revoked_at.is_some() {
return Err(Error::Signing(format!(
"{} row for did {did:?} is already revoked",
table.name()
)));
}
let now_ms = epoch_ms_now();
let update_sql = match table {
MembershipTable::KnownCallers => {
"UPDATE xrpc_known_callers
SET revoked_at = ?1, revoked_by_moderator = ?2
WHERE did = ?3"
}
MembershipTable::TrustedPdses => {
"UPDATE xrpc_trusted_pdses
SET revoked_at = ?1, revoked_by_moderator = ?2
WHERE did = ?3"
}
};
sqlx::query(update_sql)
.bind(now_ms)
.bind(revoked_by_moderator)
.bind(did)
.execute(&mut *tx)
.await
.map_err(|e| Error::Signing(format!("{} revoke: {e}", table.name())))?;
let audit_action = match table {
MembershipTable::KnownCallers => "xrpc_known_caller_revoked",
MembershipTable::TrustedPdses => "xrpc_trusted_pds_revoked",
};
append_in_tx(
&mut tx,
&AuditRowForAppend {
created_at: now_ms,
action: audit_action.to_string(),
actor_did: revoked_by_moderator.to_string(),
target: Some(did.to_string()),
target_cid: None,
outcome: "success".to_string(),
reason: None,
},
)
.await?;
tx.commit().await?;
Ok(())
}
struct MembershipRowForHashing<'a> {
did: &'a str,
note: Option<&'a str>,
added_by_moderator: &'a str,
added_at: i64,
}
fn compute_membership_row_hash(
prev_hash: &[u8; 32],
row: &MembershipRowForHashing<'_>,
) -> Result<[u8; 32]> {
let canonical = encode(&row_to_lex_value(row))?;
Ok(compute_chain_hash(prev_hash, &canonical))
}
fn row_to_lex_value(row: &MembershipRowForHashing<'_>) -> LexValue {
let mut m = BTreeMap::new();
m.insert("did".to_string(), LexValue::String(row.did.to_string()));
if let Some(note) = row.note {
m.insert("note".to_string(), LexValue::String(note.to_string()));
}
m.insert(
"added_by_moderator".to_string(),
LexValue::String(row.added_by_moderator.to_string()),
);
m.insert("added_at".to_string(), LexValue::Integer(row.added_at));
LexValue::Map(m)
}
pub(crate) fn recompute_membership_row_hash(
prev_hash: &[u8; 32],
did: &str,
note: Option<&str>,
added_by_moderator: &str,
added_at: i64,
) -> Result<[u8; 32]> {
compute_membership_row_hash(
prev_hash,
&MembershipRowForHashing {
did,
note,
added_by_moderator,
added_at,
},
)
}
#[cfg(test)]
#[allow(unused_imports, dead_code)]
mod tests {
use super::*;
use crate::storage;
use tempfile::tempdir;
async fn fresh_pool() -> Pool<Sqlite> {
let dir = tempdir().unwrap();
let path = dir.path().join("membership-test.db");
let pool = storage::open(&path).await.unwrap();
Box::leak(Box::new(dir));
pool
}
const M_DID: &str = "did:plc:moderator0000000000000000";
#[tokio::test]
async fn add_known_caller_persists_active_row() {
let pool = fresh_pool().await;
add_known_caller(&pool, "did:plc:alice", Some("alice"), M_DID)
.await
.unwrap();
assert!(is_known_caller(&pool, "did:plc:alice").await.unwrap());
assert!(!is_known_caller(&pool, "did:plc:bob").await.unwrap());
}
#[tokio::test]
async fn list_known_callers_active_only_excludes_revoked() {
let pool = fresh_pool().await;
add_known_caller(&pool, "did:plc:alice", None, M_DID)
.await
.unwrap();
add_known_caller(&pool, "did:plc:bob", None, M_DID)
.await
.unwrap();
revoke_known_caller(&pool, "did:plc:alice", M_DID)
.await
.unwrap();
let active = list_known_callers(&pool, false).await.unwrap();
assert_eq!(active.len(), 1);
assert_eq!(active[0].did, "did:plc:bob");
let all = list_known_callers(&pool, true).await.unwrap();
assert_eq!(all.len(), 2);
}
#[tokio::test]
async fn revoke_known_caller_marks_inactive_and_writes_audit_log() {
let pool = fresh_pool().await;
add_known_caller(&pool, "did:plc:alice", Some("alice"), M_DID)
.await
.unwrap();
revoke_known_caller(&pool, "did:plc:alice", M_DID)
.await
.unwrap();
assert!(!is_known_caller(&pool, "did:plc:alice").await.unwrap());
let audit_count: i64 = sqlx::query_scalar!(
"SELECT COUNT(*) FROM audit_log
WHERE action = 'xrpc_known_caller_revoked' AND target = 'did:plc:alice'"
)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(audit_count, 1);
}
#[tokio::test]
async fn revoke_unknown_did_errors() {
let pool = fresh_pool().await;
let res = revoke_known_caller(&pool, "did:plc:nope", M_DID).await;
assert!(res.is_err());
}
#[tokio::test]
async fn double_revoke_errors() {
let pool = fresh_pool().await;
add_known_caller(&pool, "did:plc:alice", None, M_DID)
.await
.unwrap();
revoke_known_caller(&pool, "did:plc:alice", M_DID)
.await
.unwrap();
let res = revoke_known_caller(&pool, "did:plc:alice", M_DID).await;
assert!(res.is_err(), "double-revoke must fail");
}
#[tokio::test]
async fn add_trusted_pds_persists_active_row() {
let pool = fresh_pool().await;
add_trusted_pds(&pool, "did:web:bsky.example.com", None, M_DID)
.await
.unwrap();
assert!(
is_trusted_pds(&pool, "did:web:bsky.example.com")
.await
.unwrap()
);
}
#[tokio::test]
async fn revoke_trusted_pds_marks_inactive() {
let pool = fresh_pool().await;
add_trusted_pds(&pool, "did:web:bsky.example.com", None, M_DID)
.await
.unwrap();
revoke_trusted_pds(&pool, "did:web:bsky.example.com", M_DID)
.await
.unwrap();
assert!(
!is_trusted_pds(&pool, "did:web:bsky.example.com")
.await
.unwrap()
);
}
#[tokio::test]
async fn known_caller_membership_does_not_imply_trusted_pds() {
let pool = fresh_pool().await;
add_known_caller(&pool, "did:plc:alice", None, M_DID)
.await
.unwrap();
assert!(is_known_caller(&pool, "did:plc:alice").await.unwrap());
assert!(!is_trusted_pds(&pool, "did:plc:alice").await.unwrap());
}
#[tokio::test]
async fn membership_rows_chain_with_audit_log() {
let pool = fresh_pool().await;
crate::audit::append::append_via_pool(
&pool,
&AuditRowForAppend {
created_at: 100,
action: "label_applied".into(),
actor_did: "did:plc:m1".into(),
target: None,
target_cid: None,
outcome: "success".into(),
reason: None,
},
)
.await
.unwrap();
let m_row = add_known_caller(&pool, "did:plc:alice", None, M_DID)
.await
.unwrap();
assert!(m_row.added_at > 100);
let stored_prev: Vec<u8> = sqlx::query_scalar!(
"SELECT prev_hash FROM xrpc_known_callers WHERE did = 'did:plc:alice'"
)
.fetch_one(&pool)
.await
.unwrap();
let audit_row_hash: Vec<u8> =
sqlx::query_scalar!(r#"SELECT row_hash AS "row_hash!" FROM audit_log WHERE id = 1"#)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(stored_prev, audit_row_hash);
}
}