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}
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    // Pack objects may be raw blobs, so this bound must be materially
341    // larger than the delta-output limit. It is also intentionally
342    // above the protocol default and loose-compression cap, while
343    // still bounding one untrusted pack record to a finite allocation.
344    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}