This repository has no description
0

Configure Feed

Select the types of activity you want to include in your feed.

core / knot2 / crates / knot-pack / src / fetch.rs
18 kB 539 lines
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}