This repository has no description
1use std::collections::HashMap;
2use std::io::{self, Write};
3
4use flate2::{Decompress, FlushDecompress, Status};
5use gix_pack::data::input;
6use gix_pack::data::{entry::Header, header};
7
8use knot_types::ObjectCount;
9
10use crate::error::{PackError, PackLimit};
11use crate::ids::{DeltaDepth, MaxObjectBytes, MaxTotalBytes, PackOffset};
12
13const MAX_OBJECTS: ObjectCount = ObjectCount::new(16_000_000);
14const MAX_OBJECT_BYTES: MaxObjectBytes = MaxObjectBytes::new(512 * 1024 * 1024);
15const MAX_TOTAL_BYTES: MaxTotalBytes = MaxTotalBytes::new(16 * 1024 * 1024 * 1024);
16const MAX_DELTA_DEPTH: DeltaDepth = DeltaDepth::new(50);
17
18#[derive(Debug, Clone, Copy, PartialEq, Eq)]
19pub struct PackLimits {
20 pub max_objects: ObjectCount,
21 pub max_object_bytes: MaxObjectBytes,
22 pub max_total_bytes: MaxTotalBytes,
23 pub max_delta_depth: DeltaDepth,
24}
25
26impl Default for PackLimits {
27 fn default() -> Self {
28 Self {
29 max_objects: MAX_OBJECTS,
30 max_object_bytes: MAX_OBJECT_BYTES,
31 max_total_bytes: MAX_TOTAL_BYTES,
32 max_delta_depth: MAX_DELTA_DEPTH,
33 }
34 }
35}
36
37pub(crate) fn malformed(message: &str) -> PackError {
38 PackError::Pack(message.to_string())
39}
40
41pub(crate) fn pack_object_count(pack: &[u8]) -> Result<ObjectCount, PackError> {
42 let head: [u8; 12] = pack
43 .get(..12)
44 .and_then(|slice| slice.try_into().ok())
45 .ok_or_else(|| malformed("packfile header is truncated"))?;
46 header::decode(&head)
47 .map(|(_version, num_objects)| ObjectCount::from(num_objects))
48 .map_err(|error| PackError::Pack(error.to_string()))
49}
50
51pub fn meter(pack: &[u8], limits: &PackLimits, kind: gix::hash::Kind) -> Result<(), PackError> {
52 if pack.len() < 12 + kind.len_in_bytes() {
53 return Err(malformed("packfile is truncated"));
54 }
55 meter_entries(
56 io::Cursor::new(pack),
57 pack_object_count(pack)?,
58 limits,
59 kind,
60 )
61 .map(|_| ())
62}
63
64pub(crate) fn meter_file(
65 pack: &gix_pack::data::File,
66 limits: &PackLimits,
67 kind: gix::hash::Kind,
68) -> Result<bool, PackError> {
69 let reader = io::BufReader::new(std::fs::File::open(pack.path())?);
70 meter_entries(reader, ObjectCount::from(pack.num_objects()), limits, kind)
71}
72
73fn meter_entries<R: io::BufRead>(
74 reader: R,
75 num_objects: ObjectCount,
76 limits: &PackLimits,
77 kind: gix::hash::Kind,
78) -> Result<bool, PackError> {
79 if num_objects > limits.max_objects {
80 return Err(PackError::LimitExceeded(PackLimit::Objects));
81 }
82 let mut entries = input::BytesToEntriesIter::new_from_header(
83 reader,
84 input::Mode::Verify,
85 input::EntryDataMode::Keep,
86 kind,
87 )
88 .map_err(|error| PackError::Pack(error.to_string()))?;
89
90 let mut total = 0u64;
91 let mut thin = false;
92 let mut base_of: HashMap<PackOffset, PackOffset> = HashMap::new();
93 entries.try_for_each(|entry| -> Result<(), PackError> {
94 let entry = entry.map_err(|error| PackError::Pack(error.to_string()))?;
95 if limits.max_object_bytes.exceeded_by(entry.decompressed_size) {
96 return Err(PackError::LimitExceeded(PackLimit::ObjectBytes));
97 }
98 total = total
99 .checked_add(entry.decompressed_size)
100 .ok_or_else(|| malformed("decompressed size overflow"))?;
101 if limits.max_total_bytes.exceeded_by(total) {
102 return Err(PackError::LimitExceeded(PackLimit::TotalBytes));
103 }
104 match entry.header {
105 Header::OfsDelta { base_distance } => {
106 let pack_offset = PackOffset::new(entry.pack_offset);
107 let base = pack_offset
108 .checked_sub_distance(base_distance)
109 .ok_or_else(|| malformed("ofs-delta base out of range"))?;
110 base_of.insert(pack_offset, base);
111 check_delta_result(&entry, limits.max_object_bytes)?;
112 }
113 Header::RefDelta { .. } => {
114 thin = true;
115 check_delta_result(&entry, limits.max_object_bytes)?;
116 }
117 _ => {}
118 }
119 Ok(())
120 })?;
121
122 check_depth(&base_of, limits.max_delta_depth)?;
123 Ok(thin)
124}
125
126fn check_delta_result(
127 entry: &input::Entry,
128 max_object_bytes: MaxObjectBytes,
129) -> Result<(), PackError> {
130 let compressed = entry
131 .compressed
132 .as_deref()
133 .ok_or_else(|| malformed("delta entry missing compressed data"))?;
134 let mut peek = HeaderPeek::new();
135 inflate_into(compressed, entry.decompressed_size, &mut peek)?;
136 if max_object_bytes.exceeded_by(delta_result_size(peek.filled())?) {
137 return Err(PackError::LimitExceeded(PackLimit::ObjectBytes));
138 }
139 Ok(())
140}
141
142pub(crate) fn inflate_into(
143 input: &[u8],
144 expected: u64,
145 out: &mut dyn Write,
146) -> Result<u64, PackError> {
147 let mut decompress = Decompress::new(true);
148 let mut scratch = [0u8; 8192];
149 let mut produced = 0u64;
150 loop {
151 let consumed = decompress.total_in() as usize;
152 let out_before = decompress.total_out();
153 let status = decompress
154 .decompress(
155 input.get(consumed..).unwrap_or_default(),
156 &mut scratch,
157 FlushDecompress::None,
158 )
159 .map_err(|error| PackError::Pack(format!("inflate: {error}")))?;
160 let written = (decompress.total_out() - out_before) as usize;
161 produced += written as u64;
162 if produced > expected {
163 return Err(malformed("object inflates beyond its declared size"));
164 }
165 out.write_all(&scratch[..written])
166 .map_err(|error| PackError::Pack(format!("inflate sink: {error}")))?;
167 match status {
168 Status::StreamEnd => break,
169 Status::Ok | Status::BufError => {
170 if decompress.total_in() as usize == consumed
171 && decompress.total_out() == out_before
172 {
173 return Err(malformed("inflate stalled or pack truncated"));
174 }
175 }
176 }
177 }
178 if produced != expected {
179 return Err(malformed("object decompressed size mismatch"));
180 }
181 Ok(decompress.total_in())
182}
183
184struct HeaderPeek {
185 bytes: [u8; 32],
186 len: usize,
187}
188
189impl HeaderPeek {
190 fn new() -> Self {
191 Self {
192 bytes: [0u8; 32],
193 len: 0,
194 }
195 }
196
197 fn filled(&self) -> &[u8] {
198 &self.bytes[..self.len]
199 }
200}
201
202impl Write for HeaderPeek {
203 fn write(&mut self, data: &[u8]) -> io::Result<usize> {
204 let take = (self.bytes.len() - self.len).min(data.len());
205 self.bytes[self.len..self.len + take].copy_from_slice(&data[..take]);
206 self.len += take;
207 Ok(data.len())
208 }
209
210 fn flush(&mut self) -> io::Result<()> {
211 Ok(())
212 }
213}
214
215fn read_delta_varint(
216 data: &[u8],
217 pos: usize,
218 shift: u32,
219 acc: u64,
220) -> Result<(u64, usize), PackError> {
221 if shift >= u64::BITS {
222 return Err(malformed("delta size header overflows"));
223 }
224 let byte = *data
225 .get(pos)
226 .ok_or_else(|| malformed("delta size header truncated"))?;
227 let acc = acc | (u64::from(byte & 0x7f) << shift);
228 if byte & 0x80 == 0 {
229 Ok((acc, pos + 1))
230 } else {
231 read_delta_varint(data, pos + 1, shift + 7, acc)
232 }
233}
234
235fn delta_result_size(header: &[u8]) -> Result<u64, PackError> {
236 let (_base_size, after_base) = read_delta_varint(header, 0, 0, 0)?;
237 let (result_size, _) = read_delta_varint(header, after_base, 0, 0)?;
238 Ok(result_size)
239}
240
241pub(crate) fn check_depth(
242 base_of: &HashMap<PackOffset, PackOffset>,
243 max: DeltaDepth,
244) -> Result<(), PackError> {
245 base_of.keys().try_for_each(|start| {
246 let mut depth = DeltaDepth::ZERO;
247 let mut cursor = *start;
248 while let Some(&base) = base_of.get(&cursor) {
249 depth = depth.deeper();
250 if depth.exceeds(max) {
251 return Err(PackError::LimitExceeded(PackLimit::DeltaDepth));
252 }
253 cursor = base;
254 }
255 Ok(())
256 })
257}