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");
}
}