diff --git a/knot2/crates/knot-pack/src/fetch.rs b/knot2/crates/knot-pack/src/fetch.rs index bdf3165b9..916406072 100644 --- a/knot2/crates/knot-pack/src/fetch.rs +++ b/knot2/crates/knot-pack/src/fetch.rs @@ -4,10 +4,11 @@ use axum::http::{HeaderMap, HeaderValue, Method, header}; use knot_git::{Filter, RefRecord, Repo}; use knot_runtime::{HttpRequest, HttpResponse, HttpTransport, NetworkError}; use knot_types::{HttpStatus, ObjectFormat, Oid, RefName}; +use tokio_stream::StreamExt; use url::Url; use crate::error::PackError; -use crate::pkt::{self, Frame}; +use crate::pkt; use crate::{HaveOids, WantOids}; #[derive(Debug, thiserror::Error)] @@ -32,6 +33,7 @@ pub enum FetchError { pub struct UpstreamRefs { pub object_format: ObjectFormat, pub head_symref: Option, + pub head_target: Option, pub refs: Vec, } @@ -109,10 +111,25 @@ pub fn parse_advertisement(body: &[u8]) -> Result { }) } -pub fn ls_refs_request(prefixes: &[&str]) -> Result, PackError> { +// git refuses non-sha1 unless we echo the format back +fn write_object_format(buf: &mut Vec, object_format: ObjectFormat) -> Result<(), PackError> { + if object_format != ObjectFormat::SHA1 { + pkt::write_data( + buf, + format!("object-format={}\n", object_format.capability()).as_bytes(), + )?; + } + Ok(()) +} + +pub fn ls_refs_request( + prefixes: &[&str], + object_format: ObjectFormat, +) -> Result, PackError> { let mut buf = Vec::new(); pkt::write_data(&mut buf, b"command=ls-refs\n")?; pkt::write_data(&mut buf, b"agent=knot/0\n")?; + write_object_format(&mut buf, object_format)?; pkt::write_delim(&mut buf)?; pkt::write_data(&mut buf, b"symrefs\n")?; prefixes.iter().try_for_each(|prefix| { @@ -128,6 +145,7 @@ pub fn parse_ls_refs(body: &[u8], object_format: ObjectFormat) -> Result Result { + refs.head_target = Some(target); refs.head_symref = attributes .find_map(|attribute| attribute.strip_prefix("symref-target:")) .and_then(|symref| RefName::new(symref).ok()); @@ -167,10 +186,15 @@ pub fn parse_ls_refs(body: &[u8], object_format: ObjectFormat) -> Result Result, PackError> { +pub fn fetch_request( + wants: &WantOids, + haves: &HaveOids, + object_format: ObjectFormat, +) -> Result, PackError> { let mut buf = Vec::new(); pkt::write_data(&mut buf, b"command=fetch\n")?; pkt::write_data(&mut buf, b"agent=knot/0\n")?; + write_object_format(&mut buf, object_format)?; pkt::write_delim(&mut buf)?; pkt::write_data(&mut buf, b"no-progress\n")?; pkt::write_data(&mut buf, b"ofs-delta\n")?; @@ -185,46 +209,91 @@ pub fn fetch_request(wants: &WantOids, haves: &HaveOids) -> Result, Pack Ok(buf) } -pub fn parse_fetch_response(body: &[u8], max_pack_bytes: u64) -> Result, FetchError> { - let (pack, in_packfile) = pkt::frames(body, None) - .map(|frame| frame.map_err(|error| protocol(error.to_string()))) - .try_fold( - (Vec::new(), false), - |(mut pack, in_packfile), frame| match (frame?.0, in_packfile) { - (Frame::Data(payload), false) => { - if let Some(message) = payload - .strip_prefix(b"ERR ".as_slice()) - .map(|rest| String::from_utf8_lossy(rest).trim_end().to_string()) - { - return Err(FetchError::Remote(message)); - } - let entered = - payload.strip_suffix(b"\n".as_slice()).unwrap_or(payload) == b"packfile"; - Ok((pack, entered)) - } - (Frame::Data(payload), true) => match payload.split_first() { +pub struct PackDemux { + pending: Vec, + in_packfile: bool, + written: u64, + limit: u64, +} + +impl PackDemux { + pub fn new(max_pack_bytes: u64) -> Self { + Self { + pending: Vec::new(), + in_packfile: false, + written: 0, + limit: max_pack_bytes, + } + } + + pub fn feed(&mut self, chunk: &[u8], out: &mut dyn Write) -> Result<(), FetchError> { + self.pending.extend_from_slice(chunk); + let mut done = 0usize; + // one sideband payload held outside `self` so the sink can run while + // the pending buffer is still being drained + let mut scratch = Vec::new(); + loop { + let rest = &self.pending[done..]; + if rest.len() < 4 { + break; + } + let line = pkt::frame_len(&rest[..4]).map_err(|error| protocol(error.to_string()))?; + if line < 4 { + done += 4; + continue; + } + if rest.len() < line { + break; + } + let payload = &rest[4..line]; + if self.in_packfile { + match payload.split_first() { Some((1, data)) => { - if pack.len() as u64 + data.len() as u64 > max_pack_bytes { - return Err(FetchError::PackTooLarge { - limit: max_pack_bytes, - }); - } - pack.extend_from_slice(data); - Ok((pack, true)) + scratch.clear(); + scratch.extend_from_slice(data); + self.admit(&scratch, out)?; } - Some((2, _)) => Ok((pack, true)), - Some((3, message)) => Err(FetchError::Remote( - String::from_utf8_lossy(message).trim_end().to_string(), - )), - _ => Err(protocol("empty sideband frame in packfile section")), - }, - (_, in_packfile) => Ok((pack, in_packfile)), - }, - )?; - if !in_packfile { - return Err(protocol("upstream response has no packfile section")); + Some((2, _)) => {} + Some((3, message)) => { + return Err(FetchError::Remote( + String::from_utf8_lossy(message).trim_end().to_string(), + )); + } + _ => return Err(protocol("empty sideband frame in packfile section")), + } + } else if let Some(message) = payload.strip_prefix(b"ERR ".as_slice()) { + return Err(FetchError::Remote( + String::from_utf8_lossy(message).trim_end().to_string(), + )); + } else { + let line = payload.strip_suffix(b"\n".as_slice()).unwrap_or(payload); + self.in_packfile |= line == b"packfile"; + } + done += line; + } + self.pending.drain(..done); + Ok(()) + } + + fn admit(&mut self, data: &[u8], out: &mut dyn Write) -> Result<(), FetchError> { + if self.written + data.len() as u64 > self.limit { + return Err(FetchError::PackTooLarge { limit: self.limit }); + } + out.write_all(data) + .map_err(|error| FetchError::Pack(error.into()))?; + self.written += data.len() as u64; + Ok(()) + } + + pub fn finish(&mut self) -> Result<(), FetchError> { + if !self.pending.is_empty() { + return Err(protocol("truncated pkt-line in upstream pack response")); + } + if !self.in_packfile { + return Err(protocol("upstream response has no packfile section")); + } + Ok(()) } - Ok(pack) } pub async fn remote_refs( @@ -252,7 +321,7 @@ pub async fn remote_refs( method: Method::POST, url: upload, headers: headers_v2(Some("application/x-git-upload-pack-request")), - body: Some(ls_refs_request(prefixes)?.into()), + body: Some(ls_refs_request(prefixes, object_format)?.into()), }, ) .await?; @@ -264,23 +333,56 @@ pub async fn remote_pack( base: &Url, wants: &WantOids, haves: &HaveOids, + object_format: ObjectFormat, max_pack_bytes: u64, ) -> Result, FetchError> { + let mut pack = Vec::new(); + remote_pack_streamed( + http, + base, + wants, + haves, + object_format, + max_pack_bytes, + &mut pack, + ) + .await?; + Ok(pack) +} + +#[allow(clippy::too_many_arguments)] +pub async fn remote_pack_streamed( + http: &dyn HttpTransport, + base: &Url, + wants: &WantOids, + haves: &HaveOids, + object_format: ObjectFormat, + max_pack_bytes: u64, + out: &mut W, +) -> Result<(), FetchError> { if wants.is_empty() { - return Ok(Vec::new()); + return Ok(()); } let upload = endpoint(base, "git-upload-pack")?; - let response = execute( - http, - HttpRequest { + let response = http + .execute_streamed(HttpRequest { method: Method::POST, url: upload, headers: headers_v2(Some("application/x-git-upload-pack-request")), - body: Some(fetch_request(wants, haves)?.into()), - }, - ) - .await?; - parse_fetch_response(&response.body, max_pack_bytes) + body: Some(fetch_request(wants, haves, object_format)?.into()), + }) + .await?; + if !response.status.is_success() { + return Err(FetchError::Status(HttpStatus::new( + response.status.as_u16(), + ))); + } + let mut body = response.body; + let mut demux = PackDemux::new(max_pack_bytes); + while let Some(chunk) = body.try_next().await? { + demux.feed(&chunk, out)?; + } + demux.finish() } pub fn local_refs(source: &Repo, prefixes: &[&str]) -> Result { @@ -291,42 +393,47 @@ pub fn local_refs(source: &Repo, prefixes: &[&str]) -> Result, +struct BoundedPack<'a> { + out: &'a mut dyn Write, limit: u64, + written: u64, overflowed: bool, } -impl Write for BoundedPack { +impl Write for BoundedPack<'_> { fn write(&mut self, data: &[u8]) -> std::io::Result { - if self.buf.len() as u64 + data.len() as u64 > self.limit { + if self.written + data.len() as u64 > self.limit { self.overflowed = true; return Err(std::io::Error::other("pack byte limit exceeded")); } - self.buf.extend_from_slice(data); + self.out.write_all(data)?; + self.written += data.len() as u64; Ok(data.len()) } fn flush(&mut self) -> std::io::Result<()> { - Ok(()) + self.out.flush() } } -pub fn local_pack( +pub fn local_pack_to( source: &Repo, wants: &WantOids, haves: &HaveOids, max_pack_bytes: u64, -) -> Result, FetchError> { + out: &mut W, +) -> Result<(), FetchError> { if wants.is_empty() { - return Ok(Vec::new()); + return Ok(()); } let oids = source .select_pack_objects_filtered( @@ -338,8 +445,9 @@ pub fn local_pack( .map_err(PackError::from)? .send; let mut out = BoundedPack { - buf: Vec::new(), + out, limit: max_pack_bytes, + written: 0, overflowed: false, }; match crate::objects::write_pack( @@ -349,7 +457,7 @@ pub fn local_pack( &mut out, source.object_format().kind(), ) { - Ok(()) => Ok(out.buf), + Ok(()) => Ok(()), Err(_) if out.overflowed => Err(FetchError::PackTooLarge { limit: max_pack_bytes, }), @@ -357,6 +465,17 @@ pub fn local_pack( } } +pub fn local_pack( + source: &Repo, + wants: &WantOids, + haves: &HaveOids, + max_pack_bytes: u64, +) -> Result, FetchError> { + let mut buf = Vec::new(); + local_pack_to(source, wants, haves, max_pack_bytes, &mut buf)?; + Ok(buf) +} + #[cfg(test)] mod tests { use super::*; @@ -430,6 +549,10 @@ mod tests { refs.head_symref.as_ref().map(RefName::as_str), Some("refs/heads/main") ); + assert_eq!( + refs.head_target.map(|oid| oid.to_string()), + Some("95d09f2b10159347eece71399a7e2e907ea3df4f".to_string()) + ); assert_eq!(refs.refs.len(), 1); assert_eq!(refs.refs[0].name.as_str(), "refs/heads/main"); assert_eq!(refs.tips().len(), 1); @@ -453,6 +576,27 @@ mod tests { )); } + #[test] + fn head_target_survives_when_the_server_omits_its_symref() { + let mut body = Vec::new(); + data( + &mut body, + b"95d09f2b10159347eece71399a7e2e907ea3df4f HEAD\n", + ); + data( + &mut body, + b"95d09f2b10159347eece71399a7e2e907ea3df4f refs/heads/master\n", + ); + pkt::write_flush(&mut body).unwrap(); + + let refs = parse_ls_refs(&body, ObjectFormat::SHA1).unwrap(); + assert_eq!(refs.head_symref, None); + assert_eq!( + refs.head_target, + refs.refs.first().map(|record| record.target) + ); + } + #[test] fn a_malformed_ref_oid_is_a_protocol_error() { let mut body = Vec::new(); @@ -475,16 +619,28 @@ mod tests { )); } + fn fed(body: &[u8], split: usize) -> Result, FetchError> { + let mut demux = PackDemux::new(1024); + let mut pack = Vec::new(); + let (head, tail) = body.split_at(split.min(body.len())); + demux.feed(head, &mut pack)?; + demux.feed(tail, &mut pack)?; + demux.finish()?; + Ok(pack) + } + #[test] - fn the_packfile_section_demuxes_data_and_drops_progress() { + fn the_packfile_section_demuxes_data_drops_progress_and_survives_split_chunks() { let mut body = Vec::new(); data(&mut body, b"packfile\n"); data(&mut body, b"\x01PACKDATA"); data(&mut body, b"\x02counting objects\n"); data(&mut body, b"\x01MORE"); pkt::write_flush(&mut body).unwrap(); - let pack = parse_fetch_response(&body, 1024).unwrap(); - assert_eq!(pack, b"PACKDATAMORE"); + for split in [0, 1, 3, 9, body.len()] { + let pack = fed(&body, split).unwrap(); + assert_eq!(pack, b"PACKDATAMORE", "split at {split}"); + } } #[test] @@ -494,7 +650,7 @@ mod tests { data(&mut body, b"\x03out of disk\n"); pkt::write_flush(&mut body).unwrap(); assert!(matches!( - parse_fetch_response(&body, 1024), + fed(&body, 0), Err(FetchError::Remote(message)) if message == "out of disk" )); @@ -502,8 +658,10 @@ mod tests { data(&mut big, b"packfile\n"); data(&mut big, b"\x01PACKDATA"); pkt::write_flush(&mut big).unwrap(); + let mut demux = PackDemux::new(4); + let mut pack = Vec::new(); assert!(matches!( - parse_fetch_response(&big, 4), + demux.feed(&big, &mut pack), Err(FetchError::PackTooLarge { limit: 4 }) )); } @@ -514,17 +672,41 @@ mod tests { data(&mut body, b"acknowledgments\n"); data(&mut body, b"NAK\n"); pkt::write_flush(&mut body).unwrap(); - assert!(matches!( - parse_fetch_response(&body, 1024), - Err(FetchError::Protocol(_)) - )); + assert!(matches!(fed(&body, 0), Err(FetchError::Protocol(_)))); + } + + #[test] + fn a_truncated_pack_response_is_refused_instead_of_silently_accepted() { + // the trailing pkt-line is cut in half, finish must refuse instead of + // reporting success over a response it never read to the end + let mut body = Vec::new(); + data(&mut body, b"packfile\n"); + data(&mut body, b"\x01PACKDATA"); + let mut demux = PackDemux::new(1024); + let mut pack = Vec::new(); + demux.feed(&body[..body.len() - 2], &mut pack).unwrap(); + assert!(matches!(demux.finish(), Err(FetchError::Protocol(_)))); + } + + #[test] + fn an_err_line_before_the_packfile_section_is_remote() { + let mut body = Vec::new(); + data(&mut body, b"ERR no such object\n"); + assert!( + matches!(fed(&body, 0), Err(FetchError::Remote(message)) if message == "no such object") + ); } #[test] fn the_fetch_request_includes_wants_haves_and_done() { let want = Oid::from_hex("95d09f2b10159347eece71399a7e2e907ea3df4f").unwrap(); let have = Oid::from_hex("2222222222222222222222222222222222222222").unwrap(); - let body = fetch_request(&WantOids::new(vec![want]), &HaveOids::new(vec![have])).unwrap(); + let body = fetch_request( + &WantOids::new(vec![want]), + &HaveOids::new(vec![have]), + ObjectFormat::SHA1, + ) + .unwrap(); let lines = pkt::data_payloads_all(&body).unwrap(); let text: Vec<&str> = lines .iter() @@ -535,5 +717,45 @@ mod tests { assert!(text.contains(&"have 2222222222222222222222222222222222222222")); assert!(text.contains(&"done")); assert!(text.contains(&"no-progress")); + assert!(!text.iter().any(|line| line.starts_with("object-format="))); + } + + #[test] + fn the_requests_repeat_a_non_default_object_format() { + let want = + Oid::from_hex("95d09f2b10159347eece71399a7e2e907ea3df4f0123456789abcdef01234567") + .unwrap(); + let text = |body: &[u8]| -> Vec { + pkt::data_payloads_all(body) + .unwrap() + .iter() + .map(|line| std::str::from_utf8(line).unwrap().trim_end().to_string()) + .collect() + }; + + let fetch = fetch_request( + &WantOids::new(vec![want]), + &HaveOids::default(), + ObjectFormat::SHA256, + ) + .unwrap(); + let lines = text(&fetch); + assert!(lines.contains(&"object-format=sha256".to_string())); + let capability = lines.iter().position(|line| line == "object-format=sha256"); + let delimiter = lines.iter().position(|line| line.as_str() == "0001"); + assert!( + capability < delimiter || delimiter.is_none(), + "the capability belongs in the request header, before the argument section" + ); + + let ls_refs = ls_refs_request(&["refs/heads/"], ObjectFormat::SHA256).unwrap(); + assert!(text(&ls_refs).contains(&"object-format=sha256".to_string())); + + let sha1 = ls_refs_request(&["refs/heads/"], ObjectFormat::SHA1).unwrap(); + assert!( + !text(&sha1) + .iter() + .any(|line| line.starts_with("object-format=")) + ); } } diff --git a/knot2/crates/knot-pack/src/lib.rs b/knot2/crates/knot-pack/src/lib.rs index 33502b0bc..a160932d7 100644 --- a/knot2/crates/knot-pack/src/lib.rs +++ b/knot2/crates/knot-pack/src/lib.rs @@ -38,7 +38,10 @@ use tokio_stream::wrappers::ReceiverStream; pub use cache::{CacheConfig, MaxCacheBytes, MaxEntryBytes}; pub use error::{PackError, PackLimit}; -pub use fetch::{FetchError, UpstreamRefs, local_pack, local_refs, remote_pack, remote_refs}; +pub use fetch::{ + FetchError, PackDemux, UpstreamRefs, local_pack, local_pack_to, local_refs, remote_pack, + remote_pack_streamed, remote_refs, +}; pub use frame::{ ReceiveFramer, UploadFramer, archive_request_complete, receive_request_complete, upload_v0_nak, }; @@ -255,6 +258,15 @@ pub fn ingest_pack( objects::index_pack(objects_dir, pack, limits, kind) } +pub fn ingest_pack_file( + objects_dir: &std::path::Path, + pack_path: &std::path::Path, + limits: &PackLimits, + kind: gix::hash::Kind, +) -> Result<(), PackError> { + objects::ingest_pack_file(objects_dir, pack_path, limits, kind) +} + #[doc(hidden)] pub mod fuzz { pub fn pkt(data: &[u8]) { diff --git a/knot2/crates/knot-pack/src/objects.rs b/knot2/crates/knot-pack/src/objects.rs index 7a08c2bcb..1ec5a5e44 100644 --- a/knot2/crates/knot-pack/src/objects.rs +++ b/knot2/crates/knot-pack/src/objects.rs @@ -1,6 +1,6 @@ use std::collections::{HashMap, HashSet}; use std::error::Error; -use std::io::Write; +use std::io::{Read, Write}; use std::path::{Path, PathBuf}; use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::{Arc, mpsc}; @@ -347,6 +347,29 @@ pub fn index_pack( index_pack_bounded(objects_dir, &file, limits, kind) } +pub fn ingest_pack_file( + objects_dir: &Path, + pack_path: &Path, + limits: &PackLimits, + kind: gix::hash::Kind, +) -> Result<(), PackError> { + let mut signature = [0u8; 4]; + let signature_len = std::fs::File::open(pack_path)? + .read(&mut signature) + .unwrap_or(0); + if signature_len == 0 { + return Err(PackError::Pack("packfile is empty".to_string())); + } + if signature != *b"PACK" { + return Err(PackError::Pack( + "packfile is missing its PACK signature".to_string(), + )); + } + let file = gix_pack::data::File::at(pack_path, kind) + .map_err(|error| PackError::Pack(error.to_string()))?; + index_pack_bounded(objects_dir, &file, limits, kind) +} + pub(crate) fn index_pack_bounded( objects_dir: &Path, pack: &gix_pack::data::File, @@ -967,6 +990,20 @@ mod tests { use super::*; + #[test] + fn an_empty_pack_file_is_not_a_successful_ingest() { + let dir = tempfile::tempdir().unwrap(); + let pack = tempfile::NamedTempFile::new().unwrap(); + let error = ingest_pack_file( + dir.path(), + pack.path(), + &PackLimits::default(), + gix::hash::Kind::Sha1, + ) + .unwrap_err(); + assert!(matches!(error, PackError::Pack(message) if message == "packfile is empty")); + } + #[test] fn interruptible_find_errors_once_the_flag_is_set() { let dir = tempfile::tempdir().unwrap(); diff --git a/knot2/crates/knot-pack/src/pkt.rs b/knot2/crates/knot-pack/src/pkt.rs index 650e44876..fb6d6022d 100644 --- a/knot2/crates/knot-pack/src/pkt.rs +++ b/knot2/crates/knot-pack/src/pkt.rs @@ -18,6 +18,13 @@ fn hex4(prefix: &[u8]) -> io::Result { .ok_or_else(|| io::Error::other("invalid pkt-line length prefix")) } +pub(crate) fn frame_len(prefix: &[u8]) -> io::Result { + match hex4(prefix)? { + 3 => Err(io::Error::other("invalid pkt-line length 3")), + n => Ok(usize::from(n)), + } +} + pub fn frames( input: &[u8], stop_after_flushes: Option, diff --git a/knot2/crates/knot-pack/tests/fetch.rs b/knot2/crates/knot-pack/tests/fetch.rs index 895ebf412..debce0eff 100644 --- a/knot2/crates/knot-pack/tests/fetch.rs +++ b/knot2/crates/knot-pack/tests/fetch.rs @@ -100,6 +100,35 @@ fn knot_server(repo_path: PathBuf) -> Arc { })) } +struct ChunkedHttp { + inner: Arc, + chunk_bytes: usize, +} + +impl HttpTransport for ChunkedHttp { + fn execute(&self, request: HttpRequest) -> knot_runtime::HttpFuture { + self.inner.execute(request) + } + + fn execute_streamed(&self, request: HttpRequest) -> knot_runtime::StreamFuture { + let inner = Arc::clone(&self.inner); + let chunk_bytes = self.chunk_bytes; + Box::pin(async move { + let response = inner.execute(request).await?; + let body = response.body; + let chunks: Vec> = body + .chunks(chunk_bytes) + .map(|chunk| Ok(axum::body::Bytes::copy_from_slice(chunk))) + .collect(); + Ok(knot_runtime::StreamedResponse { + status: response.status, + headers: response.headers, + body: Box::pin(tokio_stream::iter(chunks)), + }) + }) + } +} + fn base_url() -> Url { Url::parse("https://kelp.oyster.cafe/did:plc:squid/uni").unwrap() } @@ -130,6 +159,7 @@ async fn clone_through(http: &dyn HttpTransport, source: &Repo, target: &Repo) { &base_url(), &WantOids::new(refs.tips()), &HaveOids::default(), + refs.object_format, CAP, ) .await @@ -225,6 +255,7 @@ async fn an_incremental_pull_completes_through_staging() { &base_url(), &WantOids::new(vec![new_tip]), &HaveOids::new(vec![old_tip]), + clone.object_format(), CAP, ) .await @@ -267,6 +298,7 @@ async fn a_tiny_pack_limit_refuses_both_remote_and_local_transfers() { &base_url(), &WantOids::new(tips.clone()), &HaveOids::default(), + source.object_format(), 16, ) .await; @@ -332,3 +364,52 @@ fn local_refs_and_pack_mirror_a_same_knot_source() { .unwrap(); assert!(closure.iter().all(|oid| fork.contains(*oid))); } + +#[tokio::test] +async fn a_streamed_fetch_lands_the_pack_on_disk_and_indexes_it_from_there() { + let dir = tempfile::tempdir().unwrap(); + let source_path = seed_source(dir.path()); + let target = Repo::create(dir.path().join("fork.git")).unwrap(); + let http = ChunkedHttp { + inner: stock_git_server(source_path.clone()), + chunk_bytes: 7, + }; + + let refs = knot_pack::remote_refs(&http, &base_url(), &["HEAD", "refs/heads/", "refs/tags/"]) + .await + .unwrap(); + let pack_path = dir.path().join("fetched.pack"); + { + let mut file = std::fs::File::create(&pack_path).unwrap(); + knot_pack::remote_pack_streamed( + &http, + &base_url(), + &WantOids::new(refs.tips()), + &HaveOids::default(), + refs.object_format, + CAP, + &mut file, + ) + .await + .unwrap(); + } + let on_disk = std::fs::read(&pack_path).unwrap(); + assert!(!on_disk.is_empty(), "the streamed fetch wrote a pack"); + assert!(on_disk.starts_with(b"PACK")); + + knot_pack::ingest_pack_file( + &target.objects_dir(), + &pack_path, + &PackLimits::default(), + target.object_format().kind(), + ) + .unwrap(); + let closure = target + .select_pack_objects( + knot_git::Wants::new(&refs.tips()), + knot_git::Haves::new(&[]), + ) + .unwrap(); + assert!(!closure.is_empty()); + assert!(closure.iter().all(|oid| target.contains(*oid))); +} diff --git a/knot2/crates/knot-runtime/src/http.rs b/knot2/crates/knot-runtime/src/http.rs index e99af37a7..9a01ac47b 100644 --- a/knot2/crates/knot-runtime/src/http.rs +++ b/knot2/crates/knot-runtime/src/http.rs @@ -5,7 +5,7 @@ use std::sync::Arc; use std::time::Duration; use bytes::Bytes; -use futures::TryStreamExt; +use futures::{StreamExt, TryStreamExt}; use http::{HeaderMap, Method, StatusCode}; use url::{Host, Url}; @@ -252,6 +252,7 @@ impl HttpTransport for ReqwestHttp { fn execute_streamed(&self, request: HttpRequest) -> StreamFuture { let client = self.client.clone(); let guard = self.block_private_addresses; + let limit = self.max_response_bytes; Box::pin(async move { if let Some(host) = guard.then(|| blocked_literal(&request.url)).flatten() { return Err(NetworkError::Blocked { host }); @@ -265,13 +266,43 @@ impl HttpTransport for ReqwestHttp { let response = builder.send().await.map_err(map_reqwest)?; let status = response.status(); let headers = response.headers().clone(); - let body: ByteStream = Box::pin(response.bytes_stream().map_err(|error| { - if error.is_timeout() { - NetworkError::Timeout(error.to_string()) - } else { - NetworkError::Body(error.to_string()) - } - })); + if response.content_length().is_some_and(|len| len > limit) { + return Err(NetworkError::TooLarge { limit }); + } + let body: ByteStream = Box::pin( + response + .bytes_stream() + .map_err(|error| { + if error.is_timeout() { + NetworkError::Timeout(error.to_string()) + } else { + NetworkError::Body(error.to_string()) + } + }) + .scan((0u64, false), move |state, item| { + let next = if state.1 { + None + } else { + Some(match item { + Ok(chunk) + if state.0.saturating_add(chunk.len() as u64) <= limit => + { + state.0 += chunk.len() as u64; + Ok(chunk) + } + Ok(_) => { + state.1 = true; + Err(NetworkError::TooLarge { limit }) + } + Err(error) => { + state.1 = true; + Err(error) + } + }) + }; + std::future::ready(next) + }), + ); Ok(StreamedResponse { status, headers, @@ -438,6 +469,20 @@ mod tests { transport.execute(HttpRequest::get(url)).await } + async fn fetch_streamed(addr: SocketAddr, limits: HttpLimits) -> Result, NetworkError> { + let transport = ReqwestHttp::new(limits).expect("client builds"); + let url = Url::parse(&format!("http://{addr}/")).expect("url"); + transport + .execute_streamed(HttpRequest::get(url)) + .await? + .body + .try_fold(Vec::new(), |mut body, chunk| async move { + body.extend_from_slice(&chunk); + Ok(body) + }) + .await + } + #[tokio::test] async fn oversized_response_is_rejected_whether_declared_or_streamed() { futures::stream::iter([true, false]) @@ -449,6 +494,15 @@ mod tests { .await; } + #[tokio::test] + async fn streamed_response_keeps_the_same_total_byte_limit() { + for with_content_length in [true, false] { + let addr = serve_body(vec![0u8; 4096], with_content_length); + let result = fetch_streamed(addr, tiny_limits(64, Duration::from_secs(2))).await; + assert!(matches!(result, Err(NetworkError::TooLarge { limit: 64 }))); + } + } + #[tokio::test] async fn small_response_within_limit_succeeds() { let addr = serve_body(b"pong".to_vec(), true);