1use std::collections::BTreeSet;
2use std::fmt;
3
4use sha2::{Digest, Sha256};
5
6use crate::{
7 classify_operation, Gate, GateDecision, GateError, GateRequest, GateValidationError,
8 OperationAccess, OPERATION_CLASSIFIER_VERSION,
9};
10
11#[derive(Clone)]
17pub struct CallerEnrollmentGate {
18 granted_actors: BTreeSet<String>,
19 grant_unattributed: bool,
20 deny_writes_for: BTreeSet<String>,
21 invalid_write_policy: bool,
22 configuration_fingerprint: String,
23}
24
25impl CallerEnrollmentGate {
26 pub fn new(granted_actors: Vec<String>, grant_unattributed: bool) -> Self {
28 let granted_actors: BTreeSet<String> = granted_actors.into_iter().collect();
29 let mut hasher = Sha256::new();
30 hasher.update(b"khive.caller-enrollment-gate.v1\0");
31 hasher.update([u8::from(grant_unattributed)]);
32 hasher.update((granted_actors.len() as u64).to_be_bytes());
33 for actor in &granted_actors {
34 hasher.update((actor.len() as u64).to_be_bytes());
35 hasher.update(actor.as_bytes());
36 }
37 let configuration_fingerprint = format!("sha256:{:x}", hasher.finalize());
38 Self {
39 granted_actors,
40 grant_unattributed,
41 deny_writes_for: BTreeSet::new(),
42 invalid_write_policy: false,
43 configuration_fingerprint,
44 }
45 }
46
47 pub fn with_write_denials(
56 granted_actors: Vec<String>,
57 grant_unattributed: bool,
58 deny_writes_for: Vec<String>,
59 ) -> Self {
60 let mut gate = Self::new(granted_actors, grant_unattributed);
61 gate.invalid_write_policy = Self::validate_write_denials(&deny_writes_for).is_err();
62 gate.deny_writes_for = deny_writes_for.into_iter().collect();
63 if !gate.deny_writes_for.is_empty() || gate.invalid_write_policy {
64 gate.configuration_fingerprint = write_policy_fingerprint(
65 &gate.configuration_fingerprint,
66 OPERATION_CLASSIFIER_VERSION,
67 gate.invalid_write_policy,
68 &gate.deny_writes_for,
69 );
70 }
71 gate
72 }
73
74 pub fn validate_write_denials(patterns: &[String]) -> Result<(), GateValidationError> {
76 if patterns.len() > 256 {
77 return Err(GateValidationError::InvalidWriteDenyPatterns(
78 "at most 256 patterns are allowed".into(),
79 ));
80 }
81 for pattern in patterns {
82 if pattern.trim().is_empty() || pattern.len() > 256 {
83 return Err(GateValidationError::InvalidWriteDenyPatterns(
84 "each pattern must be non-blank and at most 256 UTF-8 bytes".into(),
85 ));
86 }
87 }
88 Ok(())
89 }
90
91 fn actor_is_granted(&self, req: &GateRequest) -> bool {
92 if req.actor.is_anonymous() {
93 self.grant_unattributed
94 } else {
95 self.granted_actors.contains(&req.actor.id)
96 }
97 }
98}
99
100fn write_policy_fingerprint(
101 enrollment: &str,
102 classifier_version: &str,
103 invalid: bool,
104 patterns: &BTreeSet<String>,
105) -> String {
106 let mut hasher = Sha256::new();
107 hasher.update(b"khive.caller-write-denials.v1\0");
108 for field in [enrollment, classifier_version] {
109 hasher.update((field.len() as u64).to_be_bytes());
110 hasher.update(field.as_bytes());
111 }
112 hasher.update([u8::from(invalid)]);
113 hasher.update((patterns.len() as u64).to_be_bytes());
114 for pattern in patterns {
115 hasher.update((pattern.len() as u64).to_be_bytes());
116 hasher.update(pattern.as_bytes());
117 }
118 format!("sha256:{:x}", hasher.finalize())
119}
120
121#[cfg(test)]
122#[path = "write_denials_tests.rs"]
123mod write_denials_tests;
124
125impl fmt::Debug for CallerEnrollmentGate {
126 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
127 f.debug_struct("CallerEnrollmentGate")
128 .field("granted_actor_count", &self.granted_actors.len())
129 .field("grant_unattributed", &self.grant_unattributed)
130 .field("write_denial_pattern_count", &self.deny_writes_for.len())
131 .field("invalid_write_policy", &self.invalid_write_policy)
132 .finish_non_exhaustive()
133 }
134}
135
136impl Gate for CallerEnrollmentGate {
137 fn check(&self, req: &GateRequest) -> Result<GateDecision, GateError> {
138 if self.invalid_write_policy {
139 return Err(GateError::Policy(
140 "invalid deny_writes_for configuration".into(),
141 ));
142 }
143 if self.actor_is_granted(req) {
144 if self
145 .deny_writes_for
146 .iter()
147 .any(|pattern| actor_matches(pattern, &req.actor.id))
148 && classify_operation(&req.verb) != Some(OperationAccess::Read)
149 {
150 return Ok(GateDecision::deny(
151 "[gate].deny_writes_for denies this operation",
152 ));
153 }
154 return Ok(GateDecision::allow());
155 }
156 let reason = if req.actor.is_anonymous() {
157 "unattributed caller is not enrolled"
158 } else {
159 "actor is not enrolled"
160 };
161 Ok(GateDecision::deny(reason))
162 }
163
164 fn impl_name(&self) -> &'static str {
165 "CallerEnrollmentGate"
166 }
167
168 fn configuration_fingerprint(&self) -> Option<&str> {
169 Some(&self.configuration_fingerprint)
170 }
171}
172
173fn actor_matches(pattern: &str, actor: &str) -> bool {
175 let pattern = pattern.as_bytes();
176 let actor = actor.as_bytes();
177 let (mut p, mut a, mut star, mut retry) = (0, 0, None, 0);
178 while a < actor.len() {
179 if p < pattern.len() && pattern[p] == b'*' {
180 star = Some(p);
181 p += 1;
182 retry = a;
183 } else if p < pattern.len() && pattern[p] == actor[a] {
184 p += 1;
185 a += 1;
186 } else if let Some(last_star) = star {
187 retry += 1;
188 a = retry;
189 p = last_star + 1;
190 } else {
191 return false;
192 }
193 }
194 while p < pattern.len() && pattern[p] == b'*' {
195 p += 1;
196 }
197 p == pattern.len()
198}