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 controller_profile_ids(
89        &self,
90        context: &AccessContext,
91        access_ids: &[AccessId],
92        user_groups: &[GroupId],
93    ) -> Vec<Option<ProfileId>> {
94        access_ids
95            .iter()
96            .map(|access_id| {
97                let record = self.records.get(access_id)?;
98                let authority = record.policy.authority();
99                if context.filter().contains(authority) {
100                    return None;
101                }
102                let controls = match authority {
103                    Authority::User(user) => user == context.user(),
104                    Authority::Group(group) => user_groups.contains(&group),
105                };
106                controls.then_some(record.profile_id)
107            })
108            .collect()
109    }
110
111    pub fn check(
112        &self,
113        context: &AccessContext,
114        access_id: AccessId,
115        expected_subsystem: SubsystemId,
116        user_groups: &[GroupId],
117        model_groups: &[GroupId],
118        groups_revision: Option<TxId>,
119    ) -> AccessCheck {
120        let Some(record) = self.records.get(&access_id) else {
121            return concealed();
122        };
123        if record.target.subsystem() != expected_subsystem {
124            return concealed();
125        }
126        let decision = evaluate(&record.policy, context, user_groups, model_groups);
127        let can_view = decision.can_view();
128        let can_edit = decision.can_edit();
129        if !can_view && !can_edit {
130            return concealed();
131        }
132        AccessCheck::new(
133            can_view,
134            can_edit,
135            can_view.then(|| record.target.clone()),
136            Some(record.revision),
137            groups_revision,
138        )
139        .expect("permission result is valid")
140    }
141
142    pub fn check_many(
143        &self,
144        context: &AccessContext,
145        access_ids: &[AccessId],
146        expected_subsystem: SubsystemId,
147        user_groups: &[GroupId],
148        model_groups: &[GroupId],
149        groups_revision: Option<TxId>,
150    ) -> Vec<AccessCheck> {
151        access_ids
152            .iter()
153            .map(|access_id| {
154                self.check(
155                    context,
156                    *access_id,
157                    expected_subsystem,
158                    user_groups,
159                    model_groups,
160                    groups_revision,
161                )
162            })
163            .collect()
164    }
165
166    pub fn resolve_visible_targets(
167        &self,
168        context: &AccessContext,
169        targets: &[Target],
170        expected_subsystem: SubsystemId,
171        user_groups: &[GroupId],
172        model_groups: &[GroupId],
173    ) -> Vec<Option<AccessId>> {
174        targets
175            .iter()
176            .map(|target| {
177                if target.subsystem() != expected_subsystem {
178                    return None;
179                }
180                let access_id = self.targets.get(target)?;
181                let record = self.records.get(access_id)?;
182                evaluate(&record.policy, context, user_groups, model_groups)
183                    .can_view()
184                    .then_some(*access_id)
185            })
186            .collect()
187    }
188
189    pub fn accepts_edit_witness(&self, access_id: AccessId, witness: Authority) -> bool {
190        self.records.get(&access_id).is_some_and(|record| {
191            record.policy.authority() == witness || record.policy.editors().contains(&witness)
192        })
193    }
194
195    fn intern(&mut self, policy: AccessPolicy) -> Arc<AccessPolicy> {
196        if let Some(shared) = self.policies.get(&policy).and_then(Weak::upgrade) {
197            return shared;
198        }
199        let shared = Arc::new(policy.clone());
200        self.policies.insert(policy, Arc::downgrade(&shared));
201        shared
202    }
203}
204
205fn concealed() -> AccessCheck {
206    AccessCheck::new(false, false, None, None, None).expect("concealed result is valid")
207}
208
209#[cfg(test)]
210mod tests;