use kcode_k1_access_types::{
AccessId, Authorizations, GroupId, ModelId, OwnerSubject, SubsystemId, Target, TxId, UserId,
ViewerSubject,
};
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum OwnerWitness {
User,
Group(GroupId),
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum AccessAction {
Create {
target: Target,
authorizations: Authorizations,
},
Replace {
access_id: AccessId,
actor: UserId,
groups_revision: Option<TxId>,
witness: OwnerWitness,
authorizations: Authorizations,
},
}
fn error(message: &str) -> String {
message.to_owned()
}
fn reserve<T>(values: &mut Vec<T>, count: usize) -> Result<(), String> {
values
.try_reserve_exact(count)
.map_err(|_| error("allocation unavailable"))
}
fn wire_len(value: usize) -> Result<u64, String> {
u64::try_from(value).map_err(|_| error("length overflow"))
}
fn add(total: &mut usize, value: usize) -> Result<(), String> {
*total = total
.checked_add(value)
.ok_or_else(|| error("length overflow"))?;
Ok(())
}
fn authorization_len(authorizations: &Authorizations) -> Result<usize, String> {
wire_len(authorizations.owners().len())?;
wire_len(authorizations.viewers().len())?;
let owner_bytes = authorizations
.owners()
.len()
.checked_mul(13)
.ok_or_else(|| error("length overflow"))?;
let mut total = 16;
add(&mut total, owner_bytes)?;
for viewer in authorizations.viewers() {
let width = match viewer {
ViewerSubject::User(_) | ViewerSubject::Group(_) => 13,
ViewerSubject::Model(_) => 33,
};
add(&mut total, width)?;
}
Ok(total)
}
fn encoded_len(action: &AccessAction) -> Result<usize, String> {
let (mut total, authorizations) = match action {
AccessAction::Create {
target,
authorizations,
} => {
wire_len(target.object_id().len())?;
let mut total = 46;
add(&mut total, target.object_id().len())?;
(total, authorizations)
}
AccessAction::Replace {
groups_revision,
witness,
authorizations,
..
} => {
let mut total = 44;
if groups_revision.is_some() {
add(&mut total, 12)?;
}
if matches!(witness, OwnerWitness::Group(_)) {
add(&mut total, 12)?;
}
(total, authorizations)
}
};
add(&mut total, authorization_len(authorizations)?)?;
Ok(total)
}
fn put_len(output: &mut Vec<u8>, value: usize) -> Result<(), String> {
output.extend_from_slice(&wire_len(value)?.to_le_bytes());
Ok(())
}
fn put_authorizations(output: &mut Vec<u8>, authorizations: &Authorizations) -> Result<(), String> {
put_len(output, authorizations.owners().len())?;
for owner in authorizations.owners() {
match owner {
OwnerSubject::User(id) => {
output.push(1);
output.extend_from_slice(id.as_tx_id().as_bytes());
}
OwnerSubject::Group(id) => {
output.push(2);
output.extend_from_slice(id.txid().as_bytes());
}
}
}
put_len(output, authorizations.viewers().len())?;
for viewer in authorizations.viewers() {
match viewer {
ViewerSubject::User(id) => {
output.push(1);
output.extend_from_slice(id.as_tx_id().as_bytes());
}
ViewerSubject::Group(id) => {
output.push(2);
output.extend_from_slice(id.txid().as_bytes());
}
ViewerSubject::Model(id) => {
output.push(3);
output.extend_from_slice(id.as_bytes());
}
}
}
Ok(())
}
pub fn encode(operation_id: [u8; 16], action: &AccessAction) -> Result<Vec<u8>, String> {
let length = encoded_len(action)?;
let mut output = Vec::new();
reserve(&mut output, length)?;
output.push(1);
output.push(match action {
AccessAction::Create { .. } => 1,
AccessAction::Replace { .. } => 2,
});
output.extend_from_slice(&operation_id);
match action {
AccessAction::Create {
target,
authorizations,
} => {
output.extend_from_slice(target.subsystem().as_bytes());
put_len(&mut output, target.object_id().len())?;
output.extend_from_slice(target.object_id());
put_authorizations(&mut output, authorizations)?;
}
AccessAction::Replace {
access_id,
actor,
groups_revision,
witness,
authorizations,
} => {
output.extend_from_slice(access_id.txid().as_bytes());
output.extend_from_slice(actor.as_tx_id().as_bytes());
match groups_revision {
None => output.push(0),
Some(revision) => {
output.push(1);
output.extend_from_slice(revision.as_bytes());
}
}
match witness {
OwnerWitness::User => output.push(1),
OwnerWitness::Group(group) => {
output.push(2);
output.extend_from_slice(group.txid().as_bytes());
}
}
put_authorizations(&mut output, authorizations)?;
}
}
debug_assert_eq!(output.len(), length);
Ok(output)
}
struct Reader<'a> {
bytes: &'a [u8],
position: usize,
}
impl<'a> Reader<'a> {
fn new(bytes: &'a [u8]) -> Self {
Self { bytes, position: 0 }
}
fn remaining(&self) -> usize {
self.bytes.len() - self.position
}
fn take(&mut self, length: usize) -> Result<&'a [u8], String> {
let end = self
.position
.checked_add(length)
.ok_or_else(|| error("length overflow"))?;
let value = self
.bytes
.get(self.position..end)
.ok_or_else(|| error("truncated payload"))?;
self.position = end;
Ok(value)
}
fn byte(&mut self) -> Result<u8, String> {
Ok(self.take(1)?[0])
}
fn array<const N: usize>(&mut self) -> Result<[u8; N], String> {
let mut value = [0; N];
value.copy_from_slice(self.take(N)?);
Ok(value)
}
fn length(&mut self) -> Result<usize, String> {
usize::try_from(u64::from_le_bytes(self.array()?)).map_err(|_| error("length overflow"))
}
fn count(&mut self, minimum_entry: usize) -> Result<usize, String> {
let count = self.length()?;
let minimum = count
.checked_mul(minimum_entry)
.ok_or_else(|| error("count overflow"))?;
if minimum > self.remaining() {
return Err(error("truncated subjects"));
}
Ok(count)
}
fn vector(&mut self, length: usize) -> Result<Vec<u8>, String> {
if length > self.remaining() {
return Err(error("truncated object id"));
}
let bytes = self.take(length)?;
let mut value = Vec::new();
reserve(&mut value, length)?;
value.extend_from_slice(bytes);
Ok(value)
}
fn finish(self) -> Result<(), String> {
if self.position == self.bytes.len() {
Ok(())
} else {
Err(error("trailing bytes"))
}
}
}
fn clone_fallible<T: Clone>(values: &[T]) -> Result<Vec<T>, String> {
let mut copy = Vec::new();
reserve(&mut copy, values.len())?;
copy.extend_from_slice(values);
Ok(copy)
}
fn parse_authorizations(reader: &mut Reader<'_>) -> Result<Authorizations, String> {
let owner_count = reader.count(13)?;
let mut owners = Vec::new();
reserve(&mut owners, owner_count)?;
for _ in 0..owner_count {
owners.push(match reader.byte()? {
1 => OwnerSubject::User(UserId::from_tx_id(TxId::from_bytes(reader.array()?))),
2 => OwnerSubject::Group(GroupId::new(TxId::from_bytes(reader.array()?))),
_ => return Err(error("invalid owner kind")),
});
}
let viewer_count = reader.count(13)?;
let mut viewers = Vec::new();
reserve(&mut viewers, viewer_count)?;
for _ in 0..viewer_count {
viewers.push(match reader.byte()? {
1 => ViewerSubject::User(UserId::from_tx_id(TxId::from_bytes(reader.array()?))),
2 => ViewerSubject::Group(GroupId::new(TxId::from_bytes(reader.array()?))),
3 => ViewerSubject::Model(ModelId::from_bytes(reader.array()?)),
_ => return Err(error("invalid viewer kind")),
});
}
let canonical = Authorizations::new(clone_fallible(&owners)?, clone_fallible(&viewers)?)?;
if canonical.owners() != owners.as_slice() || canonical.viewers() != viewers.as_slice() {
return Err(error("noncanonical authorizations"));
}
Ok(canonical)
}
pub fn decode(payload: &[u8]) -> Result<([u8; 16], AccessAction), String> {
let mut reader = Reader::new(payload);
if reader.byte()? != 1 {
return Err(error("unsupported version"));
}
let kind = reader.byte()?;
let operation_id = reader.array()?;
let action = match kind {
1 => {
let subsystem = SubsystemId::from_bytes(reader.array()?)
.map_err(|value| format!("invalid subsystem: {value}"))?;
let object_length = reader.length()?;
let target = Target::new(subsystem, reader.vector(object_length)?);
let authorizations = parse_authorizations(&mut reader)?;
AccessAction::Create {
target,
authorizations,
}
}
2 => {
let access_id = AccessId::new(TxId::from_bytes(reader.array()?));
let actor = UserId::from_tx_id(TxId::from_bytes(reader.array()?));
let groups_revision = match reader.byte()? {
0 => None,
1 => Some(TxId::from_bytes(reader.array()?)),
_ => return Err(error("invalid revision presence")),
};
let witness = match reader.byte()? {
1 => OwnerWitness::User,
2 => OwnerWitness::Group(GroupId::new(TxId::from_bytes(reader.array()?))),
_ => return Err(error("invalid witness kind")),
};
let authorizations = parse_authorizations(&mut reader)?;
AccessAction::Replace {
access_id,
actor,
groups_revision,
witness,
authorizations,
}
}
_ => return Err(error("invalid action kind")),
};
reader.finish()?;
Ok((operation_id, action))
}
#[cfg(test)]
mod tests {
use super::*;
fn tx(value: u8) -> TxId {
TxId::from_bytes([value; 12])
}
fn user(value: u8) -> UserId {
UserId::from_tx_id(tx(value))
}
fn group(value: u8) -> GroupId {
GroupId::new(tx(value))
}
fn authorizations() -> Authorizations {
Authorizations::new(
vec![OwnerSubject::User(user(1)), OwnerSubject::Group(group(2))],
vec![
ViewerSubject::User(user(3)),
ViewerSubject::Group(group(4)),
ViewerSubject::Model(ModelId::from_bytes([5; 32])),
],
)
.unwrap()
}
fn raw_authorizations(owners: &[(u8, u8)], viewers: &[(u8, u8)]) -> Vec<u8> {
let subsystem = SubsystemId::from_str("s").unwrap();
let mut payload = vec![1, 1];
payload.extend_from_slice(&[0; 16]);
payload.extend_from_slice(subsystem.as_bytes());
payload.extend_from_slice(&0u64.to_le_bytes());
payload.extend_from_slice(&u64::try_from(owners.len()).unwrap().to_le_bytes());
for &(kind, id) in owners {
payload.push(kind);
payload.resize(payload.len() + 12, id);
}
payload.extend_from_slice(&u64::try_from(viewers.len()).unwrap().to_le_bytes());
for &(kind, id) in viewers {
payload.push(kind);
let width = if kind == 3 { 32 } else { 12 };
payload.resize(payload.len() + width, id);
}
payload
}
fn altered(mut payload: Vec<u8>, position: usize, value: u8) -> Vec<u8> {
payload[position] = value;
payload
}
#[test]
fn create_round_trip_preserves_arbitrary_object_and_padding() {
let subsystem = SubsystemId::from_str("padded").unwrap();
let object = vec![0, 255, 128, 0, 1];
let action = AccessAction::Create {
target: Target::new(subsystem, object.clone()),
authorizations: authorizations(),
};
let payload = encode([7; 16], &action).unwrap();
assert_eq!(&payload[18..38], subsystem.as_bytes());
assert!(payload[24..38].iter().all(|byte| *byte == 0));
assert_eq!(&payload[46..51], object.as_slice());
assert_eq!(decode(&payload).unwrap(), ([7; 16], action));
}
#[test]
fn replace_round_trips_direct_group_and_revision() {
let actions = [
AccessAction::Replace {
access_id: AccessId::new(tx(6)),
actor: user(1),
groups_revision: None,
witness: OwnerWitness::User,
authorizations: authorizations(),
},
AccessAction::Replace {
access_id: AccessId::new(tx(7)),
actor: user(1),
groups_revision: Some(tx(8)),
witness: OwnerWitness::Group(group(2)),
authorizations: authorizations(),
},
];
for action in actions {
let payload = encode([9; 16], &action).unwrap();
assert_eq!(decode(&payload).unwrap(), ([9; 16], action));
}
}
#[test]
fn rejects_malformed_frames_and_discriminants() {
let valid = raw_authorizations(&[(1, 1)], &[]);
for end in 0..valid.len() {
assert!(decode(&valid[..end]).is_err());
}
assert!(decode(&altered(valid.clone(), 0, 2)).is_err());
assert!(decode(&altered(valid.clone(), 1, 9)).is_err());
let mut trailing = valid.clone();
trailing.push(0);
assert!(decode(&trailing).is_err());
assert!(decode(&altered(valid, 20, 1)).is_err());
let direct = AccessAction::Replace {
access_id: AccessId::new(tx(1)),
actor: user(1),
groups_revision: None,
witness: OwnerWitness::User,
authorizations: authorizations(),
};
let payload = encode([0; 16], &direct).unwrap();
assert!(decode(&altered(payload.clone(), 42, 2)).is_err());
assert!(decode(&altered(payload, 43, 9)).is_err());
}
#[test]
fn rejects_noncanonical_authorizations() {
let cases = [
raw_authorizations(&[], &[]),
raw_authorizations(&[(1, 2), (1, 1)], &[]),
raw_authorizations(&[(1, 1), (1, 1)], &[]),
raw_authorizations(&[(1, 1)], &[(1, 1)]),
raw_authorizations(&[(2, 1)], &[(2, 1)]),
raw_authorizations(&[(1, 1)], &[(2, 2), (1, 3)]),
raw_authorizations(&[(1, 1)], &[(3, 2), (3, 2)]),
raw_authorizations(&[(9, 1)], &[]),
raw_authorizations(&[(1, 1)], &[(9, 1)]),
];
for payload in cases {
assert!(decode(&payload).is_err());
}
}
#[test]
fn rejects_overflowing_lengths_and_counts() {
let valid = raw_authorizations(&[(1, 1)], &[]);
let mut object = valid.clone();
object[38..46].copy_from_slice(&u64::MAX.to_le_bytes());
assert!(decode(&object).is_err());
let mut owners = valid.clone();
owners[46..54].copy_from_slice(&u64::MAX.to_le_bytes());
assert!(decode(&owners).is_err());
let mut viewers = valid;
viewers[67..75].copy_from_slice(&u64::MAX.to_le_bytes());
assert!(decode(&viewers).is_err());
}
}