use std::{
collections::{HashMap, HashSet},
fs,
hash::Hash,
path::{Path, PathBuf},
sync::{
Mutex, RwLock,
atomic::{AtomicBool, Ordering},
},
time::{Duration, Instant},
};
pub use kcode_k1_invites::UserId;
use kcode_k1_transaction::{SubsystemId, Transaction};
use kcode_k1_txn_ordering::K1TxnOrdering;
pub use kcode_k1_txn_ordering::TxId;
use rusqlite::{Connection, Error as SqlError, TransactionBehavior, params};
const DATABASE_NAME: &str = "groups.sqlite3";
const SUBSYSTEM: &str = "k1-groups-subsystem";
const CREATE_METADATA: &str = "CREATE TABLE metadata (singleton INTEGER PRIMARY KEY CHECK (singleton = 1), schema_version INTEGER NOT NULL, last_applied_txid BLOB)";
const CREATE_GROUPS: &str =
"CREATE TABLE groups (group_id BLOB PRIMARY KEY, revision BLOB NOT NULL)";
const CREATE_USERS: &str = "CREATE TABLE user_memberships (group_id BLOB NOT NULL, user_id BLOB NOT NULL, role INTEGER NOT NULL, PRIMARY KEY (group_id, user_id), FOREIGN KEY (group_id) REFERENCES groups(group_id) ON DELETE CASCADE)";
const CREATE_MODELS: &str = "CREATE TABLE model_memberships (group_id BLOB NOT NULL, model_id BLOB NOT NULL, PRIMARY KEY (group_id, model_id), FOREIGN KEY (group_id) REFERENCES groups(group_id) ON DELETE CASCADE)";
const CREATE_USER_INDEX: &str =
"CREATE INDEX user_memberships_by_user ON user_memberships (user_id, group_id)";
const CREATE_MODEL_INDEX: &str =
"CREATE INDEX model_memberships_by_model ON model_memberships (model_id, group_id)";
#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct GroupId(TxId);
impl GroupId {
pub const fn new(txid: TxId) -> Self {
Self(txid)
}
pub const fn txid(self) -> TxId {
self.0
}
}
#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub struct ModelId([u8; 32]);
impl ModelId {
pub const fn from_bytes(bytes: [u8; 32]) -> Self {
Self(bytes)
}
pub const fn as_bytes(&self) -> &[u8; 32] {
&self.0
}
pub const fn into_bytes(self) -> [u8; 32] {
self.0
}
}
#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)]
pub enum GroupRole {
User,
Admin,
Owner,
}
impl GroupRole {
fn stored(self) -> i64 {
match self {
Self::User => 0,
Self::Admin => 1,
Self::Owner => 2,
}
}
fn from_stored(value: i64) -> Option<Self> {
match value {
0 => Some(Self::User),
1 => Some(Self::Admin),
2 => Some(Self::Owner),
_ => None,
}
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct GroupUser {
user_id: UserId,
role: GroupRole,
}
impl GroupUser {
pub fn user_id(&self) -> UserId {
self.user_id
}
pub fn role(&self) -> GroupRole {
self.role
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct GroupRevision {
group_id: GroupId,
txid: TxId,
}
impl GroupRevision {
pub fn group_id(&self) -> GroupId {
self.group_id
}
pub fn txid(&self) -> TxId {
self.txid
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct Group {
id: GroupId,
revision: GroupRevision,
users: Vec<GroupUser>,
models: Vec<ModelId>,
}
impl Group {
pub fn id(&self) -> GroupId {
self.id
}
pub fn revision(&self) -> &GroupRevision {
&self.revision
}
pub fn users(&self) -> &[GroupUser] {
&self.users
}
pub fn models(&self) -> &[ModelId] {
&self.models
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct GroupMemberships {
revision: Option<TxId>,
user_groups: Vec<GroupId>,
model_groups: Vec<GroupId>,
shared_groups: Vec<GroupId>,
}
impl GroupMemberships {
pub fn revision(&self) -> Option<TxId> {
self.revision
}
pub fn user_groups(&self) -> &[GroupId] {
&self.user_groups
}
pub fn model_groups(&self) -> &[GroupId] {
&self.model_groups
}
pub fn shared_groups(&self) -> &[GroupId] {
&self.shared_groups
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum GroupAction {
Create {
owner: UserId,
},
SetUserRole {
group: GroupId,
actor: UserId,
user: UserId,
role: Option<GroupRole>,
},
SetModelMembership {
group: GroupId,
actor: UserId,
model: ModelId,
present: bool,
},
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum ApplyOutcome {
Applied(GroupRevision),
Unchanged(GroupRevision),
Rejected(String),
}
#[derive(Default)]
struct State {
cursor: Option<TxId>,
groups: HashMap<GroupId, GroupState>,
user_groups: HashMap<UserId, HashSet<GroupId>>,
model_groups: HashMap<ModelId, HashSet<GroupId>>,
}
struct GroupState {
revision: TxId,
users: HashMap<UserId, GroupRole>,
models: HashSet<ModelId>,
owners: usize,
}
pub struct Projection {
state: RwLock<State>,
apply_lane: Mutex<()>,
connection: Mutex<Connection>,
unavailable: AtomicBool,
}
impl Projection {
pub fn open(root: &Path, ordering: &K1TxnOrdering) -> Result<(Self, Option<TxId>), String> {
let started = Instant::now();
let result = Self::open_inner(root, ordering);
let elapsed = started.elapsed();
if elapsed > Duration::from_millis(100) {
let outcome = if result.is_ok() { "ready" } else { "error" };
eprintln!(
"level=warn module=kcode-k1-groups-projection operation=open elapsed_us={} outcome={outcome}",
elapsed.as_micros()
);
}
result
}
fn open_inner(root: &Path, ordering: &K1TxnOrdering) -> Result<(Self, Option<TxId>), String> {
fs::create_dir_all(root)
.map_err(|error| format!("create projection root {}: {error}", root.display()))?;
let database = root.join(DATABASE_NAME);
let (connection, state) = load_or_rebuild(&database, ordering)?;
let cursor = state.cursor;
Ok((
Self {
state: RwLock::new(state),
apply_lane: Mutex::new(()),
connection: Mutex::new(connection),
unavailable: AtomicBool::new(false),
},
cursor,
))
}
pub fn apply(&self, callback_txid: TxId, action: GroupAction) -> Result<ApplyOutcome, String> {
self.ensure_available()?;
let _lane = self.apply_lane.lock().map_err(|_| {
self.mark_unavailable();
"projection apply lane is unavailable until reopen".to_owned()
})?;
self.ensure_available()?;
let prepared = self.prepare(callback_txid, action)?;
{
let mut connection = self.connection.lock().map_err(|_| {
self.mark_unavailable();
"projection database lane is unavailable until reopen".to_owned()
})?;
if let Err(error) = persist_prepared(&mut connection, callback_txid, &prepared) {
self.mark_unavailable();
return Err(error);
}
}
self.publish(callback_txid, prepared)
}
pub fn get(&self, group: GroupId) -> Result<Option<Group>, String> {
self.ensure_available()?;
let state = self.state.read().map_err(|_| {
self.mark_unavailable();
"projection state is unavailable until reopen".to_owned()
})?;
self.ensure_available()?;
let Some(stored) = state.groups.get(&group) else {
return Ok(None);
};
let mut users = Vec::new();
users
.try_reserve(stored.users.len())
.map_err(|error| format!("reserve group users: {error}"))?;
users.extend(stored.users.iter().map(|(user_id, role)| GroupUser {
user_id: *user_id,
role: *role,
}));
let mut models = Vec::new();
models
.try_reserve(stored.models.len())
.map_err(|error| format!("reserve group models: {error}"))?;
models.extend(stored.models.iter().copied());
Ok(Some(Group {
id: group,
revision: GroupRevision {
group_id: group,
txid: stored.revision,
},
users,
models,
}))
}
pub fn groups_for_user(&self, user: UserId) -> Result<Vec<GroupId>, String> {
self.ensure_available()?;
let state = self.state.read().map_err(|_| {
self.mark_unavailable();
"projection state is unavailable until reopen".to_owned()
})?;
self.ensure_available()?;
copy_set(state.user_groups.get(&user), "user groups")
}
pub fn groups_for_model(&self, model: ModelId) -> Result<Vec<GroupId>, String> {
self.ensure_available()?;
let state = self.state.read().map_err(|_| {
self.mark_unavailable();
"projection state is unavailable until reopen".to_owned()
})?;
self.ensure_available()?;
copy_set(state.model_groups.get(&model), "model groups")
}
pub fn memberships(&self, user: UserId, model: ModelId) -> Result<GroupMemberships, String> {
self.ensure_available()?;
let state = self.state.read().map_err(|_| {
self.mark_unavailable();
"projection state is unavailable until reopen".to_owned()
})?;
self.ensure_available()?;
let user_set = state.user_groups.get(&user);
let model_set = state.model_groups.get(&model);
let user_groups = copy_set(user_set, "membership user groups")?;
let model_groups = copy_set(model_set, "membership model groups")?;
let mut shared_groups = Vec::new();
let shared_capacity = user_set
.map(HashSet::len)
.unwrap_or_default()
.min(model_set.map(HashSet::len).unwrap_or_default());
shared_groups
.try_reserve(shared_capacity)
.map_err(|error| format!("reserve shared groups: {error}"))?;
if let (Some(users), Some(models)) = (user_set, model_set) {
let (smaller, larger) = if users.len() <= models.len() {
(users, models)
} else {
(models, users)
};
shared_groups.extend(
smaller
.iter()
.filter(|group| larger.contains(group))
.copied(),
);
}
Ok(GroupMemberships {
revision: state.cursor,
user_groups,
model_groups,
shared_groups,
})
}
pub fn clear(&self) -> Result<(), String> {
self.ensure_available()?;
let _lane = self.apply_lane.lock().map_err(|_| {
self.mark_unavailable();
"projection apply lane is unavailable until reopen".to_owned()
})?;
self.ensure_available()?;
{
let mut connection = self.connection.lock().map_err(|_| {
self.mark_unavailable();
"projection database lane is unavailable until reopen".to_owned()
})?;
if let Err(error) = persist_clear(&mut connection) {
self.mark_unavailable();
return Err(error);
}
}
let old_state = {
let mut state = self.state.write().map_err(|_| {
self.mark_unavailable();
"projection state is unavailable until reopen".to_owned()
})?;
std::mem::take(&mut *state)
};
drop(old_state);
Ok(())
}
fn prepare(&self, callback_txid: TxId, action: GroupAction) -> Result<Prepared, String> {
let mut state = self.state.write().map_err(|_| {
self.mark_unavailable();
"projection state is unavailable until reopen".to_owned()
})?;
self.ensure_available()?;
match action {
GroupAction::Create { owner } => {
let group = GroupId::new(callback_txid);
if state.groups.contains_key(&group) {
return Ok(Prepared::rejected("group already exists"));
}
state
.groups
.try_reserve(1)
.map_err(|error| format!("reserve group index: {error}"))?;
let new_reverse = if let Some(reverse) = state.user_groups.get_mut(&owner) {
reverse
.try_reserve(1)
.map_err(|error| format!("reserve owner reverse membership: {error}"))?;
None
} else {
state
.user_groups
.try_reserve(1)
.map_err(|error| format!("reserve user reverse index: {error}"))?;
let mut reverse = HashSet::new();
reverse
.try_reserve(1)
.map_err(|error| format!("reserve owner reverse membership: {error}"))?;
reverse.insert(group);
Some(reverse)
};
let mut users = HashMap::new();
users
.try_reserve(1)
.map_err(|error| format!("reserve owner membership: {error}"))?;
users.insert(owner, GroupRole::Owner);
Ok(Prepared {
result: PreparedResult::Applied(GroupRevision {
group_id: group,
txid: callback_txid,
}),
mutation: Mutation::Create {
group,
owner,
stored: GroupState {
revision: callback_txid,
users,
models: HashSet::new(),
owners: 1,
},
new_reverse,
},
})
}
GroupAction::SetUserRole {
group,
actor,
user,
role,
} => prepare_user_change(&mut state, group, actor, user, role, callback_txid),
GroupAction::SetModelMembership {
group,
actor,
model,
present,
} => prepare_model_change(&mut state, group, actor, model, present, callback_txid),
}
}
fn publish(&self, callback_txid: TxId, prepared: Prepared) -> Result<ApplyOutcome, String> {
let mut state = self.state.write().map_err(|_| {
self.mark_unavailable();
"projection state is unavailable until reopen".to_owned()
})?;
if !mutation_matches(&state, &prepared.mutation) {
self.mark_unavailable();
return Err(
"projection state contradicted committed database state; reopen required"
.to_owned(),
);
}
match prepared.mutation {
Mutation::None => {}
Mutation::Create {
group,
owner,
stored,
new_reverse,
} => {
state.groups.insert(group, stored);
if let Some(reverse) = new_reverse {
state.user_groups.insert(owner, reverse);
} else {
state
.user_groups
.get_mut(&owner)
.expect("validated owner reverse index")
.insert(group);
}
}
Mutation::User {
group,
user,
from,
to,
previous_owners,
new_reverse,
..
} => {
let stored = state.groups.get_mut(&group).expect("validated group");
match to {
Some(role) => {
stored.users.insert(user, role);
}
None => {
stored.users.remove(&user);
}
}
stored.owners = owner_count_after(previous_owners, from, to);
stored.revision = callback_txid;
match (from, to) {
(None, Some(_)) => {
if let Some(reverse) = new_reverse {
state.user_groups.insert(user, reverse);
} else {
state
.user_groups
.get_mut(&user)
.expect("validated user reverse index")
.insert(group);
}
}
(Some(_), None) => {
let reverse = state
.user_groups
.get_mut(&user)
.expect("validated user reverse index");
reverse.remove(&group);
if reverse.is_empty() {
state.user_groups.remove(&user);
}
}
_ => {}
}
}
Mutation::Model {
group,
model,
from,
to,
new_reverse,
..
} => {
let stored = state.groups.get_mut(&group).expect("validated group");
if to {
stored.models.insert(model);
} else {
stored.models.remove(&model);
}
stored.revision = callback_txid;
match (from, to) {
(false, true) => {
if let Some(reverse) = new_reverse {
state.model_groups.insert(model, reverse);
} else {
state
.model_groups
.get_mut(&model)
.expect("validated model reverse index")
.insert(group);
}
}
(true, false) => {
let reverse = state
.model_groups
.get_mut(&model)
.expect("validated model reverse index");
reverse.remove(&group);
if reverse.is_empty() {
state.model_groups.remove(&model);
}
}
_ => {}
}
}
}
state.cursor = Some(callback_txid);
Ok(prepared.result.into_public())
}
fn ensure_available(&self) -> Result<(), String> {
if self.unavailable.load(Ordering::SeqCst) {
Err("projection is unavailable until reopen".to_owned())
} else {
Ok(())
}
}
fn mark_unavailable(&self) {
self.unavailable.store(true, Ordering::SeqCst);
}
}
fn prepare_user_change(
state: &mut State,
group: GroupId,
actor: UserId,
user: UserId,
to: Option<GroupRole>,
callback_txid: TxId,
) -> Result<Prepared, String> {
let Some(stored) = state.groups.get(&group) else {
return Ok(Prepared::rejected("group does not exist"));
};
let Some(actor_role) = stored.users.get(&actor).copied() else {
return Ok(Prepared::rejected("actor is not authorized"));
};
let from = stored.users.get(&user).copied();
let permitted = match actor_role {
GroupRole::User => false,
GroupRole::Admin => matches!(
(from, to),
(None, Some(GroupRole::User))
| (Some(GroupRole::User), None)
| (Some(GroupRole::User), Some(GroupRole::User))
),
GroupRole::Owner => true,
};
if !permitted {
return Ok(Prepared::rejected(match actor_role {
GroupRole::Admin => "administrator transition is not permitted",
_ => "actor is not authorized",
}));
}
if actor_role == GroupRole::Owner
&& from == Some(GroupRole::Owner)
&& to != Some(GroupRole::Owner)
&& stored.owners == 1
{
return Ok(Prepared::rejected(
"final owner cannot be removed or demoted",
));
}
let previous_revision = stored.revision;
let previous_owners = stored.owners;
if from == to {
return Ok(Prepared::unchanged(group, previous_revision));
}
let mut new_reverse = None;
if from.is_none() && to.is_some() {
state
.groups
.get_mut(&group)
.expect("existing group")
.users
.try_reserve(1)
.map_err(|error| format!("reserve group user membership: {error}"))?;
if let Some(reverse) = state.user_groups.get_mut(&user) {
reverse
.try_reserve(1)
.map_err(|error| format!("reserve user reverse membership: {error}"))?;
} else {
state
.user_groups
.try_reserve(1)
.map_err(|error| format!("reserve user reverse index: {error}"))?;
let mut reverse = HashSet::new();
reverse
.try_reserve(1)
.map_err(|error| format!("reserve user reverse membership: {error}"))?;
reverse.insert(group);
new_reverse = Some(reverse);
}
}
Ok(Prepared {
result: PreparedResult::Applied(GroupRevision {
group_id: group,
txid: callback_txid,
}),
mutation: Mutation::User {
group,
user,
from,
to,
previous_revision,
previous_owners,
new_reverse,
},
})
}
fn prepare_model_change(
state: &mut State,
group: GroupId,
actor: UserId,
model: ModelId,
to: bool,
callback_txid: TxId,
) -> Result<Prepared, String> {
let Some(stored) = state.groups.get(&group) else {
return Ok(Prepared::rejected("group does not exist"));
};
if stored.users.get(&actor) != Some(&GroupRole::Owner) {
return Ok(Prepared::rejected("actor is not authorized"));
}
let from = stored.models.contains(&model);
let previous_revision = stored.revision;
if from == to {
return Ok(Prepared::unchanged(group, previous_revision));
}
let mut new_reverse = None;
if to {
state
.groups
.get_mut(&group)
.expect("existing group")
.models
.try_reserve(1)
.map_err(|error| format!("reserve group model membership: {error}"))?;
if let Some(reverse) = state.model_groups.get_mut(&model) {
reverse
.try_reserve(1)
.map_err(|error| format!("reserve model reverse membership: {error}"))?;
} else {
state
.model_groups
.try_reserve(1)
.map_err(|error| format!("reserve model reverse index: {error}"))?;
let mut reverse = HashSet::new();
reverse
.try_reserve(1)
.map_err(|error| format!("reserve model reverse membership: {error}"))?;
reverse.insert(group);
new_reverse = Some(reverse);
}
}
Ok(Prepared {
result: PreparedResult::Applied(GroupRevision {
group_id: group,
txid: callback_txid,
}),
mutation: Mutation::Model {
group,
model,
from,
to,
previous_revision,
new_reverse,
},
})
}
struct Prepared {
result: PreparedResult,
mutation: Mutation,
}
impl Prepared {
fn rejected(reason: &str) -> Self {
Self {
result: PreparedResult::Rejected(reason.to_owned()),
mutation: Mutation::None,
}
}
fn unchanged(group_id: GroupId, txid: TxId) -> Self {
Self {
result: PreparedResult::Unchanged(GroupRevision { group_id, txid }),
mutation: Mutation::None,
}
}
}
enum PreparedResult {
Applied(GroupRevision),
Unchanged(GroupRevision),
Rejected(String),
}
impl PreparedResult {
fn into_public(self) -> ApplyOutcome {
match self {
Self::Applied(revision) => ApplyOutcome::Applied(revision),
Self::Unchanged(revision) => ApplyOutcome::Unchanged(revision),
Self::Rejected(reason) => ApplyOutcome::Rejected(reason),
}
}
}
enum Mutation {
None,
Create {
group: GroupId,
owner: UserId,
stored: GroupState,
new_reverse: Option<HashSet<GroupId>>,
},
User {
group: GroupId,
user: UserId,
from: Option<GroupRole>,
to: Option<GroupRole>,
previous_revision: TxId,
previous_owners: usize,
new_reverse: Option<HashSet<GroupId>>,
},
Model {
group: GroupId,
model: ModelId,
from: bool,
to: bool,
previous_revision: TxId,
new_reverse: Option<HashSet<GroupId>>,
},
}
fn mutation_matches(state: &State, mutation: &Mutation) -> bool {
match mutation {
Mutation::None => true,
Mutation::Create {
group,
owner,
new_reverse,
..
} => {
if state.groups.contains_key(group) {
return false;
}
match new_reverse {
Some(reverse) => !state.user_groups.contains_key(owner) && reverse.contains(group),
None => state
.user_groups
.get(owner)
.is_some_and(|groups| !groups.contains(group)),
}
}
Mutation::User {
group,
user,
from,
to,
previous_revision,
previous_owners,
new_reverse,
} => {
let Some(stored) = state.groups.get(group) else {
return false;
};
if stored.revision != *previous_revision
|| stored.owners != *previous_owners
|| stored.users.get(user).copied() != *from
{
return false;
}
match (from, to) {
(None, Some(_)) => match new_reverse {
Some(reverse) => {
!state.user_groups.contains_key(user) && reverse.contains(group)
}
None => state
.user_groups
.get(user)
.is_some_and(|groups| !groups.contains(group)),
},
(Some(_), None) => state
.user_groups
.get(user)
.is_some_and(|groups| groups.contains(group)),
_ => true,
}
}
Mutation::Model {
group,
model,
from,
to,
previous_revision,
new_reverse,
} => {
let Some(stored) = state.groups.get(group) else {
return false;
};
if stored.revision != *previous_revision || stored.models.contains(model) != *from {
return false;
}
match (from, to) {
(false, true) => match new_reverse {
Some(reverse) => {
!state.model_groups.contains_key(model) && reverse.contains(group)
}
None => state
.model_groups
.get(model)
.is_some_and(|groups| !groups.contains(group)),
},
(true, false) => state
.model_groups
.get(model)
.is_some_and(|groups| groups.contains(group)),
_ => true,
}
}
}
}
fn owner_count_after(previous: usize, from: Option<GroupRole>, to: Option<GroupRole>) -> usize {
match (from == Some(GroupRole::Owner), to == Some(GroupRole::Owner)) {
(false, true) => previous + 1,
(true, false) => previous - 1,
_ => previous,
}
}
fn persist_prepared(
connection: &mut Connection,
callback_txid: TxId,
prepared: &Prepared,
) -> Result<(), String> {
let transaction = connection
.transaction_with_behavior(TransactionBehavior::Immediate)
.map_err(|error| format!("begin projection apply transaction: {error}"))?;
match &prepared.mutation {
Mutation::None => {}
Mutation::Create {
group,
owner,
stored,
..
} => {
expect_one(
transaction.execute(
"INSERT INTO groups (group_id, revision) VALUES (?1, ?2)",
params![
group.txid().as_bytes().as_slice(),
stored.revision.as_bytes().as_slice()
],
),
"insert group",
)?;
let owner_txid = owner.as_tx_id();
expect_one(
transaction.execute(
"INSERT INTO user_memberships (group_id, user_id, role) VALUES (?1, ?2, ?3)",
params![
group.txid().as_bytes().as_slice(),
owner_txid.as_bytes().as_slice(),
GroupRole::Owner.stored()
],
),
"insert owner membership",
)?;
}
Mutation::User {
group,
user,
from,
to,
..
} => {
let user_txid = user.as_tx_id();
match (from, to) {
(None, Some(role)) => expect_one(
transaction.execute(
"INSERT INTO user_memberships (group_id, user_id, role) VALUES (?1, ?2, ?3)",
params![
group.txid().as_bytes().as_slice(),
user_txid.as_bytes().as_slice(),
role.stored()
],
),
"insert user membership",
)?,
(Some(_), Some(role)) => expect_one(
transaction.execute(
"UPDATE user_memberships SET role = ?3 WHERE group_id = ?1 AND user_id = ?2",
params![
group.txid().as_bytes().as_slice(),
user_txid.as_bytes().as_slice(),
role.stored()
],
),
"update user membership",
)?,
(Some(_), None) => expect_one(
transaction.execute(
"DELETE FROM user_memberships WHERE group_id = ?1 AND user_id = ?2",
params![
group.txid().as_bytes().as_slice(),
user_txid.as_bytes().as_slice()
],
),
"delete user membership",
)?,
(None, None) => {
return Err("prepared user mutation has no state change".to_owned());
}
}
update_group_revision(&transaction, *group, callback_txid)?;
}
Mutation::Model {
group,
model,
from,
to,
..
} => {
match (from, to) {
(false, true) => expect_one(
transaction.execute(
"INSERT INTO model_memberships (group_id, model_id) VALUES (?1, ?2)",
params![
group.txid().as_bytes().as_slice(),
model.as_bytes().as_slice()
],
),
"insert model membership",
)?,
(true, false) => expect_one(
transaction.execute(
"DELETE FROM model_memberships WHERE group_id = ?1 AND model_id = ?2",
params![
group.txid().as_bytes().as_slice(),
model.as_bytes().as_slice()
],
),
"delete model membership",
)?,
_ => return Err("prepared model mutation has no state change".to_owned()),
}
update_group_revision(&transaction, *group, callback_txid)?;
}
}
expect_one(
transaction.execute(
"UPDATE metadata SET last_applied_txid = ?1 WHERE singleton = 1",
params![callback_txid.as_bytes().as_slice()],
),
"advance projection cursor",
)?;
transaction
.commit()
.map_err(|error| format!("commit projection apply transaction: {error}"))
}
fn update_group_revision(
transaction: &rusqlite::Transaction<'_>,
group: GroupId,
callback_txid: TxId,
) -> Result<(), String> {
expect_one(
transaction.execute(
"UPDATE groups SET revision = ?2 WHERE group_id = ?1",
params![
group.txid().as_bytes().as_slice(),
callback_txid.as_bytes().as_slice()
],
),
"update group revision",
)
}
fn expect_one(result: Result<usize, SqlError>, operation: &str) -> Result<(), String> {
let changed = result.map_err(|error| format!("{operation}: {error}"))?;
if changed == 1 {
Ok(())
} else {
Err(format!("{operation}: expected one row, changed {changed}"))
}
}
fn persist_clear(connection: &mut Connection) -> Result<(), String> {
let transaction = connection
.transaction_with_behavior(TransactionBehavior::Immediate)
.map_err(|error| format!("begin projection clear transaction: {error}"))?;
transaction
.execute("DELETE FROM user_memberships", [])
.map_err(|error| format!("clear user memberships: {error}"))?;
transaction
.execute("DELETE FROM model_memberships", [])
.map_err(|error| format!("clear model memberships: {error}"))?;
transaction
.execute("DELETE FROM groups", [])
.map_err(|error| format!("clear groups: {error}"))?;
expect_one(
transaction.execute(
"UPDATE metadata SET last_applied_txid = NULL WHERE singleton = 1",
[],
),
"clear projection cursor",
)?;
transaction
.commit()
.map_err(|error| format!("commit projection clear transaction: {error}"))
}
fn copy_set<K>(set: Option<&HashSet<K>>, operation: &str) -> Result<Vec<K>, String>
where
K: Copy + Eq + Hash,
{
let Some(set) = set else {
return Ok(Vec::new());
};
let mut values = Vec::new();
values
.try_reserve(set.len())
.map_err(|error| format!("reserve {operation}: {error}"))?;
values.extend(set.iter().copied());
Ok(values)
}
enum LoadError {
Recoverable,
Fatal(String),
}
fn load_or_rebuild(
database: &Path,
ordering: &K1TxnOrdering,
) -> Result<(Connection, State), String> {
let exists = database.try_exists().map_err(|error| {
format!(
"inspect projection database {}: {error}",
database.display()
)
})?;
if !exists {
return initialize_database(database);
}
match load_existing(database, ordering) {
Ok(loaded) => Ok(loaded),
Err(LoadError::Recoverable) => {
remove_database_files(database)?;
initialize_database(database)
}
Err(LoadError::Fatal(error)) => Err(error),
}
}
fn initialize_database(database: &Path) -> Result<(Connection, State), String> {
let mut connection = Connection::open(database)
.map_err(|error| format!("create projection database {}: {error}", database.display()))?;
let journal: String = connection
.query_row("PRAGMA journal_mode = WAL", [], |row| row.get(0))
.map_err(|error| format!("enable projection WAL: {error}"))?;
if !journal.eq_ignore_ascii_case("wal") {
return Err(format!("enable projection WAL: SQLite selected {journal}"));
}
connection
.execute_batch("PRAGMA synchronous = FULL; PRAGMA foreign_keys = ON;")
.map_err(|error| format!("configure projection database: {error}"))?;
let transaction = connection
.transaction_with_behavior(TransactionBehavior::Immediate)
.map_err(|error| format!("begin projection initialization: {error}"))?;
for (statement, operation) in [
(CREATE_METADATA, "create metadata schema"),
(CREATE_GROUPS, "create groups schema"),
(CREATE_USERS, "create user membership schema"),
(CREATE_MODELS, "create model membership schema"),
(CREATE_USER_INDEX, "create user reverse index"),
(CREATE_MODEL_INDEX, "create model reverse index"),
] {
transaction
.execute(statement, [])
.map_err(|error| format!("{operation}: {error}"))?;
}
expect_one(
transaction.execute(
"INSERT INTO metadata (singleton, schema_version, last_applied_txid) VALUES (1, 1, NULL)",
[],
),
"initialize projection metadata",
)?;
transaction
.commit()
.map_err(|error| format!("commit projection initialization: {error}"))?;
Ok((connection, State::default()))
}
fn load_existing(
database: &Path,
ordering: &K1TxnOrdering,
) -> Result<(Connection, State), LoadError> {
let connection = Connection::open(database).map_err(|error| {
classify_validation_error(
error,
&format!("open projection database {}", database.display()),
)
})?;
let journal: String = validate_sql(
connection.query_row("PRAGMA journal_mode", [], |row| row.get(0)),
"read projection journal mode",
)?;
if !journal.eq_ignore_ascii_case("wal") {
return Err(LoadError::Recoverable);
}
validate_sql(
connection.execute_batch("PRAGMA synchronous = FULL; PRAGMA foreign_keys = ON;"),
"configure projection database",
)?;
validate_quick_check(&connection)?;
validate_schema(&connection)?;
validate_foreign_keys(&connection)?;
let state = load_state(&connection, ordering)?;
Ok((connection, state))
}
fn validate_quick_check(connection: &Connection) -> Result<(), LoadError> {
let mut statement = validate_sql(
connection.prepare("PRAGMA quick_check"),
"prepare projection quick_check",
)?;
let mut rows = validate_sql(statement.query([]), "run projection quick_check")?;
let Some(row) = validate_sql(rows.next(), "read projection quick_check")? else {
return Err(LoadError::Recoverable);
};
let result: String = validate_sql(row.get(0), "decode projection quick_check")?;
if result != "ok" || validate_sql(rows.next(), "finish projection quick_check")?.is_some() {
return Err(LoadError::Recoverable);
}
Ok(())
}
#[derive(Eq, Ord, PartialEq, PartialOrd)]
struct SchemaObject {
name: String,
kind: String,
table: String,
sql: Option<String>,
}
fn validate_schema(connection: &Connection) -> Result<(), LoadError> {
let mut statement = validate_sql(
connection.prepare(
"SELECT name, type, tbl_name, sql FROM sqlite_master ORDER BY name, type, tbl_name",
),
"prepare projection schema validation",
)?;
let mut rows = validate_sql(statement.query([]), "query projection schema")?;
let mut actual = Vec::new();
while let Some(row) = validate_sql(rows.next(), "read projection schema")? {
actual.push(SchemaObject {
name: validate_sql(row.get(0), "decode projection schema name")?,
kind: validate_sql(row.get(1), "decode projection schema type")?,
table: validate_sql(row.get(2), "decode projection schema table")?,
sql: validate_sql(row.get(3), "decode projection schema SQL")?,
});
}
let mut expected = vec![
schema("groups", "table", "groups", Some(CREATE_GROUPS)),
schema("metadata", "table", "metadata", Some(CREATE_METADATA)),
schema(
"model_memberships",
"table",
"model_memberships",
Some(CREATE_MODELS),
),
schema(
"model_memberships_by_model",
"index",
"model_memberships",
Some(CREATE_MODEL_INDEX),
),
schema("sqlite_autoindex_groups_1", "index", "groups", None),
schema(
"sqlite_autoindex_model_memberships_1",
"index",
"model_memberships",
None,
),
schema(
"sqlite_autoindex_user_memberships_1",
"index",
"user_memberships",
None,
),
schema(
"user_memberships",
"table",
"user_memberships",
Some(CREATE_USERS),
),
schema(
"user_memberships_by_user",
"index",
"user_memberships",
Some(CREATE_USER_INDEX),
),
];
expected.sort();
if actual != expected {
return Err(LoadError::Recoverable);
}
Ok(())
}
fn schema(name: &str, kind: &str, table: &str, sql: Option<&str>) -> SchemaObject {
SchemaObject {
name: name.to_owned(),
kind: kind.to_owned(),
table: table.to_owned(),
sql: sql.map(str::to_owned),
}
}
fn validate_foreign_keys(connection: &Connection) -> Result<(), LoadError> {
let mut statement = validate_sql(
connection.prepare("PRAGMA foreign_key_check"),
"prepare projection foreign key validation",
)?;
let mut rows = validate_sql(statement.query([]), "run projection foreign key validation")?;
if validate_sql(rows.next(), "read projection foreign key validation")?.is_some() {
return Err(LoadError::Recoverable);
}
Ok(())
}
fn load_state(connection: &Connection, ordering: &K1TxnOrdering) -> Result<State, LoadError> {
let mut state = State::default();
let mut metadata = validate_sql(
connection.prepare(
"SELECT singleton, schema_version, last_applied_txid FROM metadata ORDER BY singleton",
),
"prepare projection metadata",
)?;
let mut metadata_rows = validate_sql(metadata.query([]), "query projection metadata")?;
let Some(row) = validate_sql(metadata_rows.next(), "read projection metadata")? else {
return Err(LoadError::Recoverable);
};
let singleton: i64 = validate_sql(row.get(0), "decode metadata singleton")?;
let version: i64 = validate_sql(row.get(1), "decode metadata version")?;
let cursor: Option<Vec<u8>> = validate_sql(row.get(2), "decode metadata cursor")?;
if singleton != 1
|| version != 1
|| validate_sql(metadata_rows.next(), "finish projection metadata")?.is_some()
{
return Err(LoadError::Recoverable);
}
state.cursor = match cursor {
Some(bytes) => Some(TxId::from_bytes(exact_bytes::<12>(&bytes)?)),
None => None,
};
let group_count = table_count(connection, "groups")?;
state
.groups
.try_reserve(group_count)
.map_err(|error| LoadError::Fatal(format!("reserve group state: {error}")))?;
let mut groups = validate_sql(
connection.prepare("SELECT group_id, revision FROM groups"),
"prepare stored groups",
)?;
let mut group_rows = validate_sql(groups.query([]), "query stored groups")?;
while let Some(row) = validate_sql(group_rows.next(), "read stored group")? {
let group_bytes: Vec<u8> = validate_sql(row.get(0), "decode stored group ID")?;
let revision_bytes: Vec<u8> = validate_sql(row.get(1), "decode stored group revision")?;
let group = GroupId::new(TxId::from_bytes(exact_bytes::<12>(&group_bytes)?));
let revision = TxId::from_bytes(exact_bytes::<12>(&revision_bytes)?);
if state
.groups
.insert(
group,
GroupState {
revision,
users: HashMap::new(),
models: HashSet::new(),
owners: 0,
},
)
.is_some()
{
return Err(LoadError::Recoverable);
}
}
let mut users = validate_sql(
connection.prepare("SELECT group_id, user_id, role FROM user_memberships"),
"prepare stored user memberships",
)?;
let mut user_rows = validate_sql(users.query([]), "query stored user memberships")?;
while let Some(row) = validate_sql(user_rows.next(), "read stored user membership")? {
let group_bytes: Vec<u8> = validate_sql(row.get(0), "decode membership group ID")?;
let user_bytes: Vec<u8> = validate_sql(row.get(1), "decode membership user ID")?;
let role_value: i64 = validate_sql(row.get(2), "decode membership role")?;
let group = GroupId::new(TxId::from_bytes(exact_bytes::<12>(&group_bytes)?));
let user = UserId::from_tx_id(TxId::from_bytes(exact_bytes::<12>(&user_bytes)?));
let Some(role) = GroupRole::from_stored(role_value) else {
return Err(LoadError::Recoverable);
};
let Some(stored) = state.groups.get_mut(&group) else {
return Err(LoadError::Recoverable);
};
stored
.users
.try_reserve(1)
.map_err(|error| LoadError::Fatal(format!("reserve stored users: {error}")))?;
if stored.users.insert(user, role).is_some() {
return Err(LoadError::Recoverable);
}
if role == GroupRole::Owner {
stored.owners = stored
.owners
.checked_add(1)
.ok_or_else(|| LoadError::Fatal("stored owner count overflow".to_owned()))?;
}
insert_reverse(
&mut state.user_groups,
user,
group,
"stored user reverse index",
)?;
}
let mut models = validate_sql(
connection.prepare("SELECT group_id, model_id FROM model_memberships"),
"prepare stored model memberships",
)?;
let mut model_rows = validate_sql(models.query([]), "query stored model memberships")?;
while let Some(row) = validate_sql(model_rows.next(), "read stored model membership")? {
let group_bytes: Vec<u8> = validate_sql(row.get(0), "decode model membership group ID")?;
let model_bytes: Vec<u8> = validate_sql(row.get(1), "decode model membership model ID")?;
let group = GroupId::new(TxId::from_bytes(exact_bytes::<12>(&group_bytes)?));
let model = ModelId::from_bytes(exact_bytes::<32>(&model_bytes)?);
let Some(stored) = state.groups.get_mut(&group) else {
return Err(LoadError::Recoverable);
};
stored
.models
.try_reserve(1)
.map_err(|error| LoadError::Fatal(format!("reserve stored models: {error}")))?;
if !stored.models.insert(model) {
return Err(LoadError::Recoverable);
}
insert_reverse(
&mut state.model_groups,
model,
group,
"stored model reverse index",
)?;
}
if state.groups.values().any(|group| group.owners == 0)
|| (!state.groups.is_empty() && state.cursor.is_none())
{
return Err(LoadError::Recoverable);
}
validate_canonical_ids(&state, ordering)?;
Ok(state)
}
fn table_count(connection: &Connection, table: &str) -> Result<usize, LoadError> {
let sql = format!("SELECT COUNT(*) FROM {table}");
let count: i64 = validate_sql(
connection.query_row(&sql, [], |row| row.get(0)),
"count projection rows",
)?;
usize::try_from(count).map_err(|_| LoadError::Recoverable)
}
fn insert_reverse<K>(
reverse: &mut HashMap<K, HashSet<GroupId>>,
key: K,
group: GroupId,
operation: &str,
) -> Result<(), LoadError>
where
K: Copy + Eq + Hash,
{
if !reverse.contains_key(&key) {
reverse
.try_reserve(1)
.map_err(|error| LoadError::Fatal(format!("reserve {operation}: {error}")))?;
}
let groups = reverse.entry(key).or_default();
groups
.try_reserve(1)
.map_err(|error| LoadError::Fatal(format!("reserve {operation}: {error}")))?;
if !groups.insert(group) {
return Err(LoadError::Recoverable);
}
Ok(())
}
fn validate_canonical_ids(state: &State, ordering: &K1TxnOrdering) -> Result<(), LoadError> {
let subsystem = SubsystemId::from_str(SUBSYSTEM)
.map_err(|error| LoadError::Fatal(format!("construct groups subsystem ID: {error}")))?;
let mut validated = HashSet::new();
validated
.try_reserve(
state
.groups
.len()
.checked_mul(2)
.and_then(|count| count.checked_add(1))
.ok_or_else(|| LoadError::Fatal("canonical ID count overflow".to_owned()))?,
)
.map_err(|error| LoadError::Fatal(format!("reserve canonical ID validation: {error}")))?;
for txid in state
.groups
.iter()
.flat_map(|(group, stored)| [group.txid(), stored.revision])
.chain(state.cursor)
{
if validated.contains(&txid) {
continue;
}
let bytes = ordering.get_txn(txid).map_err(|error| {
LoadError::Fatal(format!("validate canonical transaction {txid}: {error}"))
})?;
let Some(bytes) = bytes else {
return Err(LoadError::Recoverable);
};
let transaction = Transaction::parse(&bytes).map_err(|_| LoadError::Recoverable)?;
if transaction.subsystem() != subsystem {
return Err(LoadError::Recoverable);
}
validated.insert(txid);
}
Ok(())
}
fn exact_bytes<const N: usize>(bytes: &[u8]) -> Result<[u8; N], LoadError> {
bytes.try_into().map_err(|_| LoadError::Recoverable)
}
fn validate_sql<T>(result: Result<T, SqlError>, operation: &str) -> Result<T, LoadError> {
result.map_err(|error| classify_validation_error(error, operation))
}
fn classify_validation_error(error: SqlError, operation: &str) -> LoadError {
let recoverable = match &error {
SqlError::SqliteFailure(code, _)
if matches!(
code.code,
rusqlite::ffi::ErrorCode::DatabaseCorrupt | rusqlite::ffi::ErrorCode::NotADatabase
) =>
{
true
}
SqlError::FromSqlConversionFailure(..)
| SqlError::IntegralValueOutOfRange(..)
| SqlError::Utf8Error(..)
| SqlError::InvalidColumnType(..)
| SqlError::QueryReturnedNoRows => true,
_ => false,
};
if recoverable {
LoadError::Recoverable
} else {
LoadError::Fatal(format!("{operation}: {error}"))
}
}
fn remove_database_files(database: &Path) -> Result<(), String> {
for path in [
sidecar_path(database, "-wal"),
sidecar_path(database, "-shm"),
database.to_path_buf(),
] {
match fs::remove_file(&path) {
Ok(()) => {}
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {}
Err(error) => {
return Err(format!(
"remove recoverable projection file {}: {error}",
path.display()
));
}
}
}
Ok(())
}
fn sidecar_path(database: &Path, suffix: &str) -> PathBuf {
let mut path = database.as_os_str().to_os_string();
path.push(suffix);
PathBuf::from(path)
}
#[cfg(test)]
mod internal_tests {
use std::{
sync::{Arc, TryLockError, mpsc},
thread,
time::{Duration, Instant},
};
use kcode_k1_txn_ordering::K1TxnOrdering;
use rusqlite::Connection;
use tempfile::TempDir;
use super::{GroupAction, GroupId, Projection, TxId, UserId};
#[test]
fn blocked_apply_does_not_block_queries() {
let temporary = TempDir::new().unwrap();
let ordering = K1TxnOrdering::open(&temporary.path().join("ordering")).unwrap();
let projection_root = temporary.path().join("projection");
let (projection, _) = Projection::open(&projection_root, &ordering).unwrap();
let projection = Arc::new(projection);
let blocker = Connection::open(projection_root.join("groups.sqlite3")).unwrap();
blocker.execute_batch("BEGIN IMMEDIATE").unwrap();
let applying = Arc::clone(&projection);
let (apply_sent, apply_received) = mpsc::channel();
let apply = thread::spawn(move || {
apply_sent
.send(applying.apply(
TxId::from_bytes([1; 12]),
GroupAction::Create {
owner: UserId::from_tx_id(TxId::from_bytes([2; 12])),
},
))
.unwrap();
});
let waiting_started = Instant::now();
loop {
match projection.apply_lane.try_lock() {
Err(TryLockError::WouldBlock) => break,
Err(TryLockError::Poisoned(_)) => panic!("apply lane poisoned"),
Ok(lane) => drop(lane),
}
assert!(waiting_started.elapsed() < Duration::from_secs(1));
thread::yield_now();
}
assert!(matches!(
apply_received.try_recv(),
Err(mpsc::TryRecvError::Empty)
));
let querying = Arc::clone(&projection);
let (query_sent, query_received) = mpsc::channel();
let query = thread::spawn(move || {
query_sent
.send(querying.get(GroupId::new(TxId::from_bytes([1; 12]))))
.unwrap();
});
let query_result = query_received.recv_timeout(Duration::from_secs(1));
blocker.execute_batch("ROLLBACK").unwrap();
assert!(query_result.unwrap().unwrap().is_none());
assert!(
apply_received
.recv_timeout(Duration::from_secs(1))
.unwrap()
.is_ok()
);
query.join().unwrap();
apply.join().unwrap();
}
}