Skip to main content

heddle_object_model/compact/
tree.rs

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