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