1#![deny(clippy::cast_possible_truncation)]
3
4use heddle_format::compression::CompressionConfig;
5
6use super::{
7 ObjectType, varint,
8 versioned_header::{HeaderChecksum, VersionedHeader},
9};
10use crate::{
11 object::{ContentHash, StateId},
12 store::{Result, StoreError},
13};
14
15pub const PACK_CHECKSUM_LEN: usize = 32;
16pub const MAX_PACK_OBJECT_OUTPUT_SIZE: usize = 1024 * 1024 * 1024;
17#[cfg(feature = "zstd")]
18pub(super) const PACK_DECOMPRESSION_INITIAL_CAP: usize = 4 * 1024 * 1024;
19
20#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Ord, PartialOrd)]
21pub enum PackObjectId {
22 Hash(ContentHash),
23 StateId(StateId),
24}
25
26impl PackObjectId {
27 pub fn encode_tagged(self, buf: &mut Vec<u8>) {
28 match self {
29 Self::Hash(hash) => {
30 buf.push(0);
31 buf.extend_from_slice(hash.as_bytes());
32 }
33 Self::StateId(state_id) => {
34 buf.push(1);
35 buf.extend_from_slice(state_id.as_bytes());
36 }
37 }
38 }
39
40 pub fn decode_tagged(data: &[u8]) -> Result<(Self, usize)> {
41 let Some(tag) = data.first().copied() else {
42 return Err(StoreError::InvalidObject(
43 "missing pack object id tag".to_string(),
44 ));
45 };
46 match tag {
47 0 => {
48 if data.len() < 33 {
49 return Err(StoreError::InvalidObject(
50 "hash pack object id truncated".to_string(),
51 ));
52 }
53 let hash = ContentHash::from_bytes(data[1..33].try_into().map_err(|_| {
54 StoreError::InvalidObject("invalid hash id length".to_string())
55 })?);
56 Ok((Self::Hash(hash), 33))
57 }
58 1 => {
59 if data.len() < 33 {
60 return Err(StoreError::InvalidObject(
61 "state id pack object id truncated".to_string(),
62 ));
63 }
64 let state_id = StateId::from_bytes(data[1..33].try_into().map_err(|_| {
65 StoreError::InvalidObject("invalid state id length".to_string())
66 })?);
67 Ok((Self::StateId(state_id), 33))
68 }
69 _ => Err(StoreError::InvalidObject(format!(
70 "unknown pack object id tag {tag}"
71 ))),
72 }
73 }
74}
75
76#[derive(Debug, Clone)]
77pub struct PackObjectRecord {
78 pub id: PackObjectId,
79 pub obj_type: ObjectType,
80 pub data: Vec<u8>,
81 pub delta_base: Option<PackObjectId>,
82 pub path_hint: Option<String>,
83}
84
85#[derive(Debug, Clone, Copy)]
86pub struct PackContainerSpec {
87 pub magic: &'static [u8; 4],
88 pub version: u32,
89}
90
91#[derive(Debug, Clone)]
92pub struct PackEntryHeader {
93 pub id: PackObjectId,
94 pub obj_type: ObjectType,
95 pub uncompressed_size: usize,
96 pub compressed_size: usize,
97 pub delta_base: Option<PackObjectId>,
98 pub header_len: usize,
99}
100
101pub fn write_container_header(buf: &mut Vec<u8>, spec: PackContainerSpec, count: u64) {
102 pack_container_header(spec).write_vec(buf, count);
103}
104
105pub fn verify_container(data: &[u8], spec: PackContainerSpec) -> Result<(u64, usize, usize)> {
106 let header = pack_container_header(spec).verify(data)?;
107 Ok((header.count, header.header_len, header.content_end))
108}
109
110pub fn verify_container_layout(
111 data: &[u8],
112 spec: PackContainerSpec,
113) -> Result<(u64, usize, usize)> {
114 let header = pack_container_header(spec).verify_layout(data)?;
115 Ok((header.count, header.header_len, header.content_end))
116}
117
118pub fn append_container_checksum(buf: &mut Vec<u8>) {
119 HeaderChecksum::Blake3Trailer.append(buf);
120}
121
122fn pack_container_header(spec: PackContainerSpec) -> VersionedHeader {
123 VersionedHeader {
124 magic: spec.magic,
125 version: spec.version,
126 checksum: HeaderChecksum::Blake3Trailer,
127 too_short: "Pack too short",
128 invalid_magic: "Invalid pack magic",
129 unsupported_version: "Unsupported pack version",
130 checksum_mismatch: "Pack checksum mismatch",
131 }
132}
133
134pub fn encode_tagged_entry(
135 buf: &mut Vec<u8>,
136 record: &PackObjectRecord,
137 stored_type: ObjectType,
138 compressed: &[u8],
139) -> Result<()> {
140 encode_tagged_entry_parts(
141 buf,
142 record.id,
143 stored_type,
144 record.data.len(),
145 record.delta_base,
146 compressed,
147 )
148}
149
150pub fn encode_tagged_entry_parts(
151 buf: &mut Vec<u8>,
152 id: PackObjectId,
153 stored_type: ObjectType,
154 uncompressed_size: usize,
155 delta_base: Option<PackObjectId>,
156 compressed: &[u8],
157) -> Result<()> {
158 id.encode_tagged(buf);
159 varint::encode_type_and_size(stored_type, uncompressed_size as u64, buf);
160 varint::encode_varint(compressed.len() as u64, buf);
161 if stored_type == ObjectType::Delta {
162 let Some(base) = delta_base else {
163 return Err(StoreError::InvalidObject(
164 "Delta entry missing base id".to_string(),
165 ));
166 };
167 base.encode_tagged(buf);
168 }
169 buf.extend_from_slice(compressed);
170 Ok(())
171}
172
173pub fn decode_tagged_entry_header(data: &[u8]) -> Result<PackEntryHeader> {
174 let (id, id_len) = PackObjectId::decode_tagged(data)?;
175 let (obj_type, uncompressed_size, type_len) = varint::decode_type_and_size(&data[id_len..])
176 .ok_or_else(|| StoreError::InvalidObject("Truncated type+size varint".to_string()))?;
177 let varint_start = id_len + type_len;
178 let (compressed_size, comp_len) = varint::decode_varint(&data[varint_start..])
179 .ok_or_else(|| StoreError::InvalidObject("Truncated compressed_size varint".to_string()))?;
180 let mut header_len = varint_start + comp_len;
181
182 let delta_base = if obj_type == ObjectType::Delta {
183 let (base, base_len) = PackObjectId::decode_tagged(&data[header_len..])?;
184 header_len += base_len;
185 Some(base)
186 } else {
187 None
188 };
189
190 Ok(PackEntryHeader {
191 id,
192 obj_type,
193 uncompressed_size: checked_decoded_size("uncompressed_size", uncompressed_size)?,
194 compressed_size: checked_decoded_size("compressed_size", compressed_size)?,
195 delta_base,
196 header_len,
197 })
198}
199
200pub fn try_decode_tagged_entry_header(data: &[u8]) -> Result<Option<PackEntryHeader>> {
201 let Some(tag) = data.first().copied() else {
202 return Ok(None);
203 };
204
205 let (id, id_len) =
206 match tag {
207 0 => {
208 if data.len() < 33 {
209 return Ok(None);
210 }
211 let hash = ContentHash::from_bytes(data[1..33].try_into().map_err(|_| {
212 StoreError::InvalidObject("invalid hash id length".to_string())
213 })?);
214 (PackObjectId::Hash(hash), 33)
215 }
216 1 => {
217 if data.len() < 33 {
218 return Ok(None);
219 }
220 let state_id = StateId::from_bytes(data[1..33].try_into().map_err(|_| {
221 StoreError::InvalidObject("invalid state id length".to_string())
222 })?);
223 (PackObjectId::StateId(state_id), 33)
224 }
225 _ => {
226 return Err(StoreError::InvalidObject(format!(
227 "unknown pack object id tag {tag}"
228 )));
229 }
230 };
231
232 let Some((obj_type, uncompressed_size, type_len)) =
233 varint::decode_type_and_size(&data[id_len..])
234 else {
235 return Ok(None);
236 };
237 let varint_start = id_len + type_len;
238 let Some((compressed_size, comp_len)) = varint::decode_varint(&data[varint_start..]) else {
239 return Ok(None);
240 };
241 let mut header_len = varint_start + comp_len;
242
243 let delta_base = if obj_type == ObjectType::Delta {
244 let Some(base_tag) = data.get(header_len).copied() else {
245 return Ok(None);
246 };
247 let (base, base_len) = match base_tag {
248 0 => {
249 let end = header_len + 33;
250 if data.len() < end {
251 return Ok(None);
252 }
253 let hash = ContentHash::from_bytes(data[header_len + 1..end].try_into().map_err(
254 |_| StoreError::InvalidObject("invalid hash id length".to_string()),
255 )?);
256 (PackObjectId::Hash(hash), 33)
257 }
258 1 => {
259 let end = header_len + 33;
260 if data.len() < end {
261 return Ok(None);
262 }
263 let state_id =
264 StateId::from_bytes(data[header_len + 1..end].try_into().map_err(|_| {
265 StoreError::InvalidObject("invalid state id length".to_string())
266 })?);
267 (PackObjectId::StateId(state_id), 33)
268 }
269 _ => {
270 return Err(StoreError::InvalidObject(format!(
271 "unknown pack object id tag {base_tag}"
272 )));
273 }
274 };
275 header_len += base_len;
276 Some(base)
277 } else {
278 None
279 };
280
281 Ok(Some(PackEntryHeader {
282 id,
283 obj_type,
284 uncompressed_size: checked_decoded_size("uncompressed_size", uncompressed_size)?,
285 compressed_size: checked_decoded_size("compressed_size", compressed_size)?,
286 delta_base,
287 header_len,
288 }))
289}
290
291fn checked_decoded_size(field: &str, size: u64) -> Result<usize> {
292 let size = usize::try_from(size).map_err(|_| {
293 StoreError::InvalidObject(format!("Decoded {field} exceeds platform limits"))
294 })?;
295 if field == "uncompressed_size" {
296 reject_pack_object_output_over_limit(size, MAX_PACK_OBJECT_OUTPUT_SIZE)?;
297 }
298 Ok(size)
299}
300
301pub fn compress_pack_payload(data: &[u8], config: &CompressionConfig) -> Result<Vec<u8>> {
302 if !config.enabled || data.len() < config.min_size {
303 return Ok(data.to_vec());
304 }
305 #[cfg(feature = "zstd")]
306 {
307 match zstd::encode_all(data, config.level) {
308 Ok(compressed) if compressed.len() < data.len() => Ok(compressed),
309 _ => Ok(data.to_vec()),
310 }
311 }
312 #[cfg(not(feature = "zstd"))]
313 {
314 let _ = config;
315 Ok(data.to_vec())
316 }
317}
318
319pub fn decompress_pack_payload(data: &[u8], expected_size: usize) -> Result<Vec<u8>> {
320 #[cfg(feature = "zstd")]
321 {
322 decompress_pack_payload_with_limit(data, expected_size, MAX_PACK_OBJECT_OUTPUT_SIZE)
323 }
324 #[cfg(not(feature = "zstd"))]
325 {
326 reject_pack_object_output_over_limit(expected_size, MAX_PACK_OBJECT_OUTPUT_SIZE)?;
327 reject_pack_object_output_over_limit(data.len(), MAX_PACK_OBJECT_OUTPUT_SIZE)?;
328 Ok(data.to_vec())
329 }
330}
331
332#[cfg(feature = "zstd")]
333pub(super) fn decompress_pack_payload_with_limit(
334 data: &[u8],
335 expected_size: usize,
336 max_output_size: usize,
337) -> Result<Vec<u8>> {
338 use std::io::Read;
339
340 reject_pack_object_output_over_limit(expected_size, max_output_size)?;
345
346 let mut decoder = zstd::stream::read::Decoder::new(data)
347 .map_err(|e| StoreError::InvalidObject(format!("zstd decode init failed: {e}")))?;
348 let capacity = initial_decompression_capacity(data.len(), expected_size, max_output_size);
349 let mut buf = Vec::with_capacity(capacity);
350 let mut chunk = [0u8; 8192];
351
352 loop {
353 let bytes_read = decoder
354 .read(&mut chunk)
355 .map_err(|e| StoreError::InvalidObject(format!("zstd decompression failed: {e}")))?;
356 if bytes_read == 0 {
357 break;
358 }
359
360 let next_len = buf.len().checked_add(bytes_read).ok_or_else(|| {
361 StoreError::InvalidObject("Pack object output size overflows".to_string())
362 })?;
363 reject_pack_object_output_over_limit(next_len, max_output_size)?;
364 buf.extend_from_slice(&chunk[..bytes_read]);
365 }
366
367 Ok(buf)
368}
369
370#[cfg(feature = "zstd")]
371fn initial_decompression_capacity(
372 compressed_len: usize,
373 expected_size: usize,
374 max_output_size: usize,
375) -> usize {
376 let hint = if expected_size > 0 {
377 expected_size
378 } else {
379 compressed_len.saturating_mul(2)
380 };
381 hint.min(PACK_DECOMPRESSION_INITIAL_CAP)
382 .min(max_output_size)
383}
384
385fn reject_pack_object_output_over_limit(size: usize, max: usize) -> Result<()> {
386 if size > max {
387 return Err(StoreError::InvalidObject(format!(
388 "Pack object output size {size} exceeds max {max}"
389 )));
390 }
391 Ok(())
392}
393
394pub fn has_zstd_magic(data: &[u8]) -> bool {
395 data.len() >= 4 && data[..4] == [0x28, 0xB5, 0x2F, 0xFD]
396}
397
398#[cfg(test)]
399mod tests {
400 use super::*;
401
402 #[test]
403 fn tagged_pack_object_ids_round_trip() {
404 let ids = [
405 PackObjectId::Hash(ContentHash::compute(b"hash-object")),
406 PackObjectId::StateId(StateId::from_bytes([7; 32])),
407 ];
408
409 for id in ids {
410 let mut encoded = Vec::new();
411 id.encode_tagged(&mut encoded);
412 let (decoded, consumed) = PackObjectId::decode_tagged(&encoded).unwrap();
413 assert_eq!(decoded, id);
414 assert_eq!(consumed, encoded.len());
415 }
416 }
417
418 #[test]
419 fn tagged_entry_header_round_trips_mixed_identity() {
420 let record = PackObjectRecord {
421 id: PackObjectId::StateId(StateId::from_bytes([8; 32])),
422 obj_type: ObjectType::State,
423 data: vec![1, 2, 3, 4, 5],
424 delta_base: None,
425 path_hint: None,
426 };
427
428 let mut encoded = Vec::new();
429 encode_tagged_entry(&mut encoded, &record, record.obj_type, &record.data).unwrap();
430 let decoded = decode_tagged_entry_header(&encoded).unwrap();
431
432 assert_eq!(decoded.id, record.id);
433 assert_eq!(decoded.obj_type, ObjectType::State);
434 assert_eq!(decoded.uncompressed_size, 5);
435 assert_eq!(decoded.compressed_size, 5);
436 assert_eq!(decoded.delta_base, None);
437 }
438
439 #[test]
440 fn tagged_entry_header_rejects_size_that_truncates_on_32_bit() {
441 let mut encoded = Vec::new();
442 PackObjectId::Hash(ContentHash::compute(b"oversized-pack-object"))
443 .encode_tagged(&mut encoded);
444 varint::encode_type_and_size(ObjectType::Blob, u64::from(u32::MAX) + 1, &mut encoded);
445 varint::encode_varint(1, &mut encoded);
446 encoded.push(0);
447
448 let result = decode_tagged_entry_header(&encoded);
449
450 let error = result.expect_err("absurd 32-bit-overflow size must be rejected");
451 assert!(
452 matches!(&error, StoreError::InvalidObject(message) if message.contains("platform limits") || message.contains("Pack object output size")),
453 "expected size-limit InvalidObject, got: {error:?}",
454 );
455 }
456
457 #[test]
458 fn tagged_entry_header_rejects_u64_max_size_when_platform_cannot_represent_it() {
459 let mut encoded = Vec::new();
460 PackObjectId::Hash(ContentHash::compute(b"u64-max-pack-object"))
461 .encode_tagged(&mut encoded);
462 varint::encode_type_and_size(ObjectType::Blob, u64::MAX, &mut encoded);
463 varint::encode_varint(1, &mut encoded);
464 encoded.push(0);
465
466 let result = decode_tagged_entry_header(&encoded);
467
468 let error = result.expect_err("absurd u64::MAX size must be rejected");
469 assert!(
470 matches!(&error, StoreError::InvalidObject(message) if message.contains("platform limits") || message.contains("Pack object output size")),
471 "expected size-limit InvalidObject, got: {error:?}",
472 );
473 }
474}