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