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