use std::{
collections::HashMap,
sync::{Arc, Weak},
};
use kcode_k1_access_policy::evaluate;
use kcode_k1_access_store::StoredAccess;
use kcode_k1_access_types::{
AccessCheck, AccessContext, AccessId, AccessPolicy, Authority, GroupId, ProfileId, SubsystemId,
Target, TxId,
};
struct AccessRecord {
target: Target,
profile_id: ProfileId,
revision: TxId,
policy: Arc<AccessPolicy>,
}
pub struct AccessState {
records: HashMap<AccessId, AccessRecord>,
targets: HashMap<Target, AccessId>,
policies: HashMap<AccessPolicy, Weak<AccessPolicy>>,
}
impl AccessState {
pub fn new(records: Vec<StoredAccess>) -> Result<Self, String> {
let mut state = Self {
records: HashMap::with_capacity(records.len()),
targets: HashMap::with_capacity(records.len()),
policies: HashMap::new(),
};
for record in records {
state.create(record)?;
}
Ok(state)
}
pub fn create(&mut self, record: StoredAccess) -> Result<(), String> {
let (access_id, target, profile_id, revision, policy) = record.into_parts();
if self.records.contains_key(&access_id) {
return Err("duplicate access ID".to_owned());
}
if self.targets.contains_key(&target) {
return Err("duplicate target".to_owned());
}
let policy = self.intern(policy);
self.targets.insert(target.clone(), access_id);
let record = AccessRecord {
target,
profile_id,
revision,
policy,
};
self.records.insert(access_id, record);
Ok(())
}
pub fn replace(
&mut self,
access_id: AccessId,
revision: TxId,
policy: AccessPolicy,
) -> Result<(), String> {
let old_policy = self
.records
.get(&access_id)
.ok_or_else(|| "unknown access ID".to_owned())?
.policy
.clone();
if old_policy.authority() != policy.authority() {
return Err("access authority is immutable".to_owned());
}
let policy = self.intern(policy);
let record = self.records.get_mut(&access_id).expect("record exists");
record.revision = revision;
record.policy = policy;
if Arc::strong_count(&old_policy) == 1 {
self.policies.remove(old_policy.as_ref());
}
Ok(())
}
pub fn profile_id(&self, access_id: AccessId) -> Option<ProfileId> {
self.records.get(&access_id).map(|record| record.profile_id)
}
pub fn controller_profile_ids(
&self,
context: &AccessContext,
access_ids: &[AccessId],
user_groups: &[GroupId],
) -> Vec<Option<ProfileId>> {
access_ids
.iter()
.map(|access_id| {
let record = self.records.get(access_id)?;
let authority = record.policy.authority();
if context.filter().contains(authority) {
return None;
}
let controls = match authority {
Authority::User(user) => user == context.user(),
Authority::Group(group) => user_groups.contains(&group),
};
controls.then_some(record.profile_id)
})
.collect()
}
pub fn check(
&self,
context: &AccessContext,
access_id: AccessId,
expected_subsystem: SubsystemId,
user_groups: &[GroupId],
model_groups: &[GroupId],
groups_revision: Option<TxId>,
) -> AccessCheck {
let Some(record) = self.records.get(&access_id) else {
return concealed();
};
if record.target.subsystem() != expected_subsystem {
return concealed();
}
let decision = evaluate(&record.policy, context, user_groups, model_groups);
let can_view = decision.can_view();
let can_edit = decision.can_edit();
if !can_view && !can_edit {
return concealed();
}
AccessCheck::new(
can_view,
can_edit,
can_view.then(|| record.target.clone()),
Some(record.revision),
groups_revision,
)
.expect("permission result is valid")
}
pub fn check_many(
&self,
context: &AccessContext,
access_ids: &[AccessId],
expected_subsystem: SubsystemId,
user_groups: &[GroupId],
model_groups: &[GroupId],
groups_revision: Option<TxId>,
) -> Vec<AccessCheck> {
access_ids
.iter()
.map(|access_id| {
self.check(
context,
*access_id,
expected_subsystem,
user_groups,
model_groups,
groups_revision,
)
})
.collect()
}
pub fn resolve_visible_targets(
&self,
context: &AccessContext,
targets: &[Target],
expected_subsystem: SubsystemId,
user_groups: &[GroupId],
model_groups: &[GroupId],
) -> Vec<Option<AccessId>> {
targets
.iter()
.map(|target| {
if target.subsystem() != expected_subsystem {
return None;
}
let access_id = self.targets.get(target)?;
let record = self.records.get(access_id)?;
evaluate(&record.policy, context, user_groups, model_groups)
.can_view()
.then_some(*access_id)
})
.collect()
}
pub fn accepts_edit_witness(&self, access_id: AccessId, witness: Authority) -> bool {
self.records.get(&access_id).is_some_and(|record| {
record.policy.authority() == witness || record.policy.editors().contains(&witness)
})
}
fn intern(&mut self, policy: AccessPolicy) -> Arc<AccessPolicy> {
if let Some(shared) = self.policies.get(&policy).and_then(Weak::upgrade) {
return shared;
}
let shared = Arc::new(policy.clone());
self.policies.insert(policy, Arc::downgrade(&shared));
shared
}
}
fn concealed() -> AccessCheck {
AccessCheck::new(false, false, None, None, None).expect("concealed result is valid")
}
#[cfg(test)]
mod tests;