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