Skip to main content

kcode_k1_access_format/
lib.rs

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}