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_rust_worktree::UnpublishedId;
5use kcode_k1_transaction_id::TxId;
6use semver::Version;
7use std::fmt::{Display, Formatter};
8
9pub const WIRE_VERSION: u8 = 1;
10const CREATE: u8 = 0;
11const FORK: u8 = 1;
12const OVERWRITE: u8 = 2;
13const PUBLISH: u8 = 3;
14const PACKAGE_HEADER_LENGTH: usize = 45;
15
16#[derive(Clone, Debug, Eq, PartialEq)]
17pub enum RustSourceTransaction {
18    Create(SourcePackage),
19    Fork(SourcePackage),
20    Overwrite {
21        id: UnpublishedId,
22        expected_revision: TxId,
23        source: SourcePackage,
24    },
25    Publish(SourcePackage),
26}
27
28impl RustSourceTransaction {
29    pub fn source(&self) -> &SourcePackage {
30        match self {
31            Self::Create(source)
32            | Self::Fork(source)
33            | Self::Overwrite { source, .. }
34            | Self::Publish(source) => source,
35        }
36    }
37
38    pub fn family(&self) -> &LibraryFamily {
39        self.source().id().family()
40    }
41}
42
43#[derive(Clone, Debug, Eq, PartialEq)]
44pub struct TransactionError(ErrorMessage);
45
46#[derive(Clone, Debug, Eq, PartialEq)]
47enum ErrorMessage {
48    Static(&'static str),
49    Package(PackageError),
50}
51
52impl TransactionError {
53    pub fn message(&self) -> &str {
54        match &self.0 {
55            ErrorMessage::Static(message) => message,
56            ErrorMessage::Package(error) => error.message(),
57        }
58    }
59}
60
61impl Display for TransactionError {
62    fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
63        formatter.write_str(self.message())
64    }
65}
66
67impl std::error::Error for TransactionError {}
68
69impl From<PackageError> for TransactionError {
70    fn from(error: PackageError) -> Self {
71        Self(ErrorMessage::Package(error))
72    }
73}
74
75pub fn encode(event: &RustSourceTransaction) -> Result<Vec<u8>, TransactionError> {
76    let metadata = matches!(event, RustSourceTransaction::Overwrite { .. })
77        .then_some(24)
78        .unwrap_or(0);
79    let total = encoded_package_length(event.source())?
80        .checked_add(2 + metadata)
81        .ok_or_else(|| error("encoded length overflows"))?;
82    let mut output = Vec::new();
83    reserve(&mut output, total)?;
84    output.push(WIRE_VERSION);
85    match event {
86        RustSourceTransaction::Create(source) => {
87            output.push(CREATE);
88            encode_package(source, &mut output)?;
89        }
90        RustSourceTransaction::Fork(source) => {
91            output.push(FORK);
92            encode_package(source, &mut output)?;
93        }
94        RustSourceTransaction::Overwrite {
95            id,
96            expected_revision,
97            source,
98        } => {
99            output.push(OVERWRITE);
100            output.extend_from_slice(id.transaction().as_bytes());
101            output.extend_from_slice(expected_revision.as_bytes());
102            encode_package(source, &mut output)?;
103        }
104        RustSourceTransaction::Publish(source) => {
105            output.push(PUBLISH);
106            encode_package(source, &mut output)?;
107        }
108    }
109    Ok(output)
110}
111
112pub fn decode(bytes: &[u8]) -> Result<RustSourceTransaction, TransactionError> {
113    let mut reader = Reader { bytes, offset: 0 };
114    if reader.byte()? != WIRE_VERSION {
115        return Err(error("unknown wire version"));
116    }
117    let kind = reader.byte()?;
118    let metadata = if kind == OVERWRITE {
119        Some((
120            UnpublishedId::new(TxId::from_bytes(reader.array()?)),
121            TxId::from_bytes(reader.array()?),
122        ))
123    } else {
124        None
125    };
126    let source = decode_package(&mut reader)?;
127    if reader.offset != bytes.len() {
128        return Err(error("trailing bytes"));
129    }
130    match (kind, metadata) {
131        (CREATE, None) => Ok(RustSourceTransaction::Create(source)),
132        (FORK, None) => Ok(RustSourceTransaction::Fork(source)),
133        (OVERWRITE, Some((id, expected_revision))) => Ok(RustSourceTransaction::Overwrite {
134            id,
135            expected_revision,
136            source,
137        }),
138        (PUBLISH, None) => Ok(RustSourceTransaction::Publish(source)),
139        _ => Err(error("unknown transaction kind")),
140    }
141}
142
143fn encoded_package_length(package: &SourcePackage) -> Result<usize, TransactionError> {
144    let mut total = PACKAGE_HEADER_LENGTH
145        .checked_add(package.id().family().logical_name().len())
146        .ok_or_else(|| error("encoded length overflows"))?;
147    for file in package.files() {
148        u16::try_from(file.path().len())
149            .map_err(|_| error("path is too long for the wire format"))?;
150        u64::try_from(file.bytes().len())
151            .map_err(|_| error("content length overflows the wire format"))?;
152        total = total
153            .checked_add(10)
154            .and_then(|value| value.checked_add(file.path().len()))
155            .and_then(|value| value.checked_add(file.bytes().len()))
156            .ok_or_else(|| error("encoded length overflows"))?;
157    }
158    Ok(total)
159}
160
161fn encode_package(package: &SourcePackage, output: &mut Vec<u8>) -> Result<(), TransactionError> {
162    let family = package.id().family();
163    let name = family.logical_name().as_bytes();
164    let name_length = u8::try_from(name.len()).map_err(|_| error("logical name is too long"))?;
165    let file_count = u64::try_from(package.files().len())
166        .map_err(|_| error("file count overflows the wire format"))?;
167    output.extend_from_slice(family.authority().transaction_id().as_bytes());
168    output.push(name_length);
169    output.extend_from_slice(name);
170    for value in [
171        package.id().version().major,
172        package.id().version().minor,
173        package.id().version().patch,
174    ] {
175        output.extend_from_slice(&value.to_be_bytes());
176    }
177    output.extend_from_slice(&file_count.to_be_bytes());
178    for file in package.files() {
179        let path_length = u16::try_from(file.path().len())
180            .map_err(|_| error("path is too long for the wire format"))?;
181        let content_length = u64::try_from(file.bytes().len())
182            .map_err(|_| error("content length overflows the wire format"))?;
183        output.extend_from_slice(&path_length.to_be_bytes());
184        output.extend_from_slice(file.path().as_bytes());
185        output.extend_from_slice(&content_length.to_be_bytes());
186        output.extend_from_slice(file.bytes());
187    }
188    Ok(())
189}
190
191fn decode_package(reader: &mut Reader<'_>) -> Result<SourcePackage, TransactionError> {
192    let authority = AuthorityId::new(TxId::from_bytes(reader.array()?));
193    let name_length = usize::from(reader.byte()?);
194    let name = reader.string(name_length, "logical name is not UTF-8")?;
195    let version = Version::new(reader.u64()?, reader.u64()?, reader.u64()?);
196    let file_count =
197        usize::try_from(reader.u64()?).map_err(|_| error("file count overflows this platform"))?;
198    let mut files = Vec::new();
199    reserve(&mut files, file_count)?;
200    for _ in 0..file_count {
201        let path_length = usize::from(reader.u16()?);
202        let path = reader.string(path_length, "file path is not UTF-8")?;
203        if files
204            .last()
205            .is_some_and(|previous: &SourceFile| previous.path() >= path.as_str())
206        {
207            return Err(error("file paths are not in strictly increasing order"));
208        }
209        let content_length = usize::try_from(reader.u64()?)
210            .map_err(|_| error("content length overflows this platform"))?;
211        files.push(SourceFile::new(path, reader.bytes(content_length)?));
212    }
213    let family = LibraryFamily::new(authority, name)?;
214    let identity = LibraryId::new(family, version)?;
215    SourcePackage::new(identity, files).map_err(Into::into)
216}
217
218fn error(message: &'static str) -> TransactionError {
219    TransactionError(ErrorMessage::Static(message))
220}
221
222fn reserve<T>(values: &mut Vec<T>, additional: usize) -> Result<(), TransactionError> {
223    values
224        .try_reserve_exact(additional)
225        .map_err(|_| error("allocation failed"))
226}
227
228struct Reader<'a> {
229    bytes: &'a [u8],
230    offset: usize,
231}
232
233impl<'a> Reader<'a> {
234    fn take(&mut self, length: usize) -> Result<&'a [u8], TransactionError> {
235        let end = self
236            .offset
237            .checked_add(length)
238            .ok_or_else(|| error("decoded offset overflows"))?;
239        let value = self
240            .bytes
241            .get(self.offset..end)
242            .ok_or_else(|| error("truncated transaction"))?;
243        self.offset = end;
244        Ok(value)
245    }
246
247    fn byte(&mut self) -> Result<u8, TransactionError> {
248        Ok(self.take(1)?[0])
249    }
250
251    fn array<const N: usize>(&mut self) -> Result<[u8; N], TransactionError> {
252        let mut value = [0; N];
253        value.copy_from_slice(self.take(N)?);
254        Ok(value)
255    }
256
257    fn u16(&mut self) -> Result<u16, TransactionError> {
258        Ok(u16::from_be_bytes(self.array()?))
259    }
260
261    fn u64(&mut self) -> Result<u64, TransactionError> {
262        Ok(u64::from_be_bytes(self.array()?))
263    }
264
265    fn bytes(&mut self, length: usize) -> Result<Vec<u8>, TransactionError> {
266        let source = self.take(length)?;
267        let mut value = Vec::new();
268        reserve(&mut value, length)?;
269        value.extend_from_slice(source);
270        Ok(value)
271    }
272
273    fn string(&mut self, length: usize, invalid: &'static str) -> Result<String, TransactionError> {
274        String::from_utf8(self.bytes(length)?).map_err(|_| error(invalid))
275    }
276}
277
278#[cfg(test)]
279mod tests {
280    use super::*;
281
282    fn package() -> SourcePackage {
283        let family =
284            LibraryFamily::new(AuthorityId::new(TxId::from_bytes([1; 12])), "alpha").unwrap();
285        let identity = LibraryId::new(family, Version::new(1, 2, 3)).unwrap();
286        let manifest = r#"[package]
287name = "k1-010101010101010101010101-alpha"
288version = "1.2.3"
289edition = "2024"
290autobins = false
291autoexamples = false
292autotests = false
293autobenches = false
294
295[lib]
296name = "alpha"
297path = "src/lib.rs"
298
299[workspace]
300resolver = "3"
301"#;
302        SourcePackage::new(
303            identity,
304            vec![
305                SourceFile::new("Cargo.toml", manifest.as_bytes().to_vec()),
306                SourceFile::new("Documentation.md", b"docs".to_vec()),
307                SourceFile::new("src/lib.rs", vec![0, 159, 255]),
308            ],
309        )
310        .unwrap()
311    }
312
313    fn events() -> Vec<RustSourceTransaction> {
314        let source = package();
315        vec![
316            RustSourceTransaction::Create(source.clone()),
317            RustSourceTransaction::Fork(source.clone()),
318            RustSourceTransaction::Overwrite {
319                id: UnpublishedId::new(TxId::from_bytes([2; 12])),
320                expected_revision: TxId::from_bytes([3; 12]),
321                source: source.clone(),
322            },
323            RustSourceTransaction::Publish(source),
324        ]
325    }
326
327    #[test]
328    fn every_event_round_trips_deterministically() {
329        for event in events() {
330            let wire = encode(&event).unwrap();
331            assert_eq!(wire, encode(&event).unwrap());
332            assert_eq!(decode(&wire).unwrap(), event);
333        }
334    }
335
336    #[test]
337    fn rejects_unknown_truncated_and_trailing_data() {
338        let wire = encode(&events().remove(0)).unwrap();
339        let mut unknown_version = wire.clone();
340        unknown_version[0] = 2;
341        assert_eq!(
342            decode(&unknown_version).unwrap_err().message(),
343            "unknown wire version"
344        );
345        let mut unknown_kind = wire.clone();
346        unknown_kind[1] = 9;
347        assert_eq!(
348            decode(&unknown_kind).unwrap_err().message(),
349            "unknown transaction kind"
350        );
351        for end in 0..wire.len() {
352            assert!(decode(&wire[..end]).is_err(), "accepted prefix {end}");
353        }
354        let mut trailing = wire;
355        trailing.push(0);
356        assert_eq!(decode(&trailing).unwrap_err().message(), "trailing bytes");
357    }
358}