Skip to main content

heddle_object_model/compact/
tree.rs

1// SPDX-License-Identifier: Apache-2.0
2
3use sley_core::{ObjectFormat as GitObjectFormat, ObjectId as GitObjectId};
4
5use super::{
6    Result, invalid,
7    io::{Reader, Writer, varint_len},
8    limits::{MAX_COMPACT_COUNT, MIN_TREE_ENTRY_BYTES, MIN_TREE_ITEM_BYTES},
9};
10use crate::object::{
11    ContentHash, EntryType, FileMode, SpoolId, StateId, Tree, TreeEntry, TreeEntryTarget,
12    TreeScheme,
13};
14
15const TREE_MAGIC: &[u8; 4] = b"HCT1";
16
17/// Whether `bytes` begin with the compact-tree frame discriminator.
18pub fn is_tree_frame(bytes: &[u8]) -> bool {
19    bytes.starts_with(TREE_MAGIC)
20}
21
22/// Exact raw-frame size for `tree`, including its per-tree entry count.
23pub fn encoded_tree_size(tree: &Tree) -> usize {
24    varint_len(tree.len())
25        + tree.len() * 2
26        + tree
27            .entries()
28            .iter()
29            .map(|entry| varint_len(entry.name().len()) + entry.name().len() + target_len(entry))
30            .sum::<usize>()
31}
32
33/// Encode name-sorted trees as columnar mode/kind/name/target payloads.
34pub fn encode_tree_frame(trees: &[Tree]) -> Result<Vec<u8>> {
35    if trees.len() > MAX_COMPACT_COUNT {
36        return Err(invalid(format!(
37            "tree frame count {} exceeds maximum {MAX_COMPACT_COUNT}",
38            trees.len()
39        )));
40    }
41    if let Some(count) = trees
42        .iter()
43        .map(Tree::len)
44        .find(|len| *len > MAX_COMPACT_COUNT)
45    {
46        return Err(invalid(format!(
47            "tree entry count {count} exceeds maximum {MAX_COMPACT_COUNT}"
48        )));
49    }
50    if trees
51        .iter()
52        .any(|tree| tree.scheme() == TreeScheme::V4Salted)
53    {
54        // HCT1 columns carry no salt; a V4 salted tree must never be written
55        // through this salt-less frame (it would silently re-hash as V3).
56        return Err(invalid(
57            "cannot encode a v4 salted tree in an HCT1 compact frame".to_string(),
58        ));
59    }
60    let mut output = Writer::new(TREE_MAGIC);
61    output.put_u64(trees.len() as u64);
62    for tree in trees {
63        tree.validate()?;
64        output.put_u64(tree.len() as u64);
65        for entry in tree.entries() {
66            output.put_u8(entry.mode().to_byte());
67        }
68        for entry in tree.entries() {
69            output.put_u8(entry.entry_type().to_byte());
70        }
71        for entry in tree.entries() {
72            output.put_bytes(entry.name().as_bytes());
73        }
74        for entry in tree.entries() {
75            encode_target(&mut output, entry);
76        }
77    }
78    Ok(output.finish())
79}
80
81/// Decode and whole-frame-verify every tree in a compact frame.
82pub fn decode_tree_frame(bytes: &[u8]) -> Result<Vec<Tree>> {
83    let mut input = Reader::verified(bytes, TREE_MAGIC)?;
84    let tree_count = input.get_count("tree frame", MIN_TREE_ITEM_BYTES)?;
85    let mut trees = Vec::with_capacity(tree_count);
86    for _ in 0..tree_count {
87        trees.push(decode_tree(&mut input)?);
88    }
89    input.finish()?;
90    Ok(trees)
91}
92
93/// Reconstruct one tree and verify its BLAKE3 typed hash.
94///
95/// The whole-frame checksum is verified first. Each tree is rebuilt from its
96/// SoA columns until `expected` matches the reconstructed typed hash.
97pub fn extract_tree(bytes: &[u8], expected: ContentHash) -> Result<Tree> {
98    let mut input = Reader::verified(bytes, TREE_MAGIC)?;
99    let tree_count = input.get_count("tree frame", MIN_TREE_ITEM_BYTES)?;
100    for _ in 0..tree_count {
101        let tree = decode_tree(&mut input)?;
102        if tree.hash() == expected {
103            return Ok(tree);
104        }
105    }
106    Err(super::CompactError::Missing)
107}
108
109fn decode_tree(input: &mut Reader<'_>) -> Result<Tree> {
110    let count = input.get_count("tree entry", MIN_TREE_ENTRY_BYTES)?;
111    let modes = (0..count)
112        .map(|_| decode_mode(input.get_u8()?))
113        .collect::<Result<Vec<_>>>()?;
114    let kinds = (0..count)
115        .map(|_| decode_kind(input.get_u8()?))
116        .collect::<Result<Vec<_>>>()?;
117    let names = (0..count)
118        .map(|_| {
119            String::from_utf8(input.get_bytes()?).map_err(|_| invalid("tree name is not UTF-8"))
120        })
121        .collect::<Result<Vec<_>>>()?;
122    let mut entries = Vec::with_capacity(count);
123    for index in 0..count {
124        entries.push(decode_entry(
125            input,
126            names[index].clone(),
127            kinds[index],
128            modes[index],
129        )?);
130    }
131    let tree = Tree::from_entries(entries);
132    tree.validate()?;
133    Ok(tree)
134}
135
136fn encode_target(output: &mut Writer, entry: &TreeEntry) {
137    match entry.target() {
138        TreeEntryTarget::Blob { hash, .. }
139        | TreeEntryTarget::Tree { hash }
140        | TreeEntryTarget::Symlink { hash } => {
141            output.put_fixed(hash.as_bytes());
142        }
143        TreeEntryTarget::Gitlink { target } => {
144            output.put_u8(git_format_tag(target.format()));
145            output.put_fixed(target.as_bytes());
146        }
147        TreeEntryTarget::Spoollink { spool_id, state_id } => {
148            output.put_bytes(spool_id.as_str().as_bytes());
149            output.put_fixed(state_id.as_bytes());
150        }
151    }
152}
153
154fn decode_entry(
155    input: &mut Reader<'_>,
156    name: String,
157    kind: EntryType,
158    mode: FileMode,
159) -> Result<TreeEntry> {
160    let entry = match kind {
161        EntryType::Blob => TreeEntry::file(
162            name,
163            ContentHash::from_bytes(input.get_fixed()?),
164            mode == FileMode::Executable,
165        )?,
166        EntryType::Tree => TreeEntry::directory(name, ContentHash::from_bytes(input.get_fixed()?))?,
167        EntryType::Symlink => {
168            TreeEntry::symlink(name, ContentHash::from_bytes(input.get_fixed()?))?
169        }
170        EntryType::Gitlink => {
171            let format = decode_git_format(input.get_u8()?)?;
172            let oid_len = match format {
173                GitObjectFormat::Sha1 => 20,
174                GitObjectFormat::Sha256 => 32,
175            };
176            let target = GitObjectId::from_raw(format, input.take(oid_len)?)
177                .map_err(|error| super::CompactError::GitObjectId(error.to_string()))?;
178            TreeEntry::gitlink(name, target)?
179        }
180        EntryType::Spoollink => {
181            let spool = String::from_utf8(input.get_bytes()?)
182                .map_err(|_| invalid("spool id is not UTF-8"))?;
183            TreeEntry::spoollink(
184                name,
185                SpoolId::parse(spool)
186                    .map_err(|error| super::CompactError::SpoolId(error.to_string()))?,
187                StateId::from_bytes(input.get_fixed()?),
188            )?
189        }
190    };
191    if entry.mode() != mode {
192        return Err(invalid(format!(
193            "tree kind/mode mismatch for {}: {kind:?}/{mode:?}",
194            entry.name()
195        )));
196    }
197    Ok(entry)
198}
199
200fn target_len(entry: &TreeEntry) -> usize {
201    match entry.target() {
202        TreeEntryTarget::Blob { .. }
203        | TreeEntryTarget::Tree { .. }
204        | TreeEntryTarget::Symlink { .. } => 32,
205        TreeEntryTarget::Gitlink { target } => 1 + target.as_bytes().len(),
206        TreeEntryTarget::Spoollink { spool_id, .. } => {
207            varint_len(spool_id.as_str().len()) + spool_id.as_str().len() + 32
208        }
209    }
210}
211
212fn decode_mode(value: u8) -> Result<FileMode> {
213    FileMode::from_byte(value).ok_or_else(|| invalid(format!("invalid tree mode {value}")))
214}
215
216fn decode_kind(value: u8) -> Result<EntryType> {
217    EntryType::from_byte(value).ok_or_else(|| invalid(format!("invalid tree kind {value}")))
218}
219
220fn git_format_tag(format: GitObjectFormat) -> u8 {
221    match format {
222        GitObjectFormat::Sha1 => 1,
223        GitObjectFormat::Sha256 => 2,
224    }
225}
226
227fn decode_git_format(value: u8) -> Result<GitObjectFormat> {
228    match value {
229        1 => Ok(GitObjectFormat::Sha1),
230        2 => Ok(GitObjectFormat::Sha256),
231        _ => Err(invalid(format!("invalid git object format {value}"))),
232    }
233}