kcode-k1-rust-transaction 0.2.0

Encode canonical K1 Rust source package transactions
Documentation
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"
        );
    }
}