Skip to main content

kcode_k1_rust_transaction/
lib.rs

1use kcode_k1_rust_package::{
2    AuthorityId, LibraryFamily, LibraryId, PackageError, SourceFile, SourcePackage,
3};
4use kcode_k1_transaction_id::TxId;
5use semver::Version;
6use std::fmt::{Display, Formatter};
7
8pub const WIRE_VERSION: u8 = 1;
9const FIXED_HEADER_LENGTH: usize = 46;
10
11#[derive(Clone, Debug, Eq, PartialEq)]
12pub struct TransactionError(ErrorMessage);
13
14#[derive(Clone, Debug, Eq, PartialEq)]
15enum ErrorMessage {
16    Static(&'static str),
17    Package(PackageError),
18}
19
20impl TransactionError {
21    pub fn message(&self) -> &str {
22        match &self.0 {
23            ErrorMessage::Static(message) => message,
24            ErrorMessage::Package(error) => error.message(),
25        }
26    }
27}
28
29impl Display for TransactionError {
30    fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
31        formatter.write_str(self.message())
32    }
33}
34
35impl std::error::Error for TransactionError {}
36
37impl From<PackageError> for TransactionError {
38    fn from(error: PackageError) -> Self {
39        Self(ErrorMessage::Package(error))
40    }
41}
42
43pub fn encode(package: &SourcePackage) -> Result<Vec<u8>, TransactionError> {
44    let family = package.id().family();
45    let name = family.logical_name().as_bytes();
46    let name_length = u8::try_from(name.len()).map_err(|_| error("logical name is too long"))?;
47    let file_count = u64::try_from(package.files().len())
48        .map_err(|_| error("file count overflows the wire format"))?;
49    let mut total = FIXED_HEADER_LENGTH
50        .checked_add(name.len())
51        .ok_or_else(|| error("encoded length overflows"))?;
52    for file in package.files() {
53        u16::try_from(file.path().len())
54            .map_err(|_| error("path is too long for the wire format"))?;
55        u64::try_from(file.bytes().len())
56            .map_err(|_| error("content length overflows the wire format"))?;
57        total = total
58            .checked_add(10)
59            .and_then(|value| value.checked_add(file.path().len()))
60            .and_then(|value| value.checked_add(file.bytes().len()))
61            .ok_or_else(|| error("encoded length overflows"))?;
62    }
63
64    let mut output = Vec::new();
65    reserve(&mut output, total)?;
66    output.push(WIRE_VERSION);
67    output.extend_from_slice(family.authority().transaction_id().as_bytes());
68    output.push(name_length);
69    output.extend_from_slice(name);
70    let version = package.id().version();
71    for value in [version.major, version.minor, version.patch] {
72        output.extend_from_slice(&value.to_be_bytes());
73    }
74    output.extend_from_slice(&file_count.to_be_bytes());
75    for file in package.files() {
76        let path_length = u16::try_from(file.path().len())
77            .map_err(|_| error("path is too long for the wire format"))?;
78        let content_length = u64::try_from(file.bytes().len())
79            .map_err(|_| error("content length overflows the wire format"))?;
80        output.extend_from_slice(&path_length.to_be_bytes());
81        output.extend_from_slice(file.path().as_bytes());
82        output.extend_from_slice(&content_length.to_be_bytes());
83        output.extend_from_slice(file.bytes());
84    }
85    Ok(output)
86}
87
88pub fn decode(bytes: &[u8]) -> Result<SourcePackage, TransactionError> {
89    let mut reader = Reader { bytes, offset: 0 };
90    if reader.byte()? != WIRE_VERSION {
91        return Err(error("unknown wire version"));
92    }
93    let authority = AuthorityId::new(TxId::from_bytes(reader.array()?));
94    let name_length = usize::from(reader.byte()?);
95    let name = reader.string(name_length, "logical name is not UTF-8")?;
96    let version = Version::new(reader.u64()?, reader.u64()?, reader.u64()?);
97    let file_count =
98        usize::try_from(reader.u64()?).map_err(|_| error("file count overflows this platform"))?;
99    let mut files = Vec::new();
100    reserve(&mut files, file_count)?;
101    for _ in 0..file_count {
102        let path_length = usize::from(reader.u16()?);
103        let path = reader.string(path_length, "file path is not UTF-8")?;
104        if files
105            .last()
106            .is_some_and(|previous: &SourceFile| previous.path() >= path.as_str())
107        {
108            return Err(error("file paths are not in strictly increasing order"));
109        }
110        let content_length = usize::try_from(reader.u64()?)
111            .map_err(|_| error("content length overflows this platform"))?;
112        let content = reader.bytes(content_length)?;
113        files.push(SourceFile::new(path, content));
114    }
115    if reader.offset != bytes.len() {
116        return Err(error("trailing bytes"));
117    }
118    let family = LibraryFamily::new(authority, name)?;
119    let id = LibraryId::new(family, version)?;
120    SourcePackage::new(id, files).map_err(Into::into)
121}
122
123fn error(message: &'static str) -> TransactionError {
124    TransactionError(ErrorMessage::Static(message))
125}
126
127fn reserve<T>(values: &mut Vec<T>, additional: usize) -> Result<(), TransactionError> {
128    values
129        .try_reserve_exact(additional)
130        .map_err(|_| error("allocation failed"))
131}
132
133struct Reader<'a> {
134    bytes: &'a [u8],
135    offset: usize,
136}
137
138impl<'a> Reader<'a> {
139    fn take(&mut self, length: usize) -> Result<&'a [u8], TransactionError> {
140        let end = self
141            .offset
142            .checked_add(length)
143            .ok_or_else(|| error("decoded offset overflows"))?;
144        let value = self
145            .bytes
146            .get(self.offset..end)
147            .ok_or_else(|| error("truncated transaction"))?;
148        self.offset = end;
149        Ok(value)
150    }
151
152    fn byte(&mut self) -> Result<u8, TransactionError> {
153        Ok(self.take(1)?[0])
154    }
155
156    fn array<const N: usize>(&mut self) -> Result<[u8; N], TransactionError> {
157        let mut value = [0; N];
158        value.copy_from_slice(self.take(N)?);
159        Ok(value)
160    }
161
162    fn u16(&mut self) -> Result<u16, TransactionError> {
163        Ok(u16::from_be_bytes(self.array()?))
164    }
165
166    fn u64(&mut self) -> Result<u64, TransactionError> {
167        Ok(u64::from_be_bytes(self.array()?))
168    }
169
170    fn bytes(&mut self, length: usize) -> Result<Vec<u8>, TransactionError> {
171        let source = self.take(length)?;
172        let mut value = Vec::new();
173        reserve(&mut value, length)?;
174        value.extend_from_slice(source);
175        Ok(value)
176    }
177
178    fn string(&mut self, length: usize, invalid: &'static str) -> Result<String, TransactionError> {
179        String::from_utf8(self.bytes(length)?).map_err(|_| error(invalid))
180    }
181}
182
183#[cfg(test)]
184mod tests {
185    use super::*;
186
187    fn package() -> SourcePackage {
188        let authority = AuthorityId::new(TxId::from_bytes([1; 12]));
189        let family = LibraryFamily::new(authority, "alpha").unwrap();
190        let id = LibraryId::new(family, Version::new(1, 2, 3)).unwrap();
191        let manifest = r#"[package]
192name = "k1-010101010101010101010101-alpha"
193version = "1.2.3"
194edition = "2024"
195autobins = false
196autoexamples = false
197autotests = false
198autobenches = false
199
200[lib]
201name = "alpha"
202path = "src/lib.rs"
203
204[workspace]
205resolver = "3"
206"#;
207        SourcePackage::new(
208            id,
209            vec![
210                SourceFile::new("src/lib.rs", vec![0, 159, 255]),
211                SourceFile::new("Documentation.md", b"docs".to_vec()),
212                SourceFile::new("Cargo.toml", manifest.as_bytes().to_vec()),
213            ],
214        )
215        .unwrap()
216    }
217
218    #[test]
219    fn round_trips_exact_bytes_deterministically() {
220        let package = package();
221        let first = encode(&package).unwrap();
222        assert_eq!(first, encode(&package).unwrap());
223        let decoded = decode(&first).unwrap();
224        assert_eq!(decoded, package);
225        assert_eq!(decoded.files().last().unwrap().bytes(), [0, 159, 255]);
226    }
227
228    #[test]
229    fn rejects_versions_truncation_and_trailing_data() {
230        let wire = encode(&package()).unwrap();
231        let mut unknown = wire.clone();
232        unknown[0] = WIRE_VERSION + 1;
233        assert_eq!(
234            decode(&unknown).unwrap_err().message(),
235            "unknown wire version"
236        );
237        for end in 0..wire.len() {
238            assert!(
239                decode(&wire[..end]).is_err(),
240                "accepted prefix of length {end}"
241            );
242        }
243        let mut trailing = wire;
244        trailing.push(0);
245        assert_eq!(decode(&trailing).unwrap_err().message(), "trailing bytes");
246    }
247
248    #[test]
249    fn rejects_out_of_order_files() {
250        let package = package();
251        let wire = encode(&package).unwrap();
252        let header = FIXED_HEADER_LENGTH + package.id().family().logical_name().len();
253        let first = &package.files()[0];
254        let first_end = header + 10 + first.path().len() + first.bytes().len();
255        let second = &package.files()[1];
256        let second_end = first_end + 10 + second.path().len() + second.bytes().len();
257        let mut reordered = Vec::new();
258        reordered.extend_from_slice(&wire[..header]);
259        reordered.extend_from_slice(&wire[first_end..second_end]);
260        reordered.extend_from_slice(&wire[header..first_end]);
261        reordered.extend_from_slice(&wire[second_end..]);
262        assert_eq!(
263            decode(&reordered).unwrap_err().message(),
264            "file paths are not in strictly increasing order"
265        );
266    }
267}