This repository has no description
1use std::net::IpAddr;
2use std::path::PathBuf;
3use std::sync::Arc;
4use std::time::Duration;
5
6use futures::StreamExt;
7use knot_acl::{KnotAcl, can_push};
8use knot_index::Resolved;
9use knot_lfs::TransferOp;
10use knot_pack::{PackError, PackLimits, RepoLookup};
11use knot_runtime::{Clock, HttpTransport};
12use knot_types::{AccountDid, ClonePath, ObjectFormat, OfferedKey, OwnerDid, RepoDid};
13use russh::Channel;
14use russh::server::Msg;
15use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
16use tokio::runtime::Handle;
17use tokio::sync::mpsc;
18
19use crate::SshState;
20
21const READ_CHUNK: usize = 64 * 1024;
22const MAX_UPLOAD_REQUEST: usize = 16 * 1024 * 1024;
23const RECEIVE_BODY_DEADLINE: Duration = Duration::from_secs(1800);
24const ARCHIVE_REQUEST_DEADLINE: Duration = Duration::from_secs(60);
25const LFS_PROGRESS_GRACE: Duration = Duration::from_secs(60);
26const LFS_PROGRESS_FLOOR_BYTES_PER_SEC: u64 = 1024;
27const LFS_STALL_TIMEOUT: Duration = Duration::from_secs(120);
28
29fn lfs_within_progress_budget(waited: Duration, moved_bytes: u64) -> bool {
30 waited
31 <= LFS_PROGRESS_GRACE + Duration::from_secs(moved_bytes / LFS_PROGRESS_FLOOR_BYTES_PER_SEC)
32}
33
34#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35enum Service {
36 Upload,
37 UploadArchive,
38 Receive,
39 Lfs(TransferOp),
40}
41
42enum ReadError {
43 Io(std::io::Error),
44 Pack(PackError),
45 TooLarge,
46 Truncated,
47 Deadline,
48}
49
50enum RepoRef {
51 Did(RepoDid),
52 OwnerPath(OwnerDid, ClonePath),
53 HandlePath(knot_types::Handle, ClonePath),
54}
55
56enum ResolvedRef {
57 Did(RepoDid),
58 OwnerPath(OwnerDid, ClonePath),
59}
60
61fn parse_exec(command: &[u8]) -> Option<(Service, RepoRef)> {
62 let text = std::str::from_utf8(command).ok()?.trim();
63 if let Some(rest) = text.strip_prefix("git-lfs-transfer ") {
64 let (path, op_token) = rest.trim().rsplit_once(' ')?;
65 let op = TransferOp::parse(op_token.trim())?;
66 return Some((Service::Lfs(op), parse_repo_path(path)?));
67 }
68 let (service, rest) = [
69 ("git-upload-pack ", Service::Upload),
70 ("git upload-pack ", Service::Upload),
71 ("git-upload-archive ", Service::UploadArchive),
72 ("git upload-archive ", Service::UploadArchive),
73 ("git-receive-pack ", Service::Receive),
74 ("git receive-pack ", Service::Receive),
75 ]
76 .into_iter()
77 .find_map(|(prefix, service)| text.strip_prefix(prefix).map(|rest| (service, rest)))?;
78 Some((service, parse_repo_path(rest)?))
79}
80
81fn parse_repo_path(raw: &str) -> Option<RepoRef> {
82 let path = raw
83 .trim()
84 .trim_matches('\'')
85 .trim_matches('"')
86 .trim_start_matches('/');
87 match path.split_once('/') {
88 Some((owner, name)) => {
89 let candidates = ClonePath::parse(name)?;
90 match knot_types::OwnerRef::parse(owner)? {
91 knot_types::OwnerRef::Did(owner) => Some(RepoRef::OwnerPath(owner, candidates)),
92 knot_types::OwnerRef::Handle(handle) => {
93 Some(RepoRef::HandlePath(handle, candidates))
94 }
95 }
96 }
97 None => Some(RepoRef::Did(RepoDid::new(path).ok()?)),
98 }
99}
100
101fn resolve_repo_ref<H: HttpTransport, C: Clock>(
102 state: &Arc<SshState<H, C>>,
103 repo_ref: ResolvedRef,
104) -> RepoLookup {
105 let candidate = match repo_ref {
106 ResolvedRef::Did(did) => RepoLookup::Hosted(did),
107 ResolvedRef::OwnerPath(owner, candidates) => RepoLookup::from_resolved(
108 state.index.resolve_clone_path(&owner, &candidates),
109 |found| found,
110 ),
111 };
112 match candidate {
113 RepoLookup::Hosted(did) => {
114 RepoLookup::from_resolved(state.index.owner_of(&did), |_| did.clone())
115 }
116 undecided => undecided,
117 }
118}
119
120pub(crate) async fn run_exec<H: HttpTransport, C: Clock>(
121 state: Arc<SshState<H, C>>,
122 key: Option<OfferedKey>,
123 channel: Channel<Msg>,
124 command: &[u8],
125 protocol_v2: bool,
126 peer: Option<IpAddr>,
127) {
128 let Some((service, repo_ref)) = parse_exec(command) else {
129 fail(channel, &state.catalog.ssh.unsupported_command.text()).await;
130 return;
131 };
132 let peer_limiter = match &service {
133 Service::Lfs(_) => state
134 .lfs
135 .as_ref()
136 .map_or(&state.peer_slots, |lfs| &lfs.peer_slots),
137 _ => &state.peer_slots,
138 };
139 let _peer_guard = match peer_limiter.admit(peer, state.atproto.now()) {
140 Ok(guard) => guard,
141 Err(refusal) => {
142 let reason = match refusal {
143 knot_resource::Refusal::RateLimited => "peer request rate exceeded",
144 knot_resource::Refusal::Saturated => "peer concurrency limit reached",
145 };
146 tracing::warn!(?peer, reason, "ssh exec rejected");
147 return fail(channel, &state.catalog.ssh.too_many_operations.text()).await;
148 }
149 };
150 let resolved_ref = match repo_ref {
151 RepoRef::Did(did) => ResolvedRef::Did(did),
152 RepoRef::OwnerPath(owner, candidates) => ResolvedRef::OwnerPath(owner, candidates),
153 RepoRef::HandlePath(owner_handle, candidates) => {
154 match state
155 .atproto
156 .resolve_handle_to_did(&owner_handle)
157 .await
158 .ok()
159 {
160 Some(did) => ResolvedRef::OwnerPath(did.into(), candidates),
161 None => {
162 fail(channel, &state.catalog.ssh.repo_not_found.text()).await;
163 return;
164 }
165 }
166 }
167 };
168 let repo_did = match resolve_repo_ref(&state, resolved_ref) {
169 RepoLookup::Hosted(did) => did,
170 RepoLookup::Unhosted => {
171 fail(channel, &state.catalog.ssh.repo_not_found.text()).await;
172 return;
173 }
174 RepoLookup::Unavailable => {
175 fail(channel, &state.catalog.ssh.index_warming.text()).await;
176 return;
177 }
178 };
179 let layout = state.layout.clone();
180 let did = repo_did.clone();
181 let opened = tokio::task::spawn_blocking(move || layout.open(&did).is_ok())
182 .await
183 .unwrap_or(false);
184 if !opened {
185 fail(channel, &state.catalog.ssh.repo_not_found.text()).await;
186 return;
187 }
188 match service {
189 Service::Upload => serve_upload(state, channel, repo_did, protocol_v2).await,
190 Service::UploadArchive => serve_upload_archive(state, channel, repo_did).await,
191 Service::Receive => serve_receive(state, key, channel, repo_did).await,
192 Service::Lfs(op) => serve_lfs(state, key, channel, repo_did, op).await,
193 }
194}
195
196async fn serve_lfs<H: HttpTransport, C: Clock>(
197 state: Arc<SshState<H, C>>,
198 key: Option<OfferedKey>,
199 mut channel: Channel<Msg>,
200 repo_did: RepoDid,
201 op: TransferOp,
202) {
203 let Some(lfs) = state.lfs.clone() else {
204 return fail(channel, &state.catalog.ssh.lfs_disabled.text()).await;
205 };
206 if op == TransferOp::Upload {
207 let pusher = resolve_pusher(&state, key.as_ref(), &repo_did).await;
208 let allowed = pusher.as_ref().is_some_and(|did| {
209 let acl = KnotAcl::new(&state.admins, state.admission, &state.index);
210 can_push(&acl, did, &repo_did).is_allowed()
211 });
212 if !allowed {
213 tracing::warn!(
214 repo = repo_did.as_str(),
215 registered = pusher.is_some(),
216 "ssh lfs upload denied"
217 );
218 let message = match pusher {
219 None => state.catalog.ssh.key_not_registered.text(),
220 Some(_) => state.catalog.ssh.push_denied.text(),
221 };
222 return fail(channel, &message).await;
223 }
224 }
225 let permit = match Arc::clone(&lfs.slots).acquire_owned().await {
226 Ok(permit) => permit,
227 Err(_) => return fail(channel, &state.catalog.ssh.shutting_down.text()).await,
228 };
229 let started = std::time::Instant::now();
230
231 let (tx, rx) = mpsc::channel::<Vec<u8>>(8);
232 let writer = Box::pin(channel.make_writer());
233 let handle = lfs.handle.clone();
234 let runtime = Handle::current();
235 let did = repo_did.clone();
236 let catalog = Arc::clone(&state.catalog);
237 let mut engine = tokio::task::spawn_blocking(move || {
238 let _permit = permit;
239 let output = std::io::BufWriter::new(MeteredWrite::new(runtime.clone(), writer));
240 knot_lfs::serve_transfer(
241 handle.store.as_ref(),
242 handle.admission.as_ref(),
243 &did,
244 op,
245 &catalog.lfs,
246 MpscRead::new(runtime, rx),
247 output,
248 )
249 });
250 let joined = {
251 let reader = channel.make_reader();
252 tokio::select! {
253 joined = &mut engine => joined,
254 () = pump_input(reader, tx) => engine.await,
255 }
256 };
257 let status = match joined {
258 Ok(Ok(())) => {
259 tracing::info!(
260 repo = repo_did.as_str(),
261 op = match op {
262 TransferOp::Upload => "upload",
263 TransferOp::Download => "download",
264 },
265 duration_ms = started.elapsed().as_millis() as u64,
266 "ssh lfs transfer finished"
267 );
268 0
269 }
270 Ok(Err(fault)) => {
271 tracing::warn!(repo = repo_did.as_str(), %fault, "ssh lfs transfer failed");
272 1
273 }
274 Err(join) => {
275 tracing::error!(repo = repo_did.as_str(), %join, "ssh lfs transfer task panicked");
276 1
277 }
278 };
279 finish(channel, status).await;
280}
281
282async fn pump_input<R: AsyncRead + Unpin>(reader: R, tx: mpsc::Sender<Vec<u8>>) {
283 use futures::TryStreamExt;
284 let _ = tokio_util::io::ReaderStream::with_capacity(reader, READ_CHUNK)
285 .map_err(|_| ())
286 .try_for_each(|chunk| {
287 let tx = &tx;
288 async move {
289 match chunk.is_empty() {
290 true => Ok(()),
291 false => tx.send(chunk.to_vec()).await.map_err(|_| ()),
292 }
293 }
294 })
295 .await;
296}
297
298fn stalled(direction: &'static str) -> std::io::Error {
299 std::io::Error::other(format!("lfs {direction} stalled past the idle timeout"))
300}
301
302struct MpscRead {
303 runtime: Handle,
304 rx: mpsc::Receiver<Vec<u8>>,
305 buffer: Vec<u8>,
306 offset: usize,
307 waited: Duration,
308 received: u64,
309}
310
311impl MpscRead {
312 fn new(runtime: Handle, rx: mpsc::Receiver<Vec<u8>>) -> Self {
313 Self {
314 runtime,
315 rx,
316 buffer: Vec::new(),
317 offset: 0,
318 waited: Duration::ZERO,
319 received: 0,
320 }
321 }
322}
323
324impl std::io::Read for MpscRead {
325 fn read(&mut self, out: &mut [u8]) -> std::io::Result<usize> {
326 if self.offset >= self.buffer.len() {
327 let started = std::time::Instant::now();
328 let rx = &mut self.rx;
329 let received = self
330 .runtime
331 // I know I know, but these aren't runtime workers here
332 .block_on(async { tokio::time::timeout(LFS_STALL_TIMEOUT, rx.recv()).await });
333 match received {
334 Ok(Some(chunk)) => {
335 self.waited += started.elapsed();
336 self.received += chunk.len() as u64;
337 if !lfs_within_progress_budget(self.waited, self.received) {
338 return Err(std::io::Error::other(
339 "lfs input trickles below the progress floor",
340 ));
341 }
342 self.buffer = chunk;
343 self.offset = 0;
344 }
345 Ok(None) => return Ok(0),
346 Err(_) => return Err(stalled("input")),
347 }
348 }
349 let take = out.len().min(self.buffer.len() - self.offset);
350 out[..take].copy_from_slice(&self.buffer[self.offset..self.offset + take]);
351 self.offset += take;
352 Ok(take)
353 }
354}
355
356struct MeteredWrite<W> {
357 runtime: Handle,
358 inner: W,
359 waited: Duration,
360 written: u64,
361}
362
363impl<W: AsyncWrite + Unpin> MeteredWrite<W> {
364 fn new(runtime: Handle, inner: W) -> Self {
365 Self {
366 runtime,
367 inner,
368 waited: Duration::ZERO,
369 written: 0,
370 }
371 }
372
373 fn charge(&mut self, started: std::time::Instant) -> std::io::Result<()> {
374 self.waited += started.elapsed();
375 match lfs_within_progress_budget(self.waited, self.written) {
376 true => Ok(()),
377 false => Err(std::io::Error::other(
378 "lfs output trickles below the progress floor",
379 )),
380 }
381 }
382}
383
384impl<W: AsyncWrite + Unpin> std::io::Write for MeteredWrite<W> {
385 fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
386 let started = std::time::Instant::now();
387 let inner = &mut self.inner;
388 let wrote = self
389 .runtime
390 .block_on(async { tokio::time::timeout(LFS_STALL_TIMEOUT, inner.write(buf)).await })
391 .map_err(|_| stalled("output"))??;
392 self.written += wrote as u64;
393 self.charge(started).map(|()| wrote)
394 }
395
396 fn flush(&mut self) -> std::io::Result<()> {
397 let started = std::time::Instant::now();
398 let inner = &mut self.inner;
399 self.runtime
400 .block_on(async { tokio::time::timeout(LFS_STALL_TIMEOUT, inner.flush()).await })
401 .map_err(|_| stalled("output"))??;
402 self.charge(started)
403 }
404}
405
406async fn serve_upload_archive<H: HttpTransport, C: Clock>(
407 state: Arc<SshState<H, C>>,
408 mut channel: Channel<Msg>,
409 repo_did: RepoDid,
410) {
411 let request = {
412 let mut reader = channel.make_reader();
413 tokio::time::timeout(ARCHIVE_REQUEST_DEADLINE, read_archive_request(&mut reader)).await
414 };
415 let request = match request {
416 Ok(Ok(request)) => request,
417 Ok(Err(())) => return fail(channel, &state.catalog.ssh.archive_malformed.text()).await,
418 Err(_) => return fail(channel, &state.catalog.ssh.archive_timeout.text()).await,
419 };
420
421 let permit = state.slots.pack.acquire().await;
422 let (tx, mut rx) = mpsc::channel::<Vec<u8>>(16);
423 let layout = state.layout.clone();
424 let did = repo_did.clone();
425 let handle = tokio::task::spawn_blocking(move || -> Result<(), PackError> {
426 let _permit = permit;
427 let repo = layout.open(&did)?;
428 let mut sink = |chunk: &[u8]| -> std::io::Result<()> {
429 tx.blocking_send(chunk.to_vec())
430 .map_err(|_| std::io::Error::other("client disconnected"))
431 };
432 knot_pack::upload_archive_streamed(&repo, &request, &mut sink)
433 });
434
435 let mut writer = channel.make_writer();
436 let mut forward = Ok(());
437 while let Some(chunk) = rx.recv().await {
438 if writer.write_all(&chunk).await.is_err() {
439 forward = Err(());
440 break;
441 }
442 }
443 drop(rx);
444 let produced = handle.await;
445 match &produced {
446 Ok(Err(error)) => {
447 tracing::warn!(repo = repo_did.as_str(), %error, "upload-archive failed")
448 }
449 Err(join) => {
450 tracing::error!(repo = repo_did.as_str(), %join, "upload-archive task panicked")
451 }
452 Ok(Ok(())) => {}
453 }
454 match (forward, produced) {
455 (Ok(()), Ok(Ok(()))) if writer.flush().await.is_ok() => finish(channel, 0).await,
456 _ => fail(channel, &state.catalog.ssh.archive_failed.text()).await,
457 }
458}
459
460async fn read_archive_request<R: AsyncRead + Unpin>(reader: &mut R) -> Result<Vec<u8>, ()> {
461 let mut buf = Vec::new();
462 loop {
463 if knot_pack::archive_request_complete(&buf).is_some() {
464 return Ok(buf);
465 }
466 match read_chunk(reader, &mut buf, MAX_UPLOAD_REQUEST).await {
467 Ok(true) => {}
468 Ok(false) | Err(_) => return Err(()),
469 }
470 }
471}
472
473async fn serve_upload<H: HttpTransport, C: Clock>(
474 state: Arc<SshState<H, C>>,
475 mut channel: Channel<Msg>,
476 repo_did: RepoDid,
477 protocol_v2: bool,
478) {
479 let advert = {
480 let layout = state.layout.clone();
481 let did = repo_did.clone();
482 tokio::task::spawn_blocking(move || -> Result<Vec<u8>, PackError> {
483 let repo = layout.open(&did)?;
484 if protocol_v2 {
485 knot_pack::advertise_upload_ssh(&repo)
486 } else {
487 knot_pack::advertise_upload_v0_ssh(&repo)
488 }
489 })
490 .await
491 };
492 let advert = match advert {
493 Ok(Ok(bytes)) => bytes,
494 _ => return fail(channel, &state.catalog.ssh.advertise_failed.text()).await,
495 };
496
497 let mut writer = channel.make_writer();
498 if writer.write_all(&advert).await.is_err() || writer.flush().await.is_err() {
499 return;
500 }
501
502 let outcome = {
503 let mut reader = channel.make_reader();
504 if protocol_v2 {
505 upload_loop_v2(&state, &repo_did, &mut reader, &mut writer).await
506 } else {
507 upload_loop_v0(&state, &repo_did, &mut reader, &mut writer).await
508 }
509 };
510 let status = match outcome {
511 Ok(()) => 0,
512 Err(()) => 1,
513 };
514 finish(channel, status).await;
515}
516
517async fn upload_loop_v2<H, C, R, W>(
518 state: &Arc<SshState<H, C>>,
519 repo_did: &RepoDid,
520 reader: &mut R,
521 writer: &mut W,
522) -> Result<(), ()>
523where
524 H: HttpTransport,
525 C: Clock,
526 R: AsyncRead + Unpin,
527 W: AsyncWriteExt + Unpin,
528{
529 let mut buf = Vec::new();
530 let mut framer = knot_pack::UploadFramer::new();
531 loop {
532 if let Some(len) = framer.advance(&buf) {
533 let request: Vec<u8> = buf.drain(..len).collect();
534 stream_upload(state, repo_did, request, writer).await?;
535 framer = knot_pack::UploadFramer::new();
536 continue;
537 }
538 match read_chunk(reader, &mut buf, MAX_UPLOAD_REQUEST).await {
539 Ok(true) => {}
540 Ok(false) => return Ok(()),
541 Err(_) => return Err(()),
542 }
543 }
544}
545
546async fn upload_loop_v0<H, C, R, W>(
547 state: &Arc<SshState<H, C>>,
548 repo_did: &RepoDid,
549 reader: &mut R,
550 writer: &mut W,
551) -> Result<(), ()>
552where
553 H: HttpTransport,
554 C: Clock,
555 R: AsyncRead + Unpin,
556 W: AsyncWriteExt + Unpin,
557{
558 let mut buf = Vec::new();
559 let mut framer = knot_pack::UploadFramer::new();
560 let mut naks_sent = 0usize;
561 loop {
562 if let Some(len) = framer.advance(&buf) {
563 let request: Vec<u8> = buf.drain(..len).collect();
564 return stream_upload(state, repo_did, request, writer).await;
565 }
566 let needed = framer.unanswered_flushes();
567 if naks_sent < needed {
568 let nak = knot_pack::upload_v0_nak();
569 if writer.write_all(&nak).await.is_err() || writer.flush().await.is_err() {
570 return Err(());
571 }
572 naks_sent += 1;
573 continue;
574 }
575 match read_chunk(reader, &mut buf, MAX_UPLOAD_REQUEST).await {
576 Ok(true) => {}
577 Ok(false) => return Ok(()),
578 Err(_) => return Err(()),
579 }
580 }
581}
582
583async fn stream_upload<H, C, W>(
584 state: &Arc<SshState<H, C>>,
585 repo_did: &RepoDid,
586 request: Vec<u8>,
587 writer: &mut W,
588) -> Result<(), ()>
589where
590 H: HttpTransport,
591 C: Clock,
592 W: AsyncWriteExt + Unpin,
593{
594 let permit = state.slots.pack.acquire().await;
595 let (tx, mut rx) = mpsc::channel::<Vec<u8>>(16);
596 let layout = state.layout.clone();
597 let did = repo_did.clone();
598 let catalog = Arc::clone(&state.catalog);
599 let knot = state.hostname.clone();
600 let handle = tokio::task::spawn_blocking(move || -> Result<(), PackError> {
601 let _permit = permit;
602 let repo = layout.open(&did)?;
603 let mut sink = |chunk: &[u8]| -> std::io::Result<()> {
604 tx.blocking_send(chunk.to_vec())
605 .map_err(|_| std::io::Error::other("client disconnected"))
606 };
607 knot_pack::upload_pack_streamed(&repo, &request, &catalog.fetch, &knot, &mut sink)
608 });
609
610 let mut forward = Ok(());
611 while let Some(chunk) = rx.recv().await {
612 if writer.write_all(&chunk).await.is_err() {
613 forward = Err(());
614 break;
615 }
616 }
617 drop(rx);
618 match (forward, handle.await) {
619 (Ok(()), Ok(Ok(()))) => writer.flush().await.map_err(|_| ()),
620 _ => Err(()),
621 }
622}
623
624async fn serve_receive<H: HttpTransport, C: Clock>(
625 state: Arc<SshState<H, C>>,
626 key: Option<OfferedKey>,
627 mut channel: Channel<Msg>,
628 repo_did: RepoDid,
629) {
630 let advert = {
631 let layout = state.layout.clone();
632 let did = repo_did.clone();
633 tokio::task::spawn_blocking(move || -> Result<(Vec<u8>, ObjectFormat), PackError> {
634 let repo = layout.open(&did)?;
635 let bytes = knot_pack::advertise_receive_ssh(&repo)?;
636 Ok((bytes, repo.object_format()))
637 })
638 .await
639 };
640 let (advert, object_format) = match advert {
641 Ok(Ok(pair)) => pair,
642 _ => return fail(channel, &state.catalog.ssh.advertise_failed.text()).await,
643 };
644
645 let mut writer = channel.make_writer();
646 if writer.write_all(&advert).await.is_err() || writer.flush().await.is_err() {
647 return;
648 }
649
650 let pusher = resolve_pusher(&state, key.as_ref(), &repo_did).await;
651 let allowed = |did: &AccountDid| {
652 let acl = KnotAcl::new(&state.admins, state.admission, &state.index);
653 can_push(&acl, did, &repo_did).is_allowed()
654 };
655 let committer = match pusher {
656 Some(did) if allowed(&did) => did,
657 Some(_) => {
658 tracing::warn!(
659 repo = repo_did.as_str(),
660 registered = true,
661 "ssh push denied"
662 );
663 return fail(channel, &state.catalog.ssh.push_denied.text()).await;
664 }
665 None => {
666 tracing::warn!(
667 repo = repo_did.as_str(),
668 registered = false,
669 "ssh push denied"
670 );
671 return fail(channel, &state.catalog.ssh.key_not_registered.text()).await;
672 }
673 };
674
675 let _receive_permit = state.slots.receive.acquire().await;
676
677 let limits = state.limits;
678 let body = {
679 let mut reader = channel.make_reader();
680 let dir = state.layout.scratch_dir().to_path_buf();
681 match tokio::time::timeout(
682 RECEIVE_BODY_DEADLINE,
683 read_receive(
684 &mut reader,
685 dir,
686 state.max_pack_bytes,
687 limits,
688 object_format,
689 ),
690 )
691 .await
692 {
693 Ok(result) => result,
694 Err(_) => Err(ReadError::Deadline),
695 }
696 };
697 let body = match body {
698 Ok(body) => body,
699 Err(ReadError::TooLarge) => {
700 return fail(channel, &state.catalog.ssh.push_too_large.text()).await;
701 }
702 Err(ReadError::Deadline) => {
703 return fail(channel, &state.catalog.ssh.receive_deadline.text()).await;
704 }
705 Err(ReadError::Pack(error)) => {
706 tracing::warn!(repo = repo_did.as_str(), %error, "receive framing failed");
707 return fail(channel, &state.catalog.ssh.malformed_pack.text()).await;
708 }
709 Err(ReadError::Io(error)) => {
710 tracing::warn!(repo = repo_did.as_str(), %error, "receive read error");
711 return fail(channel, &state.catalog.ssh.receive_read_error.text()).await;
712 }
713 Err(ReadError::Truncated) => {
714 return fail(channel, &state.catalog.ssh.receive_ended_early.text()).await;
715 }
716 };
717 if body.is_empty() {
718 return finish(channel, 0).await;
719 }
720
721 let _pack_permit = state.slots.pack.acquire().await;
722 let landed = knot_receive::land(knot_receive::Push {
723 layout: &state.layout,
724 repo_did: &repo_did,
725 received: body,
726 limits: state.limits,
727 knot_actor: state.knot_actor.clone(),
728 committer,
729 events: Arc::clone(&state.events),
730 index: &state.index,
731 atproto: &state.atproto,
732 resolve_slots: &state.slots.resolve,
733 appview: &state.appview,
734 maintenance: &state.maintenance,
735 hostname: &state.hostname,
736 languages_push_budget: state.languages_push_budget,
737 catalog: Arc::clone(&state.catalog),
738 ci_logs: state.ci_logs.clone(),
739 })
740 .await;
741 match landed {
742 Ok(framed) => {
743 let _ = writer.write_all(&framed).await;
744 let _ = writer.flush().await;
745 finish(channel, 0).await;
746 }
747 Err(error) => {
748 tracing::warn!(repo = repo_did.as_str(), %error, "receive-pack failed");
749 fail(channel, &state.catalog.ssh.receive_failed.text()).await;
750 }
751 }
752}
753
754pub(crate) async fn run_greeting<H: HttpTransport, C: Clock>(
755 state: Arc<SshState<H, C>>,
756 key: Option<OfferedKey>,
757 channel: Channel<Msg>,
758) {
759 let who = greeting_identity(&state, key.as_ref()).await;
760 let greeting = state.catalog.ssh.greeting.lines(|key| match key {
761 knot_messages::GreetingKey::User => who.clone(),
762 knot_messages::GreetingKey::Knot => state.hostname.as_str().to_string(),
763 });
764 if greeting.is_empty() {
765 return finish(channel, 0).await;
766 }
767 let body = greeting.join("\r\n");
768 let _ = channel
769 .extended_data_bytes(1, format!("{body}\r\n").into_bytes())
770 .await;
771 finish(channel, 0).await;
772}
773
774async fn greeting_identity<H: HttpTransport, C: Clock>(
775 state: &Arc<SshState<H, C>>,
776 key: Option<&OfferedKey>,
777) -> String {
778 let Some(did) = key.and_then(|key| state.roster.did_for(key)) else {
779 return "there".to_string();
780 };
781 match knot_receive::resolve_handle(&state.atproto, &state.slots.resolve, &did).await {
782 Some(handle) => format!("@{}", handle.as_str()),
783 None => did.as_str().to_string(),
784 }
785}
786
787async fn resolve_pusher<H: HttpTransport, C: Clock>(
788 state: &Arc<SshState<H, C>>,
789 key: Option<&OfferedKey>,
790 repo: &RepoDid,
791) -> Option<AccountDid> {
792 let key = key?;
793 let owner = match state.index.owner_of(repo) {
794 Resolved::Ready(Some(owner)) => Some(AccountDid::from(owner)),
795 _ => None,
796 };
797 {
798 let index = Arc::clone(&state.index);
799 let target = repo.clone();
800 let _ = tokio::task::spawn_blocking(move || index.ensure_collaborators(&target)).await;
801 }
802 let collaborators = match state.index.collaborators_of(repo) {
803 Resolved::Ready(collaborators) => collaborators,
804 _ => Vec::new(),
805 };
806 let candidates: Vec<AccountDid> = owner.into_iter().chain(collaborators).collect();
807 if let Resolved::Ready(Some(cached)) = state.index.owner_of_key(key)
808 && candidates.contains(&cached)
809 {
810 return Some(cached);
811 }
812 let _permit = state.slots.resolve.acquire().await;
813 let matches = futures::stream::iter(candidates).filter_map(|did| async move {
814 let keys = state.atproto.resolve_pubkeys(&did).await.ok()?;
815 keys.iter()
816 .for_each(|resolved| state.index.cache_key(resolved.clone(), &did));
817 keys.iter().any(|resolved| resolved == key).then_some(did)
818 });
819 futures::pin_mut!(matches);
820 matches.next().await
821}
822
823async fn read_chunk<R: AsyncRead + Unpin>(
824 reader: &mut R,
825 buf: &mut Vec<u8>,
826 limit: usize,
827) -> Result<bool, ReadError> {
828 let mut chunk = [0u8; READ_CHUNK];
829 let read = reader.read(&mut chunk).await.map_err(ReadError::Io)?;
830 if read == 0 {
831 return Ok(false);
832 }
833 buf.extend_from_slice(&chunk[..read]);
834 if buf.len() > limit {
835 return Err(ReadError::TooLarge);
836 }
837 Ok(true)
838}
839
840async fn read_receive<R: AsyncRead + Unpin>(
841 reader: &mut R,
842 dir: PathBuf,
843 limit: knot_pack::MaxWireBytes,
844 limits: PackLimits,
845 format: ObjectFormat,
846) -> Result<knot_pack::ReceivedPack, ReadError> {
847 let (tx, rx) = mpsc::channel::<Vec<u8>>(8);
848 let mut framer =
849 tokio::task::spawn_blocking(move || frame_receive(rx, dir, limit, limits, format));
850 let mut chunk = [0u8; READ_CHUNK];
851 let mut io_error = None;
852 loop {
853 tokio::select! {
854 biased;
855 framed = &mut framer => return join_framed(framed, io_error),
856 read = reader.read(&mut chunk) => match read {
857 Ok(0) => break,
858 Ok(read) => {
859 if tx.send(chunk[..read].to_vec()).await.is_err() {
860 break;
861 }
862 }
863 Err(error) => {
864 io_error = Some(error);
865 break;
866 }
867 },
868 }
869 }
870 drop(tx);
871 join_framed(framer.await, io_error)
872}
873
874fn join_framed(
875 framed: Result<Result<knot_pack::ReceivedPack, ReadError>, tokio::task::JoinError>,
876 io_error: Option<std::io::Error>,
877) -> Result<knot_pack::ReceivedPack, ReadError> {
878 match framed {
879 Ok(Ok(body)) => Ok(body),
880 Ok(Err(ReadError::Truncated)) => {
881 Err(io_error.map(ReadError::Io).unwrap_or(ReadError::Truncated))
882 }
883 Ok(Err(other)) => Err(other),
884 Err(_) => Err(ReadError::Truncated),
885 }
886}
887
888fn read_error(error: knot_pack::ReceiveReadError) -> ReadError {
889 match error {
890 knot_pack::ReceiveReadError::Io(error) => ReadError::Io(error),
891 knot_pack::ReceiveReadError::Pack(error) => ReadError::Pack(error),
892 knot_pack::ReceiveReadError::TooLarge => ReadError::TooLarge,
893 knot_pack::ReceiveReadError::Truncated => ReadError::Truncated,
894 }
895}
896
897fn frame_receive(
898 mut rx: mpsc::Receiver<Vec<u8>>,
899 dir: PathBuf,
900 limit: knot_pack::MaxWireBytes,
901 limits: PackLimits,
902 format: ObjectFormat,
903) -> Result<knot_pack::ReceivedPack, ReadError> {
904 let mut receiver =
905 knot_pack::PackReceiver::new(&dir, limit, limits, format.kind()).map_err(ReadError::Io)?;
906 loop {
907 match rx.blocking_recv() {
908 Some(chunk) => {
909 if receiver.write(&chunk).map_err(read_error)? {
910 return receiver.finish().map_err(read_error);
911 }
912 }
913 None => return receiver.finish().map_err(read_error),
914 }
915 }
916}
917
918async fn fail(channel: Channel<Msg>, message: &str) {
919 let _ = channel
920 .extended_data_bytes(1, format!("{message}\n").into_bytes())
921 .await;
922 finish(channel, 1).await;
923}
924
925async fn finish(channel: Channel<Msg>, status: u32) {
926 let _ = channel.exit_status(status).await;
927 let _ = channel.eof().await;
928 let _ = channel.close().await;
929}
930
931#[cfg(test)]
932mod tests {
933 use super::*;
934
935 #[test]
936 fn the_lfs_progress_budget_spares_slow_links_and_cuts_trickles() {
937 assert!(lfs_within_progress_budget(Duration::from_secs(59), 0));
938 assert!(!lfs_within_progress_budget(Duration::from_secs(61), 0));
939 assert!(lfs_within_progress_budget(
940 Duration::from_secs(50_000),
941 5 * 1024 * 1024 * 1024
942 ));
943 assert!(!lfs_within_progress_budget(Duration::from_secs(1_000), 10));
944 }
945
946 #[test]
947 fn the_repo_path_parser_separates_dids_from_handles() {
948 assert!(matches!(
949 parse_repo_path("did:plc:nel/squid"),
950 Some(RepoRef::OwnerPath(..))
951 ));
952 assert!(matches!(
953 parse_repo_path("nel.pet/squid"),
954 Some(RepoRef::HandlePath(..))
955 ));
956 assert!(matches!(
957 parse_repo_path("did:plc:barnacle"),
958 Some(RepoRef::Did(_))
959 ));
960 assert!(parse_repo_path("did:nonsense/squid").is_none());
961 assert!(parse_repo_path("nel.pet").is_none());
962 }
963}