This repository has no description
1use std::io::Write;
2
3use axum::http::{HeaderMap, HeaderValue, Method, header};
4use knot_git::{Filter, RefRecord, Repo};
5use knot_runtime::{HttpRequest, HttpResponse, HttpTransport, NetworkError};
6use knot_types::{HttpStatus, Oid, RefName};
7use url::Url;
8
9use crate::error::PackError;
10use crate::pkt::{self, Frame};
11use crate::{HaveOids, WantOids};
12
13#[derive(Debug, thiserror::Error)]
14pub enum FetchError {
15 #[error("upstream url: {0}")]
16 Url(String),
17 #[error("upstream network: {0}")]
18 Network(#[from] NetworkError),
19 #[error("upstream returned http status {0}")]
20 Status(HttpStatus),
21 #[error("upstream protocol: {0}")]
22 Protocol(String),
23 #[error("upstream reported: {0}")]
24 Remote(String),
25 #[error("fetched pack exceeds {limit} bytes")]
26 PackTooLarge { limit: u64 },
27 #[error(transparent)]
28 Pack(#[from] PackError),
29}
30
31#[derive(Debug, Clone, PartialEq, Eq)]
32pub struct UpstreamRefs {
33 pub head_symref: Option<RefName>,
34 pub refs: Vec<RefRecord>,
35}
36
37impl UpstreamRefs {
38 pub fn tips(&self) -> Vec<Oid> {
39 let mut tips: Vec<Oid> = self.refs.iter().map(|record| record.target).collect();
40 tips.sort_unstable();
41 tips.dedup();
42 tips
43 }
44
45 pub fn find(&self, name: &RefName) -> Option<Oid> {
46 self.refs
47 .iter()
48 .find(|record| record.name == *name)
49 .map(|record| record.target)
50 }
51}
52
53fn protocol(message: impl Into<String>) -> FetchError {
54 FetchError::Protocol(message.into())
55}
56
57fn endpoint(base: &Url, suffix: &str) -> Result<Url, FetchError> {
58 let trimmed = base.as_str().trim_end_matches('/');
59 Url::parse(&format!("{trimmed}/{suffix}")).map_err(|error| FetchError::Url(error.to_string()))
60}
61
62fn headers_v2(content_type: Option<&'static str>) -> HeaderMap {
63 let mut headers = HeaderMap::new();
64 headers.insert("git-protocol", HeaderValue::from_static("version=2"));
65 if let Some(content_type) = content_type {
66 headers.insert(header::CONTENT_TYPE, HeaderValue::from_static(content_type));
67 }
68 headers
69}
70
71async fn execute(
72 http: &dyn HttpTransport,
73 request: HttpRequest,
74) -> Result<HttpResponse, FetchError> {
75 let response = http.execute(request).await?;
76 if !response.status.is_success() {
77 return Err(FetchError::Status(HttpStatus::new(
78 response.status.as_u16(),
79 )));
80 }
81 Ok(response)
82}
83
84pub fn parse_advertisement(body: &[u8]) -> Result<(), FetchError> {
85 let lines = pkt::data_payloads_all(body).map_err(|error| protocol(error.to_string()))?;
86 let lines: Vec<&str> = lines
87 .iter()
88 .map(|line| std::str::from_utf8(line).unwrap_or_default().trim_end())
89 .collect();
90 let has = |name: &str| {
91 lines
92 .iter()
93 .any(|line| *line == name || line.starts_with(&format!("{name}=")))
94 };
95 if !has("version 2") {
96 return Err(protocol("upstream doesn't speak git protocol v2"));
97 }
98 if !has("ls-refs") || !has("fetch") {
99 return Err(protocol("upstream is missing ls-refs or fetch v2 command"));
100 }
101 Ok(())
102}
103
104pub fn ls_refs_request(prefixes: &[&str]) -> Result<Vec<u8>, PackError> {
105 let mut buf = Vec::new();
106 pkt::write_data(&mut buf, b"command=ls-refs\n")?;
107 pkt::write_data(&mut buf, b"agent=knot/0\n")?;
108 pkt::write_delim(&mut buf)?;
109 pkt::write_data(&mut buf, b"symrefs\n")?;
110 prefixes.iter().try_for_each(|prefix| {
111 pkt::write_data(&mut buf, format!("ref-prefix {prefix}\n").as_bytes())
112 })?;
113 pkt::write_flush(&mut buf)?;
114 Ok(buf)
115}
116
117pub fn parse_ls_refs(body: &[u8]) -> Result<UpstreamRefs, FetchError> {
118 let lines = pkt::data_payloads(body).map_err(|error| protocol(error.to_string()))?;
119 lines.iter().try_fold(
120 UpstreamRefs {
121 head_symref: None,
122 refs: Vec::new(),
123 },
124 |mut refs, line| {
125 let text = std::str::from_utf8(line)
126 .map_err(|_| protocol("ref line isn't utf-8"))?
127 .trim_end();
128 if let Some(message) = text.strip_prefix("ERR ") {
129 return Err(FetchError::Remote(message.to_string()));
130 }
131 let (oid, rest) = text
132 .split_once(' ')
133 .ok_or_else(|| protocol(format!("malformed ref line: {text}")))?;
134 let target =
135 Oid::from_hex(oid).map_err(|_| protocol(format!("malformed ref oid: {oid}")))?;
136 let mut attributes = rest.split(' ');
137 match attributes.next() {
138 Some("HEAD") => {
139 refs.head_symref = attributes
140 .find_map(|attribute| attribute.strip_prefix("symref-target:"))
141 .and_then(|symref| RefName::new(symref).ok());
142 }
143 Some(name) => {
144 if let Ok(name) = RefName::new(name) {
145 refs.refs.push(RefRecord { name, target });
146 }
147 }
148 None => return Err(protocol(format!("malformed ref line: {text}"))),
149 }
150 Ok(refs)
151 },
152 )
153}
154
155pub fn fetch_request(wants: &WantOids, haves: &HaveOids) -> Result<Vec<u8>, PackError> {
156 let mut buf = Vec::new();
157 pkt::write_data(&mut buf, b"command=fetch\n")?;
158 pkt::write_data(&mut buf, b"agent=knot/0\n")?;
159 pkt::write_delim(&mut buf)?;
160 pkt::write_data(&mut buf, b"no-progress\n")?;
161 pkt::write_data(&mut buf, b"ofs-delta\n")?;
162 wants
163 .iter()
164 .try_for_each(|want| pkt::write_data(&mut buf, format!("want {want}\n").as_bytes()))?;
165 haves
166 .iter()
167 .try_for_each(|have| pkt::write_data(&mut buf, format!("have {have}\n").as_bytes()))?;
168 pkt::write_data(&mut buf, b"done\n")?;
169 pkt::write_flush(&mut buf)?;
170 Ok(buf)
171}
172
173pub fn parse_fetch_response(body: &[u8], max_pack_bytes: u64) -> Result<Vec<u8>, FetchError> {
174 let (pack, in_packfile) = pkt::frames(body, None)
175 .map(|frame| frame.map_err(|error| protocol(error.to_string())))
176 .try_fold(
177 (Vec::new(), false),
178 |(mut pack, in_packfile), frame| match (frame?.0, in_packfile) {
179 (Frame::Data(payload), false) => {
180 if let Some(message) = payload
181 .strip_prefix(b"ERR ".as_slice())
182 .map(|rest| String::from_utf8_lossy(rest).trim_end().to_string())
183 {
184 return Err(FetchError::Remote(message));
185 }
186 let entered =
187 payload.strip_suffix(b"\n".as_slice()).unwrap_or(payload) == b"packfile";
188 Ok((pack, entered))
189 }
190 (Frame::Data(payload), true) => match payload.split_first() {
191 Some((1, data)) => {
192 if pack.len() as u64 + data.len() as u64 > max_pack_bytes {
193 return Err(FetchError::PackTooLarge {
194 limit: max_pack_bytes,
195 });
196 }
197 pack.extend_from_slice(data);
198 Ok((pack, true))
199 }
200 Some((2, _)) => Ok((pack, true)),
201 Some((3, message)) => Err(FetchError::Remote(
202 String::from_utf8_lossy(message).trim_end().to_string(),
203 )),
204 _ => Err(protocol("empty sideband frame in packfile section")),
205 },
206 (_, in_packfile) => Ok((pack, in_packfile)),
207 },
208 )?;
209 if !in_packfile {
210 return Err(protocol("upstream response has no packfile section"));
211 }
212 Ok(pack)
213}
214
215pub async fn remote_refs(
216 http: &dyn HttpTransport,
217 base: &Url,
218 prefixes: &[&str],
219) -> Result<UpstreamRefs, FetchError> {
220 let advertise = endpoint(base, "info/refs?service=git-upload-pack")?;
221 let response = execute(
222 http,
223 HttpRequest {
224 method: Method::GET,
225 url: advertise,
226 headers: headers_v2(None),
227 body: None,
228 },
229 )
230 .await?;
231 parse_advertisement(&response.body)?;
232
233 let upload = endpoint(base, "git-upload-pack")?;
234 let response = execute(
235 http,
236 HttpRequest {
237 method: Method::POST,
238 url: upload,
239 headers: headers_v2(Some("application/x-git-upload-pack-request")),
240 body: Some(ls_refs_request(prefixes)?.into()),
241 },
242 )
243 .await?;
244 parse_ls_refs(&response.body)
245}
246
247pub async fn remote_pack(
248 http: &dyn HttpTransport,
249 base: &Url,
250 wants: &WantOids,
251 haves: &HaveOids,
252 max_pack_bytes: u64,
253) -> Result<Vec<u8>, FetchError> {
254 if wants.is_empty() {
255 return Ok(Vec::new());
256 }
257 let upload = endpoint(base, "git-upload-pack")?;
258 let response = execute(
259 http,
260 HttpRequest {
261 method: Method::POST,
262 url: upload,
263 headers: headers_v2(Some("application/x-git-upload-pack-request")),
264 body: Some(fetch_request(wants, haves)?.into()),
265 },
266 )
267 .await?;
268 parse_fetch_response(&response.body, max_pack_bytes)
269}
270
271pub fn local_refs(source: &Repo, prefixes: &[&str]) -> Result<UpstreamRefs, FetchError> {
272 let refs = source
273 .advertised_refs()
274 .map_err(PackError::from)?
275 .iter()
276 .filter(|record| crate::upload::matches_prefix(record.name.as_str(), prefixes))
277 .cloned()
278 .collect();
279 Ok(UpstreamRefs {
280 head_symref: source.head().map(|head| head.name),
281 refs,
282 })
283}
284
285struct BoundedPack {
286 buf: Vec<u8>,
287 limit: u64,
288 overflowed: bool,
289}
290
291impl Write for BoundedPack {
292 fn write(&mut self, data: &[u8]) -> std::io::Result<usize> {
293 if self.buf.len() as u64 + data.len() as u64 > self.limit {
294 self.overflowed = true;
295 return Err(std::io::Error::other("pack byte limit exceeded"));
296 }
297 self.buf.extend_from_slice(data);
298 Ok(data.len())
299 }
300
301 fn flush(&mut self) -> std::io::Result<()> {
302 Ok(())
303 }
304}
305
306pub fn local_pack(
307 source: &Repo,
308 wants: &WantOids,
309 haves: &HaveOids,
310 max_pack_bytes: u64,
311) -> Result<Vec<u8>, FetchError> {
312 if wants.is_empty() {
313 return Ok(Vec::new());
314 }
315 let oids = source
316 .select_pack_objects_filtered(
317 wants.wants(),
318 haves.haves(),
319 Filter::None,
320 crate::upload::selection_budget(),
321 )
322 .map_err(PackError::from)?
323 .send;
324 let mut out = BoundedPack {
325 buf: Vec::new(),
326 limit: max_pack_bytes,
327 overflowed: false,
328 };
329 match crate::objects::write_pack(
330 &source.objects_dir(),
331 oids,
332 None,
333 &mut out,
334 source.object_format().kind(),
335 ) {
336 Ok(()) => Ok(out.buf),
337 Err(_) if out.overflowed => Err(FetchError::PackTooLarge {
338 limit: max_pack_bytes,
339 }),
340 Err(error) => Err(FetchError::Pack(error)),
341 }
342}
343
344#[cfg(test)]
345mod tests {
346 use super::*;
347
348 fn data(buf: &mut Vec<u8>, line: &[u8]) {
349 pkt::write_data(buf, line).unwrap();
350 }
351
352 #[test]
353 fn the_v2_advertisement_is_accepted_and_v0_is_refused() {
354 let scan = tempfile::tempdir().unwrap();
355 let repo = knot_git::Layout::new(scan.path())
356 .create(&knot_types::RepoDid::new("did:plc:squid").unwrap())
357 .unwrap();
358 let v2 = crate::upload::advertise(&repo).unwrap();
359 assert!(parse_advertisement(&v2).is_ok());
360
361 let mut v0 = Vec::new();
362 data(&mut v0, b"# service=git-upload-pack\n");
363 pkt::write_flush(&mut v0).unwrap();
364 data(
365 &mut v0,
366 b"95d09f2b10159347eece71399a7e2e907ea3df4f HEAD\0side-band-64k\n",
367 );
368 pkt::write_flush(&mut v0).unwrap();
369 assert!(matches!(
370 parse_advertisement(&v0),
371 Err(FetchError::Protocol(_))
372 ));
373 }
374
375 #[test]
376 fn ls_refs_lines_parse_with_symref_and_skip_head() {
377 let mut body = Vec::new();
378 data(
379 &mut body,
380 b"95d09f2b10159347eece71399a7e2e907ea3df4f HEAD symref-target:refs/heads/main\n",
381 );
382 data(
383 &mut body,
384 b"95d09f2b10159347eece71399a7e2e907ea3df4f refs/heads/main\n",
385 );
386 pkt::write_flush(&mut body).unwrap();
387 let refs = parse_ls_refs(&body).unwrap();
388 assert_eq!(
389 refs.head_symref.as_ref().map(RefName::as_str),
390 Some("refs/heads/main")
391 );
392 assert_eq!(refs.refs.len(), 1);
393 assert_eq!(refs.refs[0].name.as_str(), "refs/heads/main");
394 assert_eq!(refs.tips().len(), 1);
395 }
396
397 #[test]
398 fn a_malformed_ref_oid_is_a_protocol_error() {
399 let mut body = Vec::new();
400 data(&mut body, b"zzzz refs/heads/main\n");
401 pkt::write_flush(&mut body).unwrap();
402 assert!(matches!(parse_ls_refs(&body), Err(FetchError::Protocol(_))));
403 }
404
405 #[test]
406 fn an_err_line_is_surfaced_as_remote() {
407 let mut body = Vec::new();
408 data(&mut body, b"ERR access denied\n");
409 pkt::write_flush(&mut body).unwrap();
410 assert!(matches!(
411 parse_ls_refs(&body),
412 Err(FetchError::Remote(message)) if message == "access denied"
413 ));
414 }
415
416 #[test]
417 fn the_packfile_section_demuxes_data_and_drops_progress() {
418 let mut body = Vec::new();
419 data(&mut body, b"packfile\n");
420 data(&mut body, b"\x01PACKDATA");
421 data(&mut body, b"\x02counting objects\n");
422 data(&mut body, b"\x01MORE");
423 pkt::write_flush(&mut body).unwrap();
424 let pack = parse_fetch_response(&body, 1024).unwrap();
425 assert_eq!(pack, b"PACKDATAMORE");
426 }
427
428 #[test]
429 fn a_sideband_error_band_is_remote_and_the_limit_holds() {
430 let mut body = Vec::new();
431 data(&mut body, b"packfile\n");
432 data(&mut body, b"\x03out of disk\n");
433 pkt::write_flush(&mut body).unwrap();
434 assert!(matches!(
435 parse_fetch_response(&body, 1024),
436 Err(FetchError::Remote(message)) if message == "out of disk"
437 ));
438
439 let mut big = Vec::new();
440 data(&mut big, b"packfile\n");
441 data(&mut big, b"\x01PACKDATA");
442 pkt::write_flush(&mut big).unwrap();
443 assert!(matches!(
444 parse_fetch_response(&big, 4),
445 Err(FetchError::PackTooLarge { limit: 4 })
446 ));
447 }
448
449 #[test]
450 fn a_response_without_a_packfile_section_is_refused() {
451 let mut body = Vec::new();
452 data(&mut body, b"acknowledgments\n");
453 data(&mut body, b"NAK\n");
454 pkt::write_flush(&mut body).unwrap();
455 assert!(matches!(
456 parse_fetch_response(&body, 1024),
457 Err(FetchError::Protocol(_))
458 ));
459 }
460
461 #[test]
462 fn the_fetch_request_includes_wants_haves_and_done() {
463 let want = Oid::from_hex("95d09f2b10159347eece71399a7e2e907ea3df4f").unwrap();
464 let have = Oid::from_hex("2222222222222222222222222222222222222222").unwrap();
465 let body = fetch_request(&WantOids::new(vec![want]), &HaveOids::new(vec![have])).unwrap();
466 let lines = pkt::data_payloads_all(&body).unwrap();
467 let text: Vec<&str> = lines
468 .iter()
469 .map(|line| std::str::from_utf8(line).unwrap().trim_end())
470 .collect();
471 assert!(text.contains(&"command=fetch"));
472 assert!(text.contains(&"want 95d09f2b10159347eece71399a7e2e907ea3df4f"));
473 assert!(text.contains(&"have 2222222222222222222222222222222222222222"));
474 assert!(text.contains(&"done"));
475 assert!(text.contains(&"no-progress"));
476 }
477}