kcode-k1-rust-transaction 0.2.1

Encode canonical K1 Rust source-state transactions
Documentation
use kcode_k1_rust_package::{
    AuthorityId, LibraryFamily, LibraryId, PackageError, SourceFile, SourcePackage,
};
use kcode_k1_rust_worktree::UnpublishedId;
use kcode_k1_transaction_id::TxId;
use semver::Version;
use std::fmt::{Display, Formatter};

pub const WIRE_VERSION: u8 = 1;
const CREATE: u8 = 0;
const FORK: u8 = 1;
const OVERWRITE: u8 = 2;
const PUBLISH: u8 = 3;
const PACKAGE_HEADER_LENGTH: usize = 45;

#[derive(Clone, Debug, Eq, PartialEq)]
pub enum RustSourceTransaction {
    Create(SourcePackage),
    Fork(SourcePackage),
    Overwrite {
        id: UnpublishedId,
        expected_revision: TxId,
        source: SourcePackage,
    },
    Publish(SourcePackage),
}

impl RustSourceTransaction {
    pub fn source(&self) -> &SourcePackage {
        match self {
            Self::Create(source)
            | Self::Fork(source)
            | Self::Overwrite { source, .. }
            | Self::Publish(source) => source,
        }
    }

    pub fn family(&self) -> &LibraryFamily {
        self.source().id().family()
    }
}

#[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(event: &RustSourceTransaction) -> Result<Vec<u8>, TransactionError> {
    let metadata = matches!(event, RustSourceTransaction::Overwrite { .. })
        .then_some(24)
        .unwrap_or(0);
    let total = encoded_package_length(event.source())?
        .checked_add(2 + metadata)
        .ok_or_else(|| error("encoded length overflows"))?;
    let mut output = Vec::new();
    reserve(&mut output, total)?;
    output.push(WIRE_VERSION);
    match event {
        RustSourceTransaction::Create(source) => {
            output.push(CREATE);
            encode_package(source, &mut output)?;
        }
        RustSourceTransaction::Fork(source) => {
            output.push(FORK);
            encode_package(source, &mut output)?;
        }
        RustSourceTransaction::Overwrite {
            id,
            expected_revision,
            source,
        } => {
            output.push(OVERWRITE);
            output.extend_from_slice(id.transaction().as_bytes());
            output.extend_from_slice(expected_revision.as_bytes());
            encode_package(source, &mut output)?;
        }
        RustSourceTransaction::Publish(source) => {
            output.push(PUBLISH);
            encode_package(source, &mut output)?;
        }
    }
    Ok(output)
}

pub fn decode(bytes: &[u8]) -> Result<RustSourceTransaction, TransactionError> {
    let mut reader = Reader { bytes, offset: 0 };
    if reader.byte()? != WIRE_VERSION {
        return Err(error("unknown wire version"));
    }
    let kind = reader.byte()?;
    let metadata = if kind == OVERWRITE {
        Some((
            UnpublishedId::new(TxId::from_bytes(reader.array()?)),
            TxId::from_bytes(reader.array()?),
        ))
    } else {
        None
    };
    let source = decode_package(&mut reader)?;
    if reader.offset != bytes.len() {
        return Err(error("trailing bytes"));
    }
    match (kind, metadata) {
        (CREATE, None) => Ok(RustSourceTransaction::Create(source)),
        (FORK, None) => Ok(RustSourceTransaction::Fork(source)),
        (OVERWRITE, Some((id, expected_revision))) => Ok(RustSourceTransaction::Overwrite {
            id,
            expected_revision,
            source,
        }),
        (PUBLISH, None) => Ok(RustSourceTransaction::Publish(source)),
        _ => Err(error("unknown transaction kind")),
    }
}

fn encoded_package_length(package: &SourcePackage) -> Result<usize, TransactionError> {
    let mut total = PACKAGE_HEADER_LENGTH
        .checked_add(package.id().family().logical_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"))?;
    }
    Ok(total)
}

fn encode_package(package: &SourcePackage, output: &mut Vec<u8>) -> Result<(), 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"))?;
    output.extend_from_slice(family.authority().transaction_id().as_bytes());
    output.push(name_length);
    output.extend_from_slice(name);
    for value in [
        package.id().version().major,
        package.id().version().minor,
        package.id().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(())
}

fn decode_package(reader: &mut Reader<'_>) -> Result<SourcePackage, TransactionError> {
    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"))?;
        files.push(SourceFile::new(path, reader.bytes(content_length)?));
    }
    let family = LibraryFamily::new(authority, name)?;
    let identity = LibraryId::new(family, version)?;
    SourcePackage::new(identity, 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 family =
            LibraryFamily::new(AuthorityId::new(TxId::from_bytes([1; 12])), "alpha").unwrap();
        let identity = 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(
            identity,
            vec![
                SourceFile::new("Cargo.toml", manifest.as_bytes().to_vec()),
                SourceFile::new("Documentation.md", b"docs".to_vec()),
                SourceFile::new("src/lib.rs", vec![0, 159, 255]),
            ],
        )
        .unwrap()
    }

    fn events() -> Vec<RustSourceTransaction> {
        let source = package();
        vec![
            RustSourceTransaction::Create(source.clone()),
            RustSourceTransaction::Fork(source.clone()),
            RustSourceTransaction::Overwrite {
                id: UnpublishedId::new(TxId::from_bytes([2; 12])),
                expected_revision: TxId::from_bytes([3; 12]),
                source: source.clone(),
            },
            RustSourceTransaction::Publish(source),
        ]
    }

    #[test]
    fn every_event_round_trips_deterministically() {
        for event in events() {
            let wire = encode(&event).unwrap();
            assert_eq!(wire, encode(&event).unwrap());
            assert_eq!(decode(&wire).unwrap(), event);
        }
    }

    #[test]
    fn rejects_unknown_truncated_and_trailing_data() {
        let wire = encode(&events().remove(0)).unwrap();
        let mut unknown_version = wire.clone();
        unknown_version[0] = 2;
        assert_eq!(
            decode(&unknown_version).unwrap_err().message(),
            "unknown wire version"
        );
        let mut unknown_kind = wire.clone();
        unknown_kind[1] = 9;
        assert_eq!(
            decode(&unknown_kind).unwrap_err().message(),
            "unknown transaction kind"
        );
        for end in 0..wire.len() {
            assert!(decode(&wire[..end]).is_err(), "accepted prefix {end}");
        }
        let mut trailing = wire;
        trailing.push(0);
        assert_eq!(decode(&trailing).unwrap_err().message(), "trailing bytes");
    }
}