This repository has no description
1use std::io::{Seek, SeekFrom, Write};
2use std::sync::atomic::AtomicBool;
3
4use gix::bstr::BString;
5use knot_types::{Oid, ParseError};
6
7use crate::error::{GitError, backend};
8use crate::objects::MAX_TREE_DEPTH;
9use crate::repo::Repo;
10
11const TAR_BLOCK: u64 = 512;
12
13knot_types::scalar_newtype! {
14 pub struct ArchiveLimit(u64);
15}
16
17impl Default for ArchiveLimit {
18 fn default() -> Self {
19 Self::new(1024 * 1024 * 1024)
20 }
21}
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
24pub enum ArchiveFormat {
25 Tar,
26 TarGz,
27 Zip,
28}
29
30impl ArchiveFormat {
31 fn gix(self) -> gix_archive::Format {
32 match self {
33 ArchiveFormat::Tar => gix_archive::Format::Tar,
34 ArchiveFormat::TarGz => gix_archive::Format::TarGz {
35 compression_level: None,
36 },
37 ArchiveFormat::Zip => gix_archive::Format::Zip {
38 compression_level: None,
39 },
40 }
41 }
42}
43
44#[derive(Debug, Clone, PartialEq, Eq)]
45pub struct ArchivePrefix(String);
46
47impl ArchivePrefix {
48 pub fn new(value: impl Into<String>) -> Result<Self, ParseError> {
49 let value = value.into();
50 let safe = !value.contains('\0')
51 && !value.starts_with(['/', '\\'])
52 && value.split(['/', '\\']).all(|component| component != "..");
53 match safe {
54 true => Ok(Self(value)),
55 false => Err(ParseError::Invalid {
56 kind: "archive prefix",
57 value,
58 }),
59 }
60 }
61
62 pub fn as_str(&self) -> &str {
63 &self.0
64 }
65}
66
67impl Repo {
68 pub fn peel_to_tree(&self, oid: Oid) -> Result<Oid, GitError> {
69 self.git()
70 .find_object(oid.object_id())
71 .map_err(backend)?
72 .peel_to_tree()
73 .map(|tree| Oid::from(tree.id))
74 .map_err(backend)
75 }
76
77 pub fn write_archive(
78 &self,
79 tree: Oid,
80 format: ArchiveFormat,
81 prefix: Option<&ArchivePrefix>,
82 limit: ArchiveLimit,
83 out: impl std::io::Write + std::io::Seek,
84 ) -> Result<(), GitError> {
85 self.bound_archive_source(tree.object_id(), limit, MAX_TREE_DEPTH, &mut 0)?;
86 let (stream, _index) = self
87 .git()
88 .worktree_stream(tree.object_id())
89 .map_err(backend)?;
90 let interrupt = AtomicBool::new(false);
91 let mut spool = BoundedSpool {
92 inner: out,
93 position: 0,
94 limit,
95 overflowed: false,
96 };
97 let written = self.git().worktree_archive(
98 stream,
99 &mut spool,
100 gix::progress::Discard,
101 &interrupt,
102 gix_archive::Options {
103 format: format.gix(),
104 tree_prefix: prefix.map(|prefix| BString::from(prefix.as_str())),
105 modification_time: 0,
106 },
107 );
108 match (written, spool.overflowed) {
109 (_, true) => Err(GitError::ArchiveTooLarge { limit }),
110 (Ok(()), false) => Ok(()),
111 (Err(error), false) => Err(backend(error)),
112 }
113 }
114
115 fn bound_archive_source(
116 &self,
117 tree: gix::ObjectId,
118 limit: ArchiveLimit,
119 nesting: usize,
120 spooled: &mut u64,
121 ) -> Result<(), GitError> {
122 if nesting == 0 {
123 return Err(GitError::DepthExceeded("tree nesting"));
124 }
125 if tree == gix::ObjectId::empty_tree(self.git().object_hash()) {
126 return Ok(());
127 }
128 let object = self.git().find_tree(tree).map_err(backend)?;
129 let decoded = object
130 .decode()
131 .map_err(|error| GitError::Decode(error.to_string()))?;
132 decoded.entries.iter().try_for_each(|entry| {
133 let oid = entry.oid.to_owned();
134 *spooled = spooled.saturating_add(TAR_BLOCK);
135 match entry.mode.kind() {
136 _ if *spooled > limit.get() => Err(GitError::ArchiveTooLarge { limit }),
137 gix::objs::tree::EntryKind::Commit => Ok(()),
138 gix::objs::tree::EntryKind::Tree => {
139 self.bound_archive_source(oid, limit, nesting - 1, spooled)
140 }
141 _ => {
142 let content = self.blob_size(Oid::from(oid))?;
143 *spooled = spooled.saturating_add(content.next_multiple_of(TAR_BLOCK));
144 match *spooled > limit.get() {
145 true => Err(GitError::ArchiveTooLarge { limit }),
146 false => Ok(()),
147 }
148 }
149 }
150 })
151 }
152}
153
154struct BoundedSpool<W> {
155 inner: W,
156 position: u64,
157 limit: ArchiveLimit,
158 overflowed: bool,
159}
160
161impl<W: Write> Write for BoundedSpool<W> {
162 fn write(&mut self, data: &[u8]) -> std::io::Result<usize> {
163 let remaining = self.limit.get().saturating_sub(self.position);
164 if data.len() as u64 > remaining {
165 self.overflowed = true;
166 return Err(std::io::Error::from(std::io::ErrorKind::WriteZero));
167 }
168 let written = self.inner.write(data)?;
169 self.position = self.position.saturating_add(written as u64);
170 Ok(written)
171 }
172
173 fn flush(&mut self) -> std::io::Result<()> {
174 self.inner.flush()
175 }
176}
177
178impl<W: Seek> Seek for BoundedSpool<W> {
179 fn seek(&mut self, pos: SeekFrom) -> std::io::Result<u64> {
180 let position = self.inner.seek(pos)?;
181 self.position = position;
182 Ok(position)
183 }
184}
185
186#[cfg(test)]
187mod tests {
188 use super::{ArchiveLimit, ArchivePrefix, BoundedSpool};
189 use std::io::{Seek, SeekFrom, Write};
190
191 fn spool(limit: u64) -> BoundedSpool<std::io::Cursor<Vec<u8>>> {
192 BoundedSpool {
193 inner: std::io::Cursor::new(Vec::new()),
194 position: 0,
195 limit: ArchiveLimit::new(limit),
196 overflowed: false,
197 }
198 }
199
200 #[test]
201 fn the_spool_refuses_the_write_that_would_pass_the_limit() {
202 let mut spool = spool(8);
203 assert!(spool.write_all(b"12345678").is_ok());
204 assert!(!spool.overflowed);
205 assert!(spool.write_all(b"9").is_err());
206 assert!(spool.overflowed);
207 assert_eq!(
208 spool.inner.into_inner(),
209 b"12345678",
210 "the refused write never reaches the inner writer"
211 );
212 }
213
214 #[test]
215 fn a_seek_backwards_re_credits_the_budget_the_zip_writer_rewinds_over() {
216 let mut spool = spool(8);
217 spool.write_all(b"12345678").unwrap();
218 spool.seek(SeekFrom::Start(4)).unwrap();
219 assert_eq!(spool.position, 4);
220 spool
221 .write_all(b"abcd")
222 .expect("rewriting bytes already counted stays within the limit");
223 assert!(!spool.overflowed);
224 }
225
226 #[test]
227 fn a_write_whose_length_would_overflow_the_position_is_refused() {
228 let mut spool = spool(u64::MAX);
229 spool.position = u64::MAX;
230 assert!(
231 spool.write_all(b"1").is_err(),
232 "the position saturates at u64::MAX, so the spool must refuse the write"
233 );
234 assert!(spool.overflowed);
235 }
236
237 #[test]
238 fn a_plain_nested_prefix_is_accepted() {
239 assert!(ArchivePrefix::new("squid-main").is_ok());
240 assert!(ArchivePrefix::new("nested/path").is_ok());
241 }
242
243 #[test]
244 fn traversal_is_rejected_across_both_separators() {
245 assert!(ArchivePrefix::new("../escape").is_err());
246 assert!(ArchivePrefix::new("nested/../escape").is_err());
247 assert!(ArchivePrefix::new("..\\escape").is_err());
248 assert!(ArchivePrefix::new("nested\\..\\escape").is_err());
249 assert!(
250 ArchivePrefix::new("dotted-..-name/").is_ok(),
251 "a component that merely contains dot-dot is not a traversal"
252 );
253 }
254
255 #[test]
256 fn absolute_and_null_bearing_prefixes_are_rejected() {
257 assert!(ArchivePrefix::new("/etc").is_err());
258 assert!(ArchivePrefix::new("\\windows").is_err());
259 assert!(ArchivePrefix::new("good\0bad").is_err());
260 }
261}