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