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;