This repository has no description
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 = ⋔
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}