use kcode_k1_rust_package::{
AuthorityId, LibraryFamily, LibraryId, PackageError, SourceFile, SourcePackage,
};
use kcode_k1_transaction_id::TxId;
use semver::Version;
use std::fmt::{Display, Formatter};
pub const WIRE_VERSION: u8 = 1;
const FIXED_HEADER_LENGTH: usize = 46;
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct TransactionError(ErrorMessage);
#[derive(Clone, Debug, Eq, PartialEq)]
enum ErrorMessage {
Static(&'static str),
Package(PackageError),
}
impl TransactionError {
pub fn message(&self) -> &str {
match &self.0 {
ErrorMessage::Static(message) => message,
ErrorMessage::Package(error) => error.message(),
}
}
}
impl Display for TransactionError {
fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
formatter.write_str(self.message())
}
}
impl std::error::Error for TransactionError {}
impl From<PackageError> for TransactionError {
fn from(error: PackageError) -> Self {
Self(ErrorMessage::Package(error))
}
}
pub fn encode(package: &SourcePackage) -> Result<Vec<u8>, TransactionError> {
let family = package.id().family();
let name = family.logical_name().as_bytes();
let name_length = u8::try_from(name.len()).map_err(|_| error("logical name is too long"))?;
let file_count = u64::try_from(package.files().len())
.map_err(|_| error("file count overflows the wire format"))?;
let mut total = FIXED_HEADER_LENGTH
.checked_add(name.len())
.ok_or_else(|| error("encoded length overflows"))?;
for file in package.files() {
u16::try_from(file.path().len())
.map_err(|_| error("path is too long for the wire format"))?;
u64::try_from(file.bytes().len())
.map_err(|_| error("content length overflows the wire format"))?;
total = total
.checked_add(10)
.and_then(|value| value.checked_add(file.path().len()))
.and_then(|value| value.checked_add(file.bytes().len()))
.ok_or_else(|| error("encoded length overflows"))?;
}
let mut output = Vec::new();
reserve(&mut output, total)?;
output.push(WIRE_VERSION);
output.extend_from_slice(family.authority().transaction_id().as_bytes());
output.push(name_length);
output.extend_from_slice(name);
let version = package.id().version();
for value in [version.major, version.minor, version.patch] {
output.extend_from_slice(&value.to_be_bytes());
}
output.extend_from_slice(&file_count.to_be_bytes());
for file in package.files() {
let path_length = u16::try_from(file.path().len())
.map_err(|_| error("path is too long for the wire format"))?;
let content_length = u64::try_from(file.bytes().len())
.map_err(|_| error("content length overflows the wire format"))?;
output.extend_from_slice(&path_length.to_be_bytes());
output.extend_from_slice(file.path().as_bytes());
output.extend_from_slice(&content_length.to_be_bytes());
output.extend_from_slice(file.bytes());
}
Ok(output)
}
pub fn decode(bytes: &[u8]) -> Result<SourcePackage, TransactionError> {
let mut reader = Reader { bytes, offset: 0 };
if reader.byte()? != WIRE_VERSION {
return Err(error("unknown wire version"));
}
let authority = AuthorityId::new(TxId::from_bytes(reader.array()?));
let name_length = usize::from(reader.byte()?);
let name = reader.string(name_length, "logical name is not UTF-8")?;
let version = Version::new(reader.u64()?, reader.u64()?, reader.u64()?);
let file_count =
usize::try_from(reader.u64()?).map_err(|_| error("file count overflows this platform"))?;
let mut files = Vec::new();
reserve(&mut files, file_count)?;
for _ in 0..file_count {
let path_length = usize::from(reader.u16()?);
let path = reader.string(path_length, "file path is not UTF-8")?;
if files
.last()
.is_some_and(|previous: &SourceFile| previous.path() >= path.as_str())
{
return Err(error("file paths are not in strictly increasing order"));
}
let content_length = usize::try_from(reader.u64()?)
.map_err(|_| error("content length overflows this platform"))?;
let content = reader.bytes(content_length)?;
files.push(SourceFile::new(path, content));
}
if reader.offset != bytes.len() {
return Err(error("trailing bytes"));
}
let family = LibraryFamily::new(authority, name)?;
let id = LibraryId::new(family, version)?;
SourcePackage::new(id, files).map_err(Into::into)
}
fn error(message: &'static str) -> TransactionError {
TransactionError(ErrorMessage::Static(message))
}
fn reserve<T>(values: &mut Vec<T>, additional: usize) -> Result<(), TransactionError> {
values
.try_reserve_exact(additional)
.map_err(|_| error("allocation failed"))
}
struct Reader<'a> {
bytes: &'a [u8],
offset: usize,
}
impl<'a> Reader<'a> {
fn take(&mut self, length: usize) -> Result<&'a [u8], TransactionError> {
let end = self
.offset
.checked_add(length)
.ok_or_else(|| error("decoded offset overflows"))?;
let value = self
.bytes
.get(self.offset..end)
.ok_or_else(|| error("truncated transaction"))?;
self.offset = end;
Ok(value)
}
fn byte(&mut self) -> Result<u8, TransactionError> {
Ok(self.take(1)?[0])
}
fn array<const N: usize>(&mut self) -> Result<[u8; N], TransactionError> {
let mut value = [0; N];
value.copy_from_slice(self.take(N)?);
Ok(value)
}
fn u16(&mut self) -> Result<u16, TransactionError> {
Ok(u16::from_be_bytes(self.array()?))
}
fn u64(&mut self) -> Result<u64, TransactionError> {
Ok(u64::from_be_bytes(self.array()?))
}
fn bytes(&mut self, length: usize) -> Result<Vec<u8>, TransactionError> {
let source = self.take(length)?;
let mut value = Vec::new();
reserve(&mut value, length)?;
value.extend_from_slice(source);
Ok(value)
}
fn string(&mut self, length: usize, invalid: &'static str) -> Result<String, TransactionError> {
String::from_utf8(self.bytes(length)?).map_err(|_| error(invalid))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn package() -> SourcePackage {
let authority = AuthorityId::new(TxId::from_bytes([1; 12]));
let family = LibraryFamily::new(authority, "alpha").unwrap();
let id = LibraryId::new(family, Version::new(1, 2, 3)).unwrap();
let manifest = r#"[package]
name = "k1-010101010101010101010101-alpha"
version = "1.2.3"
edition = "2024"
autobins = false
autoexamples = false
autotests = false
autobenches = false
[lib]
name = "alpha"
path = "src/lib.rs"
[workspace]
resolver = "3"
"#;
SourcePackage::new(
id,
vec![
SourceFile::new("src/lib.rs", vec![0, 159, 255]),
SourceFile::new("Documentation.md", b"docs".to_vec()),
SourceFile::new("Cargo.toml", manifest.as_bytes().to_vec()),
],
)
.unwrap()
}
#[test]
fn round_trips_exact_bytes_deterministically() {
let package = package();
let first = encode(&package).unwrap();
assert_eq!(first, encode(&package).unwrap());
let decoded = decode(&first).unwrap();
assert_eq!(decoded, package);
assert_eq!(decoded.files().last().unwrap().bytes(), [0, 159, 255]);
}
#[test]
fn rejects_versions_truncation_and_trailing_data() {
let wire = encode(&package()).unwrap();
let mut unknown = wire.clone();
unknown[0] = WIRE_VERSION + 1;
assert_eq!(
decode(&unknown).unwrap_err().message(),
"unknown wire version"
);
for end in 0..wire.len() {
assert!(
decode(&wire[..end]).is_err(),
"accepted prefix of length {end}"
);
}
let mut trailing = wire;
trailing.push(0);
assert_eq!(decode(&trailing).unwrap_err().message(), "trailing bytes");
}
#[test]
fn rejects_out_of_order_files() {
let package = package();
let wire = encode(&package).unwrap();
let header = FIXED_HEADER_LENGTH + package.id().family().logical_name().len();
let first = &package.files()[0];
let first_end = header + 10 + first.path().len() + first.bytes().len();
let second = &package.files()[1];
let second_end = first_end + 10 + second.path().len() + second.bytes().len();
let mut reordered = Vec::new();
reordered.extend_from_slice(&wire[..header]);
reordered.extend_from_slice(&wire[first_end..second_end]);
reordered.extend_from_slice(&wire[header..first_end]);
reordered.extend_from_slice(&wire[second_end..]);
assert_eq!(
decode(&reordered).unwrap_err().message(),
"file paths are not in strictly increasing order"
);
}
}