kcode_k1_rust_transaction/
lib.rs1use 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}