1use kcode_k1_access_types::{
2 AccessId, Authorizations, GroupId, ModelId, OwnerSubject, SubsystemId, Target, TxId, UserId,
3 ViewerSubject,
4};
5#[derive(Clone, Debug, Eq, PartialEq)]
6pub enum OwnerWitness {
7 User,
8 Group(GroupId),
9}
10#[derive(Clone, Debug, Eq, PartialEq)]
11pub enum AccessAction {
12 Create {
13 target: Target,
14 authorizations: Authorizations,
15 },
16 Replace {
17 access_id: AccessId,
18 actor: UserId,
19 groups_revision: Option<TxId>,
20 witness: OwnerWitness,
21 authorizations: Authorizations,
22 },
23 EnsureDiscovery {
24 access_id: AccessId,
25 },
26}
27fn error(message: &str) -> String {
28 message.to_owned()
29}
30fn reserve<T>(values: &mut Vec<T>, count: usize) -> Result<(), String> {
31 values
32 .try_reserve_exact(count)
33 .map_err(|_| error("allocation unavailable"))
34}
35fn wire_len(value: usize) -> Result<u64, String> {
36 u64::try_from(value).map_err(|_| error("length overflow"))
37}
38fn add(total: &mut usize, value: usize) -> Result<(), String> {
39 *total = total
40 .checked_add(value)
41 .ok_or_else(|| error("length overflow"))?;
42 Ok(())
43}
44fn authorization_len(authorizations: &Authorizations) -> Result<usize, String> {
45 wire_len(authorizations.owners().len())?;
46 wire_len(authorizations.viewers().len())?;
47 let owner_bytes = authorizations
48 .owners()
49 .len()
50 .checked_mul(13)
51 .ok_or_else(|| error("length overflow"))?;
52 let mut total = 16;
53 add(&mut total, owner_bytes)?;
54 for viewer in authorizations.viewers() {
55 let width = match viewer {
56 ViewerSubject::User(_) | ViewerSubject::Group(_) => 13,
57 ViewerSubject::Model(_) => 33,
58 };
59 add(&mut total, width)?;
60 }
61 Ok(total)
62}
63fn encoded_len(action: &AccessAction) -> Result<usize, String> {
64 let (mut total, authorizations) = match action {
65 AccessAction::Create {
66 target,
67 authorizations,
68 } => {
69 wire_len(target.object_id().len())?;
70 let mut total = 46;
71 add(&mut total, target.object_id().len())?;
72 (total, authorizations)
73 }
74 AccessAction::Replace {
75 groups_revision,
76 witness,
77 authorizations,
78 ..
79 } => {
80 let mut total = 44;
81 if groups_revision.is_some() {
82 add(&mut total, 12)?;
83 }
84 if matches!(witness, OwnerWitness::Group(_)) {
85 add(&mut total, 12)?;
86 }
87 (total, authorizations)
88 }
89 AccessAction::EnsureDiscovery { .. } => return Ok(30),
90 };
91 add(&mut total, authorization_len(authorizations)?)?;
92 Ok(total)
93}
94fn put_len(output: &mut Vec<u8>, value: usize) -> Result<(), String> {
95 output.extend_from_slice(&wire_len(value)?.to_le_bytes());
96 Ok(())
97}
98fn put_authorizations(output: &mut Vec<u8>, authorizations: &Authorizations) -> Result<(), String> {
99 put_len(output, authorizations.owners().len())?;
100 for owner in authorizations.owners() {
101 match owner {
102 OwnerSubject::User(id) => {
103 output.push(1);
104 output.extend_from_slice(id.as_tx_id().as_bytes());
105 }
106 OwnerSubject::Group(id) => {
107 output.push(2);
108 output.extend_from_slice(id.txid().as_bytes());
109 }
110 }
111 }
112 put_len(output, authorizations.viewers().len())?;
113 for viewer in authorizations.viewers() {
114 match viewer {
115 ViewerSubject::User(id) => {
116 output.push(1);
117 output.extend_from_slice(id.as_tx_id().as_bytes());
118 }
119 ViewerSubject::Group(id) => {
120 output.push(2);
121 output.extend_from_slice(id.txid().as_bytes());
122 }
123 ViewerSubject::Model(id) => {
124 output.push(3);
125 output.extend_from_slice(id.as_bytes());
126 }
127 }
128 }
129 Ok(())
130}
131pub fn encode(operation_id: [u8; 16], action: &AccessAction) -> Result<Vec<u8>, String> {
132 let length = encoded_len(action)?;
133 let mut output = Vec::new();
134 reserve(&mut output, length)?;
135 output.push(1);
136 output.push(match action {
137 AccessAction::Create { .. } => 1,
138 AccessAction::Replace { .. } => 2,
139 AccessAction::EnsureDiscovery { .. } => 3,
140 });
141 output.extend_from_slice(&operation_id);
142 match action {
143 AccessAction::Create {
144 target,
145 authorizations,
146 } => {
147 output.extend_from_slice(target.subsystem().as_bytes());
148 put_len(&mut output, target.object_id().len())?;
149 output.extend_from_slice(target.object_id());
150 put_authorizations(&mut output, authorizations)?;
151 }
152 AccessAction::Replace {
153 access_id,
154 actor,
155 groups_revision,
156 witness,
157 authorizations,
158 } => {
159 output.extend_from_slice(access_id.txid().as_bytes());
160 output.extend_from_slice(actor.as_tx_id().as_bytes());
161 match groups_revision {
162 None => output.push(0),
163 Some(revision) => {
164 output.push(1);
165 output.extend_from_slice(revision.as_bytes());
166 }
167 }
168 match witness {
169 OwnerWitness::User => output.push(1),
170 OwnerWitness::Group(group) => {
171 output.push(2);
172 output.extend_from_slice(group.txid().as_bytes());
173 }
174 }
175 put_authorizations(&mut output, authorizations)?;
176 }
177 AccessAction::EnsureDiscovery { access_id } => {
178 output.extend_from_slice(access_id.txid().as_bytes())
179 }
180 }
181 debug_assert_eq!(output.len(), length);
182 Ok(output)
183}
184struct Reader<'a> {
185 bytes: &'a [u8],
186 position: usize,
187}
188impl<'a> Reader<'a> {
189 fn new(bytes: &'a [u8]) -> Self {
190 Self { bytes, position: 0 }
191 }
192 fn remaining(&self) -> usize {
193 self.bytes.len() - self.position
194 }
195 fn take(&mut self, length: usize) -> Result<&'a [u8], String> {
196 let end = self
197 .position
198 .checked_add(length)
199 .ok_or_else(|| error("length overflow"))?;
200 let value = self
201 .bytes
202 .get(self.position..end)
203 .ok_or_else(|| error("truncated payload"))?;
204 self.position = end;
205 Ok(value)
206 }
207 fn byte(&mut self) -> Result<u8, String> {
208 Ok(self.take(1)?[0])
209 }
210 fn array<const N: usize>(&mut self) -> Result<[u8; N], String> {
211 let mut value = [0; N];
212 value.copy_from_slice(self.take(N)?);
213 Ok(value)
214 }
215 fn length(&mut self) -> Result<usize, String> {
216 usize::try_from(u64::from_le_bytes(self.array()?)).map_err(|_| error("length overflow"))
217 }
218 fn count(&mut self, minimum_entry: usize) -> Result<usize, String> {
219 let count = self.length()?;
220 let minimum = count
221 .checked_mul(minimum_entry)
222 .ok_or_else(|| error("count overflow"))?;
223 if minimum > self.remaining() {
224 return Err(error("truncated subjects"));
225 }
226 Ok(count)
227 }
228 fn vector(&mut self, length: usize) -> Result<Vec<u8>, String> {
229 if length > self.remaining() {
230 return Err(error("truncated object id"));
231 }
232 let bytes = self.take(length)?;
233 let mut value = Vec::new();
234 reserve(&mut value, length)?;
235 value.extend_from_slice(bytes);
236 Ok(value)
237 }
238 fn finish(self) -> Result<(), String> {
239 if self.position == self.bytes.len() {
240 Ok(())
241 } else {
242 Err(error("trailing bytes"))
243 }
244 }
245}
246fn clone_fallible<T: Clone>(values: &[T]) -> Result<Vec<T>, String> {
247 let mut copy = Vec::new();
248 reserve(&mut copy, values.len())?;
249 copy.extend_from_slice(values);
250 Ok(copy)
251}
252fn parse_authorizations(reader: &mut Reader<'_>) -> Result<Authorizations, String> {
253 let owner_count = reader.count(13)?;
254 let mut owners = Vec::new();
255 reserve(&mut owners, owner_count)?;
256 for _ in 0..owner_count {
257 owners.push(match reader.byte()? {
258 1 => OwnerSubject::User(UserId::from_tx_id(TxId::from_bytes(reader.array()?))),
259 2 => OwnerSubject::Group(GroupId::new(TxId::from_bytes(reader.array()?))),
260 _ => return Err(error("invalid owner kind")),
261 });
262 }
263 let viewer_count = reader.count(13)?;
264 let mut viewers = Vec::new();
265 reserve(&mut viewers, viewer_count)?;
266 for _ in 0..viewer_count {
267 viewers.push(match reader.byte()? {
268 1 => ViewerSubject::User(UserId::from_tx_id(TxId::from_bytes(reader.array()?))),
269 2 => ViewerSubject::Group(GroupId::new(TxId::from_bytes(reader.array()?))),
270 3 => ViewerSubject::Model(ModelId::from_bytes(reader.array()?)),
271 _ => return Err(error("invalid viewer kind")),
272 });
273 }
274 let canonical = Authorizations::new(clone_fallible(&owners)?, clone_fallible(&viewers)?)?;
275 if canonical.owners() != owners.as_slice() || canonical.viewers() != viewers.as_slice() {
276 return Err(error("noncanonical authorizations"));
277 }
278 Ok(canonical)
279}
280pub fn decode(payload: &[u8]) -> Result<([u8; 16], AccessAction), String> {
281 let mut reader = Reader::new(payload);
282 if reader.byte()? != 1 {
283 return Err(error("unsupported version"));
284 }
285 let kind = reader.byte()?;
286 let operation_id = reader.array()?;
287 let action = match kind {
288 1 => {
289 let subsystem = SubsystemId::from_bytes(reader.array()?)
290 .map_err(|value| format!("invalid subsystem: {value}"))?;
291 let object_length = reader.length()?;
292 let target = Target::new(subsystem, reader.vector(object_length)?);
293 let authorizations = parse_authorizations(&mut reader)?;
294 AccessAction::Create {
295 target,
296 authorizations,
297 }
298 }
299 2 => {
300 let access_id = AccessId::new(TxId::from_bytes(reader.array()?));
301 let actor = UserId::from_tx_id(TxId::from_bytes(reader.array()?));
302 let groups_revision = match reader.byte()? {
303 0 => None,
304 1 => Some(TxId::from_bytes(reader.array()?)),
305 _ => return Err(error("invalid revision presence")),
306 };
307 let witness = match reader.byte()? {
308 1 => OwnerWitness::User,
309 2 => OwnerWitness::Group(GroupId::new(TxId::from_bytes(reader.array()?))),
310 _ => return Err(error("invalid witness kind")),
311 };
312 let authorizations = parse_authorizations(&mut reader)?;
313 AccessAction::Replace {
314 access_id,
315 actor,
316 groups_revision,
317 witness,
318 authorizations,
319 }
320 }
321 3 => AccessAction::EnsureDiscovery {
322 access_id: AccessId::new(TxId::from_bytes(reader.array()?)),
323 },
324 _ => return Err(error("invalid action kind")),
325 };
326 reader.finish()?;
327 Ok((operation_id, action))
328}
329#[cfg(test)]
330mod tests {
331 use super::*;
332 fn tx(value: u8) -> TxId {
333 TxId::from_bytes([value; 12])
334 }
335 fn user(value: u8) -> UserId {
336 UserId::from_tx_id(tx(value))
337 }
338 fn group(value: u8) -> GroupId {
339 GroupId::new(tx(value))
340 }
341 fn authorizations() -> Authorizations {
342 Authorizations::new(
343 vec![OwnerSubject::User(user(1)), OwnerSubject::Group(group(2))],
344 vec![
345 ViewerSubject::User(user(3)),
346 ViewerSubject::Group(group(4)),
347 ViewerSubject::Model(ModelId::from_bytes([5; 32])),
348 ],
349 )
350 .unwrap()
351 }
352 fn raw_authorizations(owners: &[(u8, u8)], viewers: &[(u8, u8)]) -> Vec<u8> {
353 let subsystem = SubsystemId::from_str("s").unwrap();
354 let mut payload = vec![1, 1];
355 payload.extend_from_slice(&[0; 16]);
356 payload.extend_from_slice(subsystem.as_bytes());
357 payload.extend_from_slice(&0u64.to_le_bytes());
358 payload.extend_from_slice(&u64::try_from(owners.len()).unwrap().to_le_bytes());
359 for &(kind, id) in owners {
360 payload.push(kind);
361 payload.resize(payload.len() + 12, id);
362 }
363 payload.extend_from_slice(&u64::try_from(viewers.len()).unwrap().to_le_bytes());
364 for &(kind, id) in viewers {
365 payload.push(kind);
366 let width = if kind == 3 { 32 } else { 12 };
367 payload.resize(payload.len() + width, id);
368 }
369 payload
370 }
371 fn altered(mut payload: Vec<u8>, position: usize, value: u8) -> Vec<u8> {
372 payload[position] = value;
373 payload
374 }
375 #[test]
376 fn create_round_trip_preserves_arbitrary_object_and_padding() {
377 let subsystem = SubsystemId::from_str("padded").unwrap();
378 let object = vec![0, 255, 128, 0, 1];
379 let action = AccessAction::Create {
380 target: Target::new(subsystem, object.clone()),
381 authorizations: authorizations(),
382 };
383 let payload = encode([7; 16], &action).unwrap();
384 assert_eq!(&payload[18..38], subsystem.as_bytes());
385 assert!(payload[24..38].iter().all(|byte| *byte == 0));
386 assert_eq!(&payload[46..51], object.as_slice());
387 assert_eq!(decode(&payload).unwrap(), ([7; 16], action));
388 }
389 #[test]
390 fn replace_round_trips_direct_group_and_revision() {
391 let actions = [
392 AccessAction::Replace {
393 access_id: AccessId::new(tx(6)),
394 actor: user(1),
395 groups_revision: None,
396 witness: OwnerWitness::User,
397 authorizations: authorizations(),
398 },
399 AccessAction::Replace {
400 access_id: AccessId::new(tx(7)),
401 actor: user(1),
402 groups_revision: Some(tx(8)),
403 witness: OwnerWitness::Group(group(2)),
404 authorizations: authorizations(),
405 },
406 ];
407 for action in actions {
408 let payload = encode([9; 16], &action).unwrap();
409 assert_eq!(decode(&payload).unwrap(), ([9; 16], action));
410 }
411 }
412 #[test]
413 fn ensure_discovery_has_exact_wire_and_round_trips() {
414 let action = AccessAction::EnsureDiscovery {
415 access_id: AccessId::new(tx(6)),
416 };
417 let payload = encode([9; 16], &action).unwrap();
418 assert_eq!(payload.len(), 30);
419 assert_eq!(
420 &payload[..18],
421 &[1, 3, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9]
422 );
423 assert_eq!(&payload[18..], tx(6).as_bytes());
424 assert_eq!(decode(&payload).unwrap(), ([9; 16], action));
425 }
426 #[test]
427 fn ensure_discovery_rejects_malformed_and_trailing_frames() {
428 let payload = encode(
429 [0; 16],
430 &AccessAction::EnsureDiscovery {
431 access_id: AccessId::new(tx(1)),
432 },
433 )
434 .unwrap();
435 for end in 0..payload.len() {
436 assert!(decode(&payload[..end]).is_err());
437 }
438 let mut trailing = payload.clone();
439 trailing.push(0);
440 assert!(decode(&trailing).is_err());
441 assert!(decode(&altered(payload.clone(), 0, 2)).is_err());
442 assert!(decode(&altered(payload, 1, 9)).is_err());
443 }
444 #[test]
445 fn rejects_malformed_frames_and_discriminants() {
446 let valid = raw_authorizations(&[(1, 1)], &[]);
447 for end in 0..valid.len() {
448 assert!(decode(&valid[..end]).is_err());
449 }
450 assert!(decode(&altered(valid.clone(), 0, 2)).is_err());
451 assert!(decode(&altered(valid.clone(), 1, 9)).is_err());
452 let mut trailing = valid.clone();
453 trailing.push(0);
454 assert!(decode(&trailing).is_err());
455 assert!(decode(&altered(valid, 20, 1)).is_err());
456 let direct = AccessAction::Replace {
457 access_id: AccessId::new(tx(1)),
458 actor: user(1),
459 groups_revision: None,
460 witness: OwnerWitness::User,
461 authorizations: authorizations(),
462 };
463 let payload = encode([0; 16], &direct).unwrap();
464 assert!(decode(&altered(payload.clone(), 42, 2)).is_err());
465 assert!(decode(&altered(payload, 43, 9)).is_err());
466 }
467 #[test]
468 fn rejects_noncanonical_authorizations() {
469 let cases = [
470 raw_authorizations(&[], &[]),
471 raw_authorizations(&[(1, 2), (1, 1)], &[]),
472 raw_authorizations(&[(1, 1), (1, 1)], &[]),
473 raw_authorizations(&[(1, 1)], &[(1, 1)]),
474 raw_authorizations(&[(2, 1)], &[(2, 1)]),
475 raw_authorizations(&[(1, 1)], &[(2, 2), (1, 3)]),
476 raw_authorizations(&[(1, 1)], &[(3, 2), (3, 2)]),
477 raw_authorizations(&[(9, 1)], &[]),
478 raw_authorizations(&[(1, 1)], &[(9, 1)]),
479 ];
480 for payload in cases {
481 assert!(decode(&payload).is_err());
482 }
483 }
484 #[test]
485 fn rejects_overflowing_lengths_and_counts() {
486 let valid = raw_authorizations(&[(1, 1)], &[]);
487 let mut object = valid.clone();
488 object[38..46].copy_from_slice(&u64::MAX.to_le_bytes());
489 assert!(decode(&object).is_err());
490 let mut owners = valid.clone();
491 owners[46..54].copy_from_slice(&u64::MAX.to_le_bytes());
492 assert!(decode(&owners).is_err());
493 let mut viewers = valid;
494 viewers[67..75].copy_from_slice(&u64::MAX.to_le_bytes());
495 assert!(decode(&viewers).is_err());
496 }
497}