Skip to main content

kcode_k1_access_state/
lib.rs

1use std::{
2    collections::HashMap,
3    sync::{Arc, Weak},
4};
5
6use kcode_k1_access_policy::evaluate;
7use kcode_k1_access_store::StoredAccess;
8use kcode_k1_access_types::{
9    AccessCheck, AccessContext, AccessId, AccessPolicy, Authority, GroupId, SubsystemId, Target,
10    TxId,
11};
12
13struct AccessRecord {
14    target: Target,
15    revision: TxId,
16    policy: Arc<AccessPolicy>,
17}
18
19pub struct AccessState {
20    records: HashMap<AccessId, AccessRecord>,
21    targets: HashMap<Target, AccessId>,
22    policies: HashMap<AccessPolicy, Weak<AccessPolicy>>,
23}
24
25impl AccessState {
26    pub fn new(records: Vec<StoredAccess>) -> Result<Self, String> {
27        let mut state = Self {
28            records: HashMap::with_capacity(records.len()),
29            targets: HashMap::with_capacity(records.len()),
30            policies: HashMap::new(),
31        };
32        for record in records {
33            state.create(record)?;
34        }
35        Ok(state)
36    }
37
38    pub fn create(&mut self, record: StoredAccess) -> Result<(), String> {
39        let (access_id, target, revision, policy) = record.into_parts();
40        if self.records.contains_key(&access_id) {
41            return Err("duplicate access ID".to_owned());
42        }
43        if self.targets.contains_key(&target) {
44            return Err("duplicate target".to_owned());
45        }
46        let policy = self.intern(policy);
47        self.targets.insert(target.clone(), access_id);
48        let record = AccessRecord {
49            target,
50            revision,
51            policy,
52        };
53        self.records.insert(access_id, record);
54        Ok(())
55    }
56
57    pub fn replace(
58        &mut self,
59        access_id: AccessId,
60        revision: TxId,
61        policy: AccessPolicy,
62    ) -> Result<(), String> {
63        let old_policy = self
64            .records
65            .get(&access_id)
66            .ok_or_else(|| "unknown access ID".to_owned())?
67            .policy
68            .clone();
69        if old_policy.authority() != policy.authority() {
70            return Err("access authority is immutable".to_owned());
71        }
72        let policy = self.intern(policy);
73        let record = self.records.get_mut(&access_id).expect("record exists");
74        record.revision = revision;
75        record.policy = policy;
76        if Arc::strong_count(&old_policy) == 1 {
77            self.policies.remove(old_policy.as_ref());
78        }
79        Ok(())
80    }
81
82    pub fn check(
83        &self,
84        context: &AccessContext,
85        access_id: AccessId,
86        expected_subsystem: SubsystemId,
87        user_groups: &[GroupId],
88        model_groups: &[GroupId],
89        groups_revision: Option<TxId>,
90    ) -> AccessCheck {
91        let Some(record) = self.records.get(&access_id) else {
92            return concealed();
93        };
94        if record.target.subsystem() != expected_subsystem {
95            return concealed();
96        }
97        let decision = evaluate(&record.policy, context, user_groups, model_groups);
98        let can_view = decision.can_view();
99        let can_edit = decision.can_edit();
100        if !can_view && !can_edit {
101            return concealed();
102        }
103        AccessCheck::new(
104            can_view,
105            can_edit,
106            can_view.then(|| record.target.clone()),
107            Some(record.revision),
108            groups_revision,
109        )
110        .expect("permission result is valid")
111    }
112
113    pub fn check_many(
114        &self,
115        context: &AccessContext,
116        access_ids: &[AccessId],
117        expected_subsystem: SubsystemId,
118        user_groups: &[GroupId],
119        model_groups: &[GroupId],
120        groups_revision: Option<TxId>,
121    ) -> Vec<AccessCheck> {
122        access_ids
123            .iter()
124            .map(|access_id| {
125                self.check(
126                    context,
127                    *access_id,
128                    expected_subsystem,
129                    user_groups,
130                    model_groups,
131                    groups_revision,
132                )
133            })
134            .collect()
135    }
136
137    pub fn resolve_visible_targets(
138        &self,
139        context: &AccessContext,
140        targets: &[Target],
141        expected_subsystem: SubsystemId,
142        user_groups: &[GroupId],
143        model_groups: &[GroupId],
144    ) -> Vec<Option<AccessId>> {
145        targets
146            .iter()
147            .map(|target| {
148                if target.subsystem() != expected_subsystem {
149                    return None;
150                }
151                let access_id = self.targets.get(target)?;
152                let record = self.records.get(access_id)?;
153                evaluate(&record.policy, context, user_groups, model_groups)
154                    .can_view()
155                    .then_some(*access_id)
156            })
157            .collect()
158    }
159
160    pub fn accepts_edit_witness(&self, access_id: AccessId, witness: Authority) -> bool {
161        self.records.get(&access_id).is_some_and(|record| {
162            record.policy.authority() == witness || record.policy.editors().contains(&witness)
163        })
164    }
165
166    fn intern(&mut self, policy: AccessPolicy) -> Arc<AccessPolicy> {
167        if let Some(shared) = self.policies.get(&policy).and_then(Weak::upgrade) {
168            return shared;
169        }
170        let shared = Arc::new(policy.clone());
171        self.policies.insert(policy, Arc::downgrade(&shared));
172        shared
173    }
174}
175
176fn concealed() -> AccessCheck {
177    AccessCheck::new(false, false, None, None, None).expect("concealed result is valid")
178}
179
180#[cfg(test)]
181mod tests;