Skip to main content

heddle_pack/store/pack/
shared.rs

1// SPDX-License-Identifier: Apache-2.0
2#![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    // Pack objects may be raw blobs, so this bound must be materially
426    // larger than the delta-output limit. It is also intentionally
427    // above the protocol default and loose-compression cap, while
428    // still bounding one untrusted pack record to a finite allocation.
429    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}