This repository has no description
0

Configure Feed

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

core / knot2 / crates / knot-xrpc / src / lfs.rs
38 kB 1168 lines
1use std::collections::{BTreeMap, BTreeSet}; 2use std::sync::{Arc, Mutex}; 3use std::time::Duration; 4 5use axum::Router; 6use axum::body::{Body, Bytes}; 7use axum::extract::rejection::BytesRejection; 8use axum::extract::{DefaultBodyLimit, Path, Request, State}; 9use axum::http::{HeaderMap, HeaderValue, StatusCode, header}; 10use axum::response::{IntoResponse, Response}; 11use axum::routing::{get, post}; 12use knot_lfs::{ 13 BATCH_MEDIA_TYPE, BatchAction, BatchActions, BatchObject, BatchObjectError, BatchOperation, 14 BatchRequest, BatchResponse, BatchResponseObject, ClaimedSize, LfsError, LfsHandle, LfsOid, 15 LfsSize, LfsStore, MAX_BATCH_OBJECTS, UploadAdmission, 16}; 17use knot_pack::{HaveOids, SocketPeer, WantOids}; 18use knot_runtime::{Clock, HttpTransport}; 19use knot_types::{AccountDid, HttpStatus, KnotServiceUrl, RepoDid}; 20use serde_json::json; 21use tokio::sync::{OwnedSemaphorePermit, Semaphore, watch}; 22use tower::ServiceExt; 23use tower_http::services::ServeFile; 24use url::Url; 25 26use crate::XrpcState; 27use crate::forks::Upstream; 28 29pub const MAX_BATCH_BYTES: usize = 1024 * 1024; 30 31const IMMUTABLE_CACHE: &str = "public, max-age=31536000, immutable"; 32 33const READINESS_PROBE_TIMEOUT: Duration = Duration::from_secs(5); 34 35struct Readiness { 36 result: watch::Sender<Option<bool>>, 37 probing: Mutex<bool>, 38} 39 40pub struct LfsWeb { 41 pub handle: LfsHandle, 42 downloads: Arc<Semaphore>, 43 readiness: Arc<Readiness>, 44} 45 46impl LfsWeb { 47 pub fn new(handle: LfsHandle, max_downloads: usize) -> Self { 48 let (result, _rx) = watch::channel(None); 49 Self { 50 handle, 51 downloads: Arc::new(Semaphore::new(max_downloads)), 52 readiness: Arc::new(Readiness { 53 result, 54 probing: Mutex::new(false), 55 }), 56 } 57 } 58 59 pub async fn ready(&self) -> bool { 60 let mut rx = self.readiness.result.subscribe(); 61 let launch = { 62 let mut probing = self.readiness.probing.lock().unwrap(); 63 match *probing { 64 false => { 65 *probing = true; 66 self.readiness.result.send_replace(None); 67 true 68 } 69 true => false, 70 } 71 }; 72 if launch { 73 Self::spawn_probe(Arc::clone(&self.readiness), Arc::clone(&self.handle.store)); 74 } 75 let settled = async { 76 match rx.wait_for(Option::is_some).await { 77 Ok(seen) => seen.unwrap_or(false), 78 Err(_) => false, 79 } 80 }; 81 tokio::time::timeout(READINESS_PROBE_TIMEOUT + Duration::from_secs(1), settled) 82 .await 83 .unwrap_or(false) 84 } 85 86 fn spawn_probe(readiness: Arc<Readiness>, store: Arc<knot_lfs::DiskStore>) { 87 tokio::spawn(async move { 88 let mut guard = ProbeGuard { 89 readiness, 90 outcome: false, 91 }; 92 let joined = tokio::task::spawn_blocking(move || store.probe_ready()).await; 93 let ok = matches!(joined, Ok(Ok(()))); 94 if !ok { 95 tracing::warn!("lfs store readiness probe failed, reporting unready"); 96 } 97 guard.outcome = ok; 98 }); 99 } 100} 101 102struct ProbeGuard { 103 readiness: Arc<Readiness>, 104 outcome: bool, 105} 106 107impl Drop for ProbeGuard { 108 fn drop(&mut self) { 109 *self.readiness.probing.lock().unwrap() = false; 110 self.readiness.result.send_replace(Some(self.outcome)); 111 } 112} 113 114pub(crate) fn routes<H: HttpTransport, C: Clock>() -> Router<Arc<XrpcState<H, C>>> { 115 let batch = Router::new() 116 .route( 117 "/{did}/{name}/info/lfs/objects/batch", 118 post(batch_named::<H, C>), 119 ) 120 .route("/{did}/info/lfs/objects/batch", post(batch_did::<H, C>)) 121 .layer(DefaultBodyLimit::max(MAX_BATCH_BYTES)); 122 let objects = Router::new() 123 .route( 124 "/{did}/{name}/info/lfs/objects/{oid}", 125 get(object_named::<H, C>).put(object_upload_named::<H, C>), 126 ) 127 .route( 128 "/{did}/info/lfs/objects/{oid}", 129 get(object_did::<H, C>).put(object_upload_did::<H, C>), 130 ) 131 .layer(DefaultBodyLimit::disable()); 132 batch.merge(objects) 133} 134 135fn fail(status: StatusCode, message: &str) -> Response { 136 ( 137 status, 138 [( 139 header::CONTENT_TYPE, 140 HeaderValue::from_static(BATCH_MEDIA_TYPE), 141 )], 142 json!({ "message": message }).to_string(), 143 ) 144 .into_response() 145} 146 147fn not_enabled() -> Response { 148 fail(StatusCode::NOT_FOUND, "LFS isn't enabled on this knot") 149} 150 151fn lfs_error(error: crate::XrpcError) -> Box<Response> { 152 Box::new(fail(error.status(), &error.to_string())) 153} 154 155fn resolve_did<H: HttpTransport, C: Clock>( 156 state: &XrpcState<H, C>, 157 segment: &crate::RepoDidSegment, 158) -> Result<RepoDid, Box<Response>> { 159 crate::resolve_repo_did(state, segment).map_err(lfs_error) 160} 161 162async fn resolve_named<H: HttpTransport, C: Clock>( 163 state: &XrpcState<H, C>, 164 owner: &crate::OwnerSegment, 165 name: &crate::RepoNameSegment, 166) -> Result<RepoDid, Box<Response>> { 167 crate::resolve_repo_named(state, owner, name) 168 .await 169 .map_err(lfs_error) 170} 171 172#[derive(serde::Deserialize)] 173#[serde(transparent)] 174struct OidSegment(String); 175 176impl OidSegment { 177 fn as_str(&self) -> &str { 178 &self.0 179 } 180} 181 182#[derive(serde::Deserialize)] 183struct RepoObjectParams { 184 did: crate::OwnerSegment, 185 name: crate::RepoNameSegment, 186 oid: OidSegment, 187} 188 189#[derive(serde::Deserialize)] 190struct DidObjectParams { 191 did: crate::RepoDidSegment, 192 oid: OidSegment, 193} 194 195async fn batch_named<H: HttpTransport, C: Clock>( 196 State(state): State<Arc<XrpcState<H, C>>>, 197 Path(crate::RepoPathParams { did, name }): Path<crate::RepoPathParams>, 198 peer: SocketPeer, 199 headers: HeaderMap, 200 body: Result<Bytes, BytesRejection>, 201) -> Response { 202 let repo = match resolve_named(&state, &did, &name).await { 203 Ok(repo) => repo, 204 Err(response) => return *response, 205 }; 206 serve_batch( 207 &state, 208 repo, 209 &format!("{}/{}", did.as_str(), name.as_str()), 210 peer, 211 &headers, 212 body, 213 ) 214 .await 215} 216 217async fn batch_did<H: HttpTransport, C: Clock>( 218 State(state): State<Arc<XrpcState<H, C>>>, 219 Path(did): Path<crate::RepoDidSegment>, 220 peer: SocketPeer, 221 headers: HeaderMap, 222 body: Result<Bytes, BytesRejection>, 223) -> Response { 224 let repo = match resolve_did(&state, &did) { 225 Ok(repo) => repo, 226 Err(response) => return *response, 227 }; 228 serve_batch(&state, repo, did.as_str(), peer, &headers, body).await 229} 230 231async fn serve_batch<H: HttpTransport, C: Clock>( 232 state: &XrpcState<H, C>, 233 repo: RepoDid, 234 path_prefix: &str, 235 peer: SocketPeer, 236 headers: &HeaderMap, 237 body: Result<Bytes, BytesRejection>, 238) -> Response { 239 let Some(lfs) = state.lfs.as_ref() else { 240 return not_enabled(); 241 }; 242 let body = match body { 243 Ok(body) => body, 244 Err(rejection) => return fail(rejection.status(), &rejection.body_text()), 245 }; 246 let request: BatchRequest = match serde_json::from_slice(&body) { 247 Ok(request) => request, 248 Err(error) => { 249 return fail( 250 StatusCode::UNPROCESSABLE_ENTITY, 251 &format!("invalid batch request: {error}"), 252 ); 253 } 254 }; 255 if request.objects.len() > MAX_BATCH_OBJECTS { 256 return fail( 257 StatusCode::UNPROCESSABLE_ENTITY, 258 &format!("batch exceeds {MAX_BATCH_OBJECTS} objects"), 259 ); 260 } 261 if !request.transfers.is_empty() 262 && !request 263 .transfers 264 .iter() 265 .any(knot_lfs::TransferAdapter::is_basic) 266 { 267 return fail( 268 StatusCode::UNPROCESSABLE_ENTITY, 269 "no mutually supported transfer adapter, this server serves basic", 270 ); 271 } 272 if let Some(algo) = &request.hash_algo 273 && !algo.is_sha256() 274 { 275 return fail( 276 StatusCode::CONFLICT, 277 &format!("unsupported hash algorithm {:?}", algo.as_str()), 278 ); 279 } 280 let base = &state.knot_service_url; 281 let objects: Result<Vec<BatchResponseObject>, Box<Response>> = match request.operation { 282 BatchOperation::Upload => match authorized_pusher(state, peer, headers, &repo).await { 283 Ok(_actor) => match probe_all(lfs, &repo, &request.objects).await { 284 Ok(stored) => request 285 .objects 286 .iter() 287 .zip(stored) 288 .map(|(object, present)| upload_action(base, path_prefix, object, present)) 289 .collect(), 290 Err(response) => Err(response), 291 }, 292 Err(response) => return *response, 293 }, 294 BatchOperation::Download => match probe_all(lfs, &repo, &request.objects).await { 295 Ok(stored) => request 296 .objects 297 .iter() 298 .zip(stored) 299 .map(|(object, stored)| downloadable(base, path_prefix, object, stored)) 300 .collect(), 301 Err(response) => Err(response), 302 }, 303 }; 304 let objects = match objects { 305 Ok(objects) => objects, 306 Err(response) => return *response, 307 }; 308 let response = BatchResponse { 309 transfer: knot_lfs::TransferAdapter::Basic, 310 objects, 311 hash_algo: Some(knot_lfs::HashAlgo::Sha256), 312 }; 313 ( 314 StatusCode::OK, 315 [( 316 header::CONTENT_TYPE, 317 HeaderValue::from_static(BATCH_MEDIA_TYPE), 318 )], 319 serde_json::to_string(&response).expect("batch response always serializes"), 320 ) 321 .into_response() 322} 323 324async fn authorized_pusher<H: HttpTransport, C: Clock>( 325 state: &XrpcState<H, C>, 326 peer: SocketPeer, 327 headers: &HeaderMap, 328 repo: &RepoDid, 329) -> Result<AccountDid, Box<Response>> { 330 crate::authenticate_and_authorize_push( 331 state, 332 peer, 333 headers, 334 repo, 335 "you aren't authorized to push to this repository", 336 ) 337 .await 338 .map_err(|error| Box::new(challenge(error))) 339} 340 341fn challenge(error: crate::XrpcError) -> Response { 342 if error.status() == StatusCode::UNAUTHORIZED { 343 unauthorized(&error.to_string()) 344 } else { 345 error.into_response() 346 } 347} 348 349fn unauthorized(message: &str) -> Response { 350 ( 351 StatusCode::UNAUTHORIZED, 352 [ 353 (header::WWW_AUTHENTICATE, crate::BASIC_CHALLENGE), 354 ( 355 header::CONTENT_TYPE, 356 HeaderValue::from_static(BATCH_MEDIA_TYPE), 357 ), 358 ], 359 json!({ "message": message }).to_string(), 360 ) 361 .into_response() 362} 363 364fn upload_action( 365 base: &KnotServiceUrl, 366 path_prefix: &str, 367 object: &BatchObject, 368 present: Option<LfsSize>, 369) -> Result<BatchResponseObject, Box<Response>> { 370 Ok(match present { 371 Some(size) => BatchResponseObject { 372 oid: object.oid.clone(), 373 size: ClaimedSize::new(size.get()), 374 authenticated: Some(true), 375 actions: None, 376 error: None, 377 }, 378 None => BatchResponseObject { 379 oid: object.oid.clone(), 380 size: object.size, 381 // I know what you're thinking about putting `authenticated: true`, 382 // but trust me TM git-lfs thinks 383 // "the href already has credentials on it" 384 // and omits the Authorization header from the following PUT req. 385 // PUT needs auth so the client would 401, re-run the batch, 386 // repeat. 387 // 388 // Having this be `None` makes git-lfs re-send the header it 389 // used on the batch-call in the first place. 390 authenticated: None, 391 actions: Some(BatchActions { 392 download: None, 393 upload: Some(BatchAction { 394 href: object_href(base, path_prefix, &object.oid)?, 395 }), 396 }), 397 error: None, 398 }, 399 }) 400} 401 402async fn probe_all( 403 lfs: &LfsWeb, 404 repo: &RepoDid, 405 objects: &[BatchObject], 406) -> Result<Vec<Option<LfsSize>>, Box<Response>> { 407 let store = Arc::clone(&lfs.handle.store); 408 let target = repo.clone(); 409 let oids: Vec<LfsOid> = objects.iter().map(|object| object.oid.clone()).collect(); 410 tokio::task::spawn_blocking(move || { 411 oids.iter() 412 .map(|oid| store.probe(&target, oid)) 413 .collect::<Result<Vec<_>, _>>() 414 }) 415 .await 416 .map_err(|_| { 417 Box::new(fail( 418 StatusCode::INTERNAL_SERVER_ERROR, 419 "store probe failed", 420 )) 421 })? 422 .map_err(|error| { 423 tracing::warn!(repo = repo.as_str(), %error, "lfs store probe failed"); 424 Box::new(fail( 425 StatusCode::INTERNAL_SERVER_ERROR, 426 "store probe failed", 427 )) 428 }) 429} 430 431fn downloadable( 432 base: &KnotServiceUrl, 433 path_prefix: &str, 434 object: &BatchObject, 435 stored: Option<LfsSize>, 436) -> Result<BatchResponseObject, Box<Response>> { 437 Ok(match stored { 438 Some(size) => BatchResponseObject { 439 oid: object.oid.clone(), 440 size: ClaimedSize::new(size.get()), 441 authenticated: Some(true), 442 actions: Some(BatchActions { 443 download: Some(BatchAction { 444 href: object_href(base, path_prefix, &object.oid)?, 445 }), 446 upload: None, 447 }), 448 error: None, 449 }, 450 None => BatchResponseObject { 451 oid: object.oid.clone(), 452 size: object.size, 453 authenticated: None, 454 actions: None, 455 error: Some(BatchObjectError { 456 code: HttpStatus::new(404), 457 message: "object not found".to_string(), 458 }), 459 }, 460 }) 461} 462 463fn object_href( 464 base: &KnotServiceUrl, 465 path_prefix: &str, 466 oid: &LfsOid, 467) -> Result<Url, Box<Response>> { 468 Url::parse(&format!( 469 "{}/{path_prefix}/info/lfs/objects/{oid}", 470 base.as_str() 471 )) 472 .map_err(|error| { 473 Box::new(fail( 474 StatusCode::INTERNAL_SERVER_ERROR, 475 &format!("cannot derive object href: {error}"), 476 )) 477 }) 478} 479 480async fn object_named<H: HttpTransport, C: Clock>( 481 State(state): State<Arc<XrpcState<H, C>>>, 482 Path(RepoObjectParams { did, name, oid }): Path<RepoObjectParams>, 483 request: Request, 484) -> Response { 485 let repo = match resolve_named(&state, &did, &name).await { 486 Ok(repo) => repo, 487 Err(response) => return *response, 488 }; 489 serve_object(&state, repo, &oid, request).await 490} 491 492async fn object_did<H: HttpTransport, C: Clock>( 493 State(state): State<Arc<XrpcState<H, C>>>, 494 Path(DidObjectParams { did, oid }): Path<DidObjectParams>, 495 request: Request, 496) -> Response { 497 let repo = match resolve_did(&state, &did) { 498 Ok(repo) => repo, 499 Err(response) => return *response, 500 }; 501 serve_object(&state, repo, &oid, request).await 502} 503 504async fn serve_object<H: HttpTransport, C: Clock>( 505 state: &XrpcState<H, C>, 506 repo: RepoDid, 507 oid_raw: &OidSegment, 508 request: Request, 509) -> Response { 510 let Some(lfs) = state.lfs.as_ref() else { 511 return not_enabled(); 512 }; 513 let Ok(oid) = LfsOid::new(oid_raw.as_str()) else { 514 return fail(StatusCode::NOT_FOUND, "object not found"); 515 }; 516 let located = { 517 let store = Arc::clone(&lfs.handle.store); 518 let target = repo.clone(); 519 let oid = oid.clone(); 520 tokio::task::spawn_blocking(move || store.object_file(&target, &oid)).await 521 }; 522 let (size, path) = match located { 523 Ok(Ok(Some((size, path)))) => (size, path), 524 Ok(Ok(None)) => return fail(StatusCode::NOT_FOUND, "object not found"), 525 Ok(Err(error)) => { 526 tracing::warn!(repo = repo.as_str(), oid = oid.as_str(), %error, "lfs store read failed"); 527 return fail(StatusCode::INTERNAL_SERVER_ERROR, "store read failed"); 528 } 529 Err(_) => return fail(StatusCode::INTERNAL_SERVER_ERROR, "store read failed"), 530 }; 531 let etag = format!("\"{oid}\""); 532 if client_holds_current(request.headers().get(header::IF_NONE_MATCH), &etag) { 533 return not_modified(&etag); 534 } 535 let request = honor_if_range(request, &etag); 536 let permit = match Arc::clone(&lfs.downloads).acquire_owned().await { 537 Ok(permit) => permit, 538 Err(_) => { 539 return fail( 540 StatusCode::SERVICE_UNAVAILABLE, 541 "server is shutting down, retry shortly", 542 ); 543 } 544 }; 545 let served = ServeFile::new(path) 546 .oneshot(request) 547 .await 548 .map(|response| response.map(Body::new)); 549 let mut response = match served { 550 Ok(response) => response, 551 Err(error) => match error {}, 552 }; 553 tracing::info!( 554 repo = repo.as_str(), 555 oid = oid.as_str(), 556 size = size.get(), 557 status = response.status().as_u16(), 558 "lfs object served over http" 559 ); 560 if response.status().is_success() { 561 let headers = response.headers_mut(); 562 headers.insert( 563 header::CONTENT_TYPE, 564 HeaderValue::from_static("application/octet-stream"), 565 ); 566 headers.insert( 567 header::CACHE_CONTROL, 568 HeaderValue::from_static(IMMUTABLE_CACHE), 569 ); 570 if let Ok(value) = HeaderValue::from_str(&etag) { 571 headers.insert(header::ETAG, value); 572 } 573 } 574 response.map(|body| { 575 Body::new(PermitBody { 576 body, 577 _permit: permit, 578 }) 579 }) 580} 581 582async fn object_upload_named<H: HttpTransport, C: Clock>( 583 State(state): State<Arc<XrpcState<H, C>>>, 584 Path(RepoObjectParams { did, name, oid }): Path<RepoObjectParams>, 585 peer: SocketPeer, 586 request: Request, 587) -> Response { 588 let repo = match resolve_named(&state, &did, &name).await { 589 Ok(repo) => repo, 590 Err(response) => return *response, 591 }; 592 serve_object_upload(&state, repo, &oid, peer, request).await 593} 594 595async fn object_upload_did<H: HttpTransport, C: Clock>( 596 State(state): State<Arc<XrpcState<H, C>>>, 597 Path(DidObjectParams { did, oid }): Path<DidObjectParams>, 598 peer: SocketPeer, 599 request: Request, 600) -> Response { 601 let repo = match resolve_did(&state, &did) { 602 Ok(repo) => repo, 603 Err(response) => return *response, 604 }; 605 serve_object_upload(&state, repo, &oid, peer, request).await 606} 607 608async fn serve_object_upload<H: HttpTransport, C: Clock>( 609 state: &XrpcState<H, C>, 610 repo: RepoDid, 611 oid_raw: &OidSegment, 612 peer: SocketPeer, 613 request: Request, 614) -> Response { 615 let Some(lfs) = state.lfs.as_ref() else { 616 return not_enabled(); 617 }; 618 let Ok(oid) = LfsOid::new(oid_raw.as_str()) else { 619 return fail(StatusCode::NOT_FOUND, "object not found"); 620 }; 621 if let Err(response) = authorized_pusher(state, peer, request.headers(), &repo).await { 622 return *response; 623 } 624 let Some(size) = content_length(request.headers()) else { 625 return fail( 626 StatusCode::LENGTH_REQUIRED, 627 "content-length is required for an lfs object upload", 628 ); 629 }; 630 let permit = match lfs.handle.admission.admit(size) { 631 Ok(permit) => permit, 632 Err(error) => return store_fault(&repo, &oid, error), 633 }; 634 let store = Arc::clone(&lfs.handle.store); 635 let target = repo.clone(); 636 let object = oid.clone(); 637 let landed = { 638 use futures::TryStreamExt; 639 let reader = tokio_util::io::StreamReader::new( 640 request 641 .into_body() 642 .into_data_stream() 643 .map_err(std::io::Error::other), 644 ); 645 tokio::task::spawn_blocking(move || { 646 let mut body = tokio_util::io::SyncIoBridge::new(reader); 647 let outcome = store.put(&target, &object, size, &mut body); 648 drop(permit); 649 outcome 650 }) 651 .await 652 }; 653 match landed { 654 Ok(Ok(())) => { 655 tracing::info!( 656 repo = repo.as_str(), 657 oid = oid.as_str(), 658 size = size.get(), 659 "lfs object stored over http" 660 ); 661 StatusCode::OK.into_response() 662 } 663 Ok(Err(error)) => store_fault(&repo, &oid, error), 664 Err(_) => fail(StatusCode::INTERNAL_SERVER_ERROR, "upload task died"), 665 } 666} 667 668fn content_length(headers: &HeaderMap) -> Option<ClaimedSize> { 669 headers 670 .get(header::CONTENT_LENGTH)? 671 .to_str() 672 .ok()? 673 .trim() 674 .parse::<u64>() 675 .ok() 676 .map(ClaimedSize::new) 677} 678 679fn store_fault(repo: &RepoDid, oid: &LfsOid, error: LfsError) -> Response { 680 let status = match &error { 681 LfsError::HashMismatch { .. } | LfsError::SizeMismatch { .. } => { 682 StatusCode::UNPROCESSABLE_ENTITY 683 } 684 LfsError::SizeLimitExceeded { .. } => StatusCode::PAYLOAD_TOO_LARGE, 685 LfsError::FreeSpaceDenied { .. } => StatusCode::INSUFFICIENT_STORAGE, 686 LfsError::BodyRead { .. } => StatusCode::BAD_REQUEST, 687 _ => StatusCode::INTERNAL_SERVER_ERROR, 688 }; 689 if status.is_server_error() { 690 tracing::warn!(repo = repo.as_str(), oid = oid.as_str(), %error, "lfs object upload failed"); 691 } 692 fail(status, &error.to_string()) 693} 694 695fn client_holds_current(if_none_match: Option<&HeaderValue>, etag: &str) -> bool { 696 if_none_match 697 .and_then(|value| value.to_str().ok()) 698 .is_some_and(|value| { 699 value 700 .split(',') 701 .map(str::trim) 702 .any(|candidate| candidate == "*" || candidate.trim_start_matches("W/") == etag) 703 }) 704} 705 706fn not_modified(etag: &str) -> Response { 707 let mut response = StatusCode::NOT_MODIFIED.into_response(); 708 let headers = response.headers_mut(); 709 headers.insert( 710 header::CACHE_CONTROL, 711 HeaderValue::from_static(IMMUTABLE_CACHE), 712 ); 713 if let Ok(value) = HeaderValue::from_str(etag) { 714 headers.insert(header::ETAG, value); 715 } 716 response 717} 718 719fn honor_if_range(mut request: Request, etag: &str) -> Request { 720 let Some(if_range) = request.headers().get(header::IF_RANGE) else { 721 return request; 722 }; 723 let matches = if_range 724 .to_str() 725 .map(|value| value == etag) 726 .unwrap_or(false); 727 let headers = request.headers_mut(); 728 headers.remove(header::IF_RANGE); 729 if !matches { 730 headers.remove(header::RANGE); 731 } 732 request 733} 734 735pub(crate) fn mirror_fork_objects<H: HttpTransport, C: Clock>( 736 state: Arc<XrpcState<H, C>>, 737 upstream: Upstream, 738 fork: RepoDid, 739 wants: WantOids, 740 haves: HaveOids, 741) -> futures::future::BoxFuture<'static, Result<Vec<LfsOid>, crate::XrpcError>> { 742 Box::pin(mirror_fork_objects_inner( 743 state, upstream, fork, wants, haves, 744 )) 745} 746 747async fn mirror_fork_objects_inner<H: HttpTransport, C: Clock>( 748 state: Arc<XrpcState<H, C>>, 749 upstream: Upstream, 750 fork: RepoDid, 751 wants: WantOids, 752 haves: HaveOids, 753) -> Result<Vec<LfsOid>, crate::XrpcError> { 754 let state = &state; 755 let upstream = &upstream; 756 let fork = &fork; 757 let Some(lfs) = state.lfs.as_ref() else { 758 return Ok(Vec::new()); 759 }; 760 let store = Arc::clone(&lfs.handle.store); 761 let admission = Arc::clone(&lfs.handle.admission); 762 763 let needed = { 764 let layout = state.layout.clone(); 765 let fork = fork.clone(); 766 let store = Arc::clone(&store); 767 tokio::task::spawn_blocking(move || -> Result<Vec<(LfsOid, ClaimedSize)>, String> { 768 let repo = layout.open(&fork).map_err(|error| error.to_string())?; 769 knot_lfs::scan_pointers(&repo, wants.wants(), haves.haves()) 770 .map_err(|error| error.to_string())? 771 .into_iter() 772 .map(|(oid, size)| match store.probe(&fork, &oid) { 773 Ok(None) => Ok(Some((oid, size))), 774 Ok(Some(_)) => Ok(None), 775 Err(fault) => Err(fault.to_string()), 776 }) 777 .filter_map(Result::transpose) 778 .collect() 779 }) 780 .await 781 }; 782 let needed = match needed { 783 Ok(Ok(needed)) => needed, 784 Ok(Err(fault)) => { 785 tracing::warn!( 786 repo = fork.as_str(), 787 fault, 788 "lfs pointer scan failed on fork" 789 ); 790 return Err(crate::XrpcError::internal("lfs pointer scan failed")); 791 } 792 Err(_) => { 793 return Err(crate::XrpcError::internal("lfs pointer scan task died")); 794 } 795 }; 796 if needed.is_empty() { 797 return Ok(Vec::new()); 798 } 799 800 let missing = match upstream { 801 Upstream::Local(source) => { 802 let source = source.clone(); 803 let fork = fork.clone(); 804 let store = Arc::clone(&store); 805 let admission = Arc::clone(&admission); 806 tokio::task::spawn_blocking(move || { 807 needed 808 .into_iter() 809 .filter_map(|(oid, size)| { 810 let copied = admission.admit(size).and_then(|_permit| { 811 store 812 .read(&source, &oid) 813 .and_then(|mut body| store.put(&fork, &oid, size, &mut body)) 814 }); 815 match copied { 816 Ok(()) => None, 817 Err(fault) => { 818 tracing::warn!( 819 source = source.as_str(), 820 oid = oid.as_str(), 821 %fault, 822 "lfs fork copy skipped an object" 823 ); 824 Some(oid) 825 } 826 } 827 }) 828 .collect() 829 }) 830 .await 831 .map_err(|_| crate::XrpcError::internal("lfs fork copy task died"))? 832 } 833 Upstream::Remote(url) => match remote_batch_url(url) { 834 Some(batch_url) => { 835 use futures::StreamExt; 836 let chunks: Vec<Vec<(LfsOid, ClaimedSize)>> = needed 837 .chunks(REMOTE_BATCH_CHUNK) 838 .map(<[(LfsOid, ClaimedSize)]>::to_vec) 839 .collect(); 840 futures::stream::iter(chunks) 841 .then(|chunk| { 842 fetch_remote_chunk( 843 Arc::clone(state), 844 Arc::clone(&store), 845 Arc::clone(&admission), 846 fork.clone(), 847 batch_url.clone(), 848 chunk, 849 ) 850 }) 851 .concat() 852 .await 853 } 854 None => needed.into_iter().map(|(oid, _)| oid).collect(), 855 }, 856 }; 857 if !missing.is_empty() { 858 tracing::warn!( 859 repo = fork.as_str(), 860 count = missing.len(), 861 "fork upstream couldn't serve every referenced lfs object" 862 ); 863 } 864 Ok(missing) 865} 866 867const REMOTE_BATCH_CHUNK: usize = 100; 868 869fn remote_batch_url(origin: &Url) -> Option<Url> { 870 let mut origin = origin.clone(); 871 origin.set_query(None); 872 origin.set_fragment(None); 873 let base = origin.as_str().trim_end_matches('/'); 874 let base = match base.ends_with(".git") { 875 true => base.to_string(), 876 false => format!("{base}.git"), 877 }; 878 Url::parse(&format!("{base}/info/lfs/objects/batch")).ok() 879} 880 881async fn fetch_remote_chunk<H: HttpTransport, C: Clock>( 882 state: Arc<XrpcState<H, C>>, 883 store: Arc<knot_lfs::DiskStore>, 884 admission: Arc<knot_lfs::StoreAdmission>, 885 fork: RepoDid, 886 batch_url: Url, 887 chunk: Vec<(LfsOid, ClaimedSize)>, 888) -> Vec<LfsOid> { 889 use futures::StreamExt; 890 let all_missing = || chunk.iter().map(|(oid, _)| oid.clone()).collect::<Vec<_>>(); 891 let request_body = BatchRequest { 892 operation: BatchOperation::Download, 893 transfers: vec![knot_lfs::TransferAdapter::Basic], 894 reference: None, 895 objects: chunk 896 .iter() 897 .map(|(oid, size)| knot_lfs::BatchObject { 898 oid: oid.clone(), 899 size: *size, 900 }) 901 .collect(), 902 hash_algo: Some(knot_lfs::HashAlgo::Sha256), 903 }; 904 let body = match serde_json::to_vec(&request_body) { 905 Ok(body) => body, 906 Err(_) => return all_missing(), 907 }; 908 let mut request = knot_runtime::HttpRequest::post(batch_url.clone(), body.into()); 909 request.headers.insert( 910 header::CONTENT_TYPE, 911 HeaderValue::from_static(BATCH_MEDIA_TYPE), 912 ); 913 request 914 .headers 915 .insert(header::ACCEPT, HeaderValue::from_static(BATCH_MEDIA_TYPE)); 916 let response = match state.git_http.execute(request).await { 917 Ok(response) if response.status.is_success() => response, 918 Ok(response) => { 919 tracing::warn!( 920 url = batch_url.as_str(), 921 status = response.status.as_u16(), 922 "upstream lfs batch refused" 923 ); 924 return all_missing(); 925 } 926 Err(fault) => { 927 tracing::warn!(url = batch_url.as_str(), %fault, "upstream lfs batch failed"); 928 return all_missing(); 929 } 930 }; 931 let parsed: BatchResponse = match serde_json::from_slice(&response.body) { 932 Ok(parsed) => parsed, 933 Err(fault) => { 934 tracing::warn!(url = batch_url.as_str(), %fault, "upstream lfs batch unparsable"); 935 return all_missing(); 936 } 937 }; 938 let declared: BTreeMap<LfsOid, ClaimedSize> = chunk.into_iter().collect(); 939 let unanswered = unanswered_oids(&declared, &parsed.objects); 940 let tagged: Vec<(BatchResponseObject, ClaimedSize)> = parsed 941 .objects 942 .into_iter() 943 .filter_map(|object| { 944 declared 945 .get(&object.oid) 946 .copied() 947 .map(|size| (object, size)) 948 }) 949 .collect(); 950 let failed: Vec<LfsOid> = futures::stream::iter(tagged) 951 .then(|(object, size)| { 952 fetch_remote_object( 953 Arc::clone(&state), 954 Arc::clone(&store), 955 Arc::clone(&admission), 956 fork.clone(), 957 object, 958 size, 959 ) 960 }) 961 .filter_map(std::future::ready) 962 .collect() 963 .await; 964 unanswered.into_iter().chain(failed).collect() 965} 966 967fn unanswered_oids( 968 declared: &BTreeMap<LfsOid, ClaimedSize>, 969 answered: &[BatchResponseObject], 970) -> Vec<LfsOid> { 971 let answered: BTreeSet<&LfsOid> = answered.iter().map(|object| &object.oid).collect(); 972 declared 973 .keys() 974 .filter(|oid| !answered.contains(oid)) 975 .cloned() 976 .collect() 977} 978 979fn href_is_fetchable(url: &Url) -> bool { 980 let scheme_ok = matches!(url.scheme(), "http" | "https"); 981 let host_ok = match url.host() { 982 Some(url::Host::Ipv4(ip)) => !knot_runtime::is_blocked_ip(ip.into()), 983 Some(url::Host::Ipv6(ip)) => !knot_runtime::is_blocked_ip(ip.into()), 984 Some(url::Host::Domain(_)) => true, 985 None => false, 986 }; 987 scheme_ok && host_ok 988} 989 990async fn fetch_remote_object<H: HttpTransport, C: Clock>( 991 state: Arc<XrpcState<H, C>>, 992 store: Arc<knot_lfs::DiskStore>, 993 admission: Arc<knot_lfs::StoreAdmission>, 994 fork: RepoDid, 995 object: BatchResponseObject, 996 declared: ClaimedSize, 997) -> Option<LfsOid> { 998 let Some(action) = object.actions.and_then(|actions| actions.download) else { 999 return Some(object.oid); 1000 }; 1001 if !href_is_fetchable(&action.href) { 1002 tracing::warn!( 1003 oid = object.oid.as_str(), 1004 href = action.href.as_str(), 1005 "lfs fork download href isn't a public http target" 1006 ); 1007 return Some(object.oid); 1008 } 1009 let permit = match admission.admit(declared) { 1010 Ok(permit) => permit, 1011 Err(fault) => { 1012 tracing::warn!(oid = object.oid.as_str(), %fault, "lfs fork download refused by admission"); 1013 return Some(object.oid); 1014 } 1015 }; 1016 let streamed = match state 1017 .git_http 1018 .execute_streamed(knot_runtime::HttpRequest::get(action.href)) 1019 .await 1020 { 1021 Ok(streamed) if streamed.status.is_success() => streamed, 1022 _ => return Some(object.oid), 1023 }; 1024 let landed = { 1025 use futures::TryStreamExt; 1026 let reader = 1027 tokio_util::io::StreamReader::new(streamed.body.map_err(std::io::Error::other)); 1028 let oid = object.oid.clone(); 1029 tokio::task::spawn_blocking(move || { 1030 let mut body = tokio_util::io::SyncIoBridge::new(reader); 1031 let outcome = store.put(&fork, &oid, declared, &mut body); 1032 drop(permit); 1033 outcome 1034 }) 1035 .await 1036 }; 1037 match landed { 1038 Ok(Ok(())) => None, 1039 Ok(Err(fault)) => { 1040 tracing::warn!(oid = object.oid.as_str(), %fault, "lfs fork download failed"); 1041 Some(object.oid) 1042 } 1043 Err(_) => Some(object.oid), 1044 } 1045} 1046 1047struct PermitBody { 1048 body: Body, 1049 _permit: OwnedSemaphorePermit, 1050} 1051 1052impl http_body::Body for PermitBody { 1053 type Data = Bytes; 1054 type Error = axum::Error; 1055 1056 fn poll_frame( 1057 mut self: std::pin::Pin<&mut Self>, 1058 cx: &mut std::task::Context<'_>, 1059 ) -> std::task::Poll<Option<Result<http_body::Frame<Self::Data>, Self::Error>>> { 1060 std::pin::Pin::new(&mut self.body).poll_frame(cx) 1061 } 1062 1063 fn is_end_stream(&self) -> bool { 1064 self.body.is_end_stream() 1065 } 1066 1067 fn size_hint(&self) -> http_body::SizeHint { 1068 self.body.size_hint() 1069 } 1070} 1071 1072#[cfg(test)] 1073mod tests { 1074 use super::*; 1075 1076 #[test] 1077 fn the_remote_batch_endpoint_matches_git_lfs_derivation() { 1078 let plain = Url::parse("https://nel.pet/did:web:witchcraft.systems/anemone").unwrap(); 1079 assert_eq!( 1080 remote_batch_url(&plain).unwrap().as_str(), 1081 "https://nel.pet/did:web:witchcraft.systems/anemone.git/info/lfs/objects/batch" 1082 ); 1083 let suffixed = Url::parse("https://nel.pet/did:plc:cuttle.git").unwrap(); 1084 assert_eq!( 1085 remote_batch_url(&suffixed).unwrap().as_str(), 1086 "https://nel.pet/did:plc:cuttle.git/info/lfs/objects/batch" 1087 ); 1088 let trailing = Url::parse("https://nel.pet/did:plc:cuttle/").unwrap(); 1089 assert_eq!( 1090 remote_batch_url(&trailing).unwrap().as_str(), 1091 "https://nel.pet/did:plc:cuttle.git/info/lfs/objects/batch" 1092 ); 1093 let decorated = Url::parse("https://nel.pet/did:plc:cuttle?ref=main#readme").unwrap(); 1094 assert_eq!( 1095 remote_batch_url(&decorated).unwrap().as_str(), 1096 "https://nel.pet/did:plc:cuttle.git/info/lfs/objects/batch" 1097 ); 1098 } 1099 1100 #[test] 1101 fn oids_the_upstream_batch_never_answers_count_as_missing() { 1102 use sha2::{Digest, Sha256}; 1103 let held = LfsOid::from_digest(Sha256::digest(b"held").into()); 1104 let ignored = LfsOid::from_digest(Sha256::digest(b"ignored").into()); 1105 let declared: BTreeMap<LfsOid, ClaimedSize> = [ 1106 (held.clone(), ClaimedSize::new(4)), 1107 (ignored.clone(), ClaimedSize::new(7)), 1108 ] 1109 .into_iter() 1110 .collect(); 1111 let answered = vec![BatchResponseObject { 1112 oid: held, 1113 size: ClaimedSize::new(4), 1114 authenticated: None, 1115 actions: None, 1116 error: None, 1117 }]; 1118 assert_eq!(unanswered_oids(&declared, &answered), vec![ignored.clone()]); 1119 assert_eq!( 1120 unanswered_oids(&declared, &[]).len(), 1121 2, 1122 "an empty upstream response must leave every oid missing" 1123 ); 1124 } 1125 1126 #[test] 1127 fn a_stale_if_range_drops_the_range_for_a_full_response() { 1128 use axum::http::Request as HttpRequest; 1129 let etag = "\"6c17f2007cbe934aee6e309b28b2fba3c119d98be6ea4156da3aa3173456ad16\""; 1130 1131 let matching = HttpRequest::builder() 1132 .header(header::IF_RANGE, etag) 1133 .header(header::RANGE, "bytes=0-9") 1134 .body(Body::empty()) 1135 .unwrap(); 1136 let kept = honor_if_range(matching, etag); 1137 assert!(kept.headers().get(header::IF_RANGE).is_none()); 1138 assert!( 1139 kept.headers().get(header::RANGE).is_some(), 1140 "a matching validator keeps the range for a 206" 1141 ); 1142 1143 let stale = HttpRequest::builder() 1144 .header(header::IF_RANGE, "\"stale\"") 1145 .header(header::RANGE, "bytes=0-9") 1146 .body(Body::empty()) 1147 .unwrap(); 1148 let full = honor_if_range(stale, etag); 1149 assert!(full.headers().get(header::IF_RANGE).is_none()); 1150 assert!( 1151 full.headers().get(header::RANGE).is_none(), 1152 "a stale validator drops the range so the client gets the whole object" 1153 ); 1154 } 1155 1156 #[test] 1157 fn revalidation_matches_the_oid_etag() { 1158 let etag = "\"6c17f2007cbe934aee6e309b28b2fba3c119d98be6ea4156da3aa3173456ad16\""; 1159 let holds = 1160 |value: &str| client_holds_current(Some(&HeaderValue::from_str(value).unwrap()), etag); 1161 assert!(holds(etag)); 1162 assert!(holds(&format!("W/{etag}"))); 1163 assert!(holds(&format!("\"other\", {etag}"))); 1164 assert!(holds("*")); 1165 assert!(!holds("\"other\"")); 1166 assert!(!client_holds_current(None, etag)); 1167 } 1168}