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;