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 archive_limit = state.archive_limit;
426 let handle = tokio::task::spawn_blocking(move || -> Result<(), PackError> {
427 let _permit = permit;
428 let repo = layout.open(&did)?;
429 let mut sink = |chunk: &[u8]| -> std::io::Result<()> {
430 tx.blocking_send(chunk.to_vec())
431 .map_err(|_| std::io::Error::other("client disconnected"))
432 };
433 knot_pack::upload_archive_streamed(&repo, &request, archive_limit, &mut sink)
434 });
435
436 let mut writer = channel.make_writer();
437 let mut forward = Ok(());
438 while let Some(chunk) = rx.recv().await {
439 if writer.write_all(&chunk).await.is_err() {
440 forward = Err(());
441 break;
442 }
443 }
444 drop(rx);
445 let produced = handle.await;
446 match &produced {
447 Ok(Err(error)) => {
448 tracing::warn!(repo = repo_did.as_str(), %error, "upload-archive failed")
449 }
450 Err(join) => {
451 tracing::error!(repo = repo_did.as_str(), %join, "upload-archive task panicked")
452 }
453 Ok(Ok(())) => {}
454 }
455 match (forward, produced) {
456 (Ok(()), Ok(Ok(()))) if writer.flush().await.is_ok() => finish(channel, 0).await,
457 _ => fail(channel, &state.catalog.ssh.archive_failed.text()).await,
458 }
459}
460
461async fn read_archive_request<R: AsyncRead + Unpin>(reader: &mut R) -> Result<Vec<u8>, ()> {
462 let mut buf = Vec::new();
463 loop {
464 if knot_pack::archive_request_complete(&buf).is_some() {
465 return Ok(buf);
466 }
467 match read_chunk(reader, &mut buf, MAX_UPLOAD_REQUEST).await {
468 Ok(true) => {}
469 Ok(false) | Err(_) => return Err(()),
470 }
471 }
472}
473
474async fn serve_upload<H: HttpTransport, C: Clock>(
475 state: Arc<SshState<H, C>>,
476 mut channel: Channel<Msg>,
477 repo_did: RepoDid,
478 protocol_v2: bool,
479) {
480 let advert = {
481 let layout = state.layout.clone();
482 let did = repo_did.clone();
483 tokio::task::spawn_blocking(move || -> Result<Vec<u8>, PackError> {
484 let repo = layout.open(&did)?;
485 if protocol_v2 {
486 knot_pack::advertise_upload_ssh(&repo)
487 } else {
488 knot_pack::advertise_upload_v0_ssh(&repo)
489 }
490 })
491 .await
492 };
493 let advert = match advert {
494 Ok(Ok(bytes)) => bytes,
495 _ => return fail(channel, &state.catalog.ssh.advertise_failed.text()).await,
496 };
497
498 let mut writer = channel.make_writer();
499 if writer.write_all(&advert).await.is_err() || writer.flush().await.is_err() {
500 return;
501 }
502
503 let outcome = {
504 let mut reader = channel.make_reader();
505 if protocol_v2 {
506 upload_loop_v2(&state, &repo_did, &mut reader, &mut writer).await
507 } else {
508 upload_loop_v0(&state, &repo_did, &mut reader, &mut writer).await
509 }
510 };
511 let status = match outcome {
512 Ok(()) => 0,
513 Err(()) => 1,
514 };
515 finish(channel, status).await;
516}
517
518async fn upload_loop_v2<H, C, R, W>(
519 state: &Arc<SshState<H, C>>,
520 repo_did: &RepoDid,
521 reader: &mut R,
522 writer: &mut W,
523) -> Result<(), ()>
524where
525 H: HttpTransport,
526 C: Clock,
527 R: AsyncRead + Unpin,
528 W: AsyncWriteExt + Unpin,
529{
530 let mut buf = Vec::new();
531 let mut framer = knot_pack::UploadFramer::new();
532 loop {
533 if let Some(len) = framer.advance(&buf) {
534 let request: Vec<u8> = buf.drain(..len).collect();
535 stream_upload(state, repo_did, request, writer).await?;
536 framer = knot_pack::UploadFramer::new();
537 continue;
538 }
539 match read_chunk(reader, &mut buf, MAX_UPLOAD_REQUEST).await {
540 Ok(true) => {}
541 Ok(false) => return Ok(()),
542 Err(_) => return Err(()),
543 }
544 }
545}
546
547async fn upload_loop_v0<H, C, R, W>(
548 state: &Arc<SshState<H, C>>,
549 repo_did: &RepoDid,
550 reader: &mut R,
551 writer: &mut W,
552) -> Result<(), ()>
553where
554 H: HttpTransport,
555 C: Clock,
556 R: AsyncRead + Unpin,
557 W: AsyncWriteExt + Unpin,
558{
559 let mut buf = Vec::new();
560 let mut framer = knot_pack::UploadFramer::new();
561 let mut naks_sent = 0usize;
562 loop {
563 if let Some(len) = framer.advance(&buf) {
564 let request: Vec<u8> = buf.drain(..len).collect();
565 return stream_upload(state, repo_did, request, writer).await;
566 }
567 let needed = framer.unanswered_flushes();
568 if naks_sent < needed {
569 let nak = knot_pack::upload_v0_nak();
570 if writer.write_all(&nak).await.is_err() || writer.flush().await.is_err() {
571 return Err(());
572 }
573 naks_sent += 1;
574 continue;
575 }
576 match read_chunk(reader, &mut buf, MAX_UPLOAD_REQUEST).await {
577 Ok(true) => {}
578 Ok(false) => return Ok(()),
579 Err(_) => return Err(()),
580 }
581 }
582}
583
584async fn stream_upload<H, C, W>(
585 state: &Arc<SshState<H, C>>,
586 repo_did: &RepoDid,
587 request: Vec<u8>,
588 writer: &mut W,
589) -> Result<(), ()>
590where
591 H: HttpTransport,
592 C: Clock,
593 W: AsyncWriteExt + Unpin,
594{
595 let permit = state.slots.pack.acquire().await;
596 let (tx, mut rx) = mpsc::channel::<Vec<u8>>(16);
597 let layout = state.layout.clone();
598 let did = repo_did.clone();
599 let catalog = Arc::clone(&state.catalog);
600 let knot = state.hostname.clone();
601 let handle = tokio::task::spawn_blocking(move || -> Result<(), PackError> {
602 let _permit = permit;
603 let repo = layout.open(&did)?;
604 let mut sink = |chunk: &[u8]| -> std::io::Result<()> {
605 tx.blocking_send(chunk.to_vec())
606 .map_err(|_| std::io::Error::other("client disconnected"))
607 };
608 knot_pack::upload_pack_streamed(&repo, &request, &catalog.fetch, &knot, &mut sink)
609 });
610
611 let mut forward = Ok(());
612 while let Some(chunk) = rx.recv().await {
613 if writer.write_all(&chunk).await.is_err() {
614 forward = Err(());
615 break;
616 }
617 }
618 drop(rx);
619 match (forward, handle.await) {
620 (Ok(()), Ok(Ok(()))) => writer.flush().await.map_err(|_| ()),
621 _ => Err(()),
622 }
623}
624
625async fn serve_receive<H: HttpTransport, C: Clock>(
626 state: Arc<SshState<H, C>>,
627 key: Option<OfferedKey>,
628 mut channel: Channel<Msg>,
629 repo_did: RepoDid,
630) {
631 let advert = {
632 let layout = state.layout.clone();
633 let did = repo_did.clone();
634 tokio::task::spawn_blocking(move || -> Result<(Vec<u8>, ObjectFormat), PackError> {
635 let repo = layout.open(&did)?;
636 let bytes = knot_pack::advertise_receive_ssh(&repo)?;
637 Ok((bytes, repo.object_format()))
638 })
639 .await
640 };
641 let (advert, object_format) = match advert {
642 Ok(Ok(pair)) => pair,
643 _ => return fail(channel, &state.catalog.ssh.advertise_failed.text()).await,
644 };
645
646 let mut writer = channel.make_writer();
647 if writer.write_all(&advert).await.is_err() || writer.flush().await.is_err() {
648 return;
649 }
650
651 let pusher = resolve_pusher(&state, key.as_ref(), &repo_did).await;
652 let allowed = |did: &AccountDid| {
653 let acl = KnotAcl::new(&state.admins, state.admission, &state.index);
654 can_push(&acl, did, &repo_did).is_allowed()
655 };
656 let committer = match pusher {
657 Some(did) if allowed(&did) => did,
658 Some(_) => {
659 tracing::warn!(
660 repo = repo_did.as_str(),
661 registered = true,
662 "ssh push denied"
663 );
664 return fail(channel, &state.catalog.ssh.push_denied.text()).await;
665 }
666 None => {
667 tracing::warn!(
668 repo = repo_did.as_str(),
669 registered = false,
670 "ssh push denied"
671 );
672 return fail(channel, &state.catalog.ssh.key_not_registered.text()).await;
673 }
674 };
675
676 let _receive_permit = state.slots.receive.acquire().await;
677
678 let limits = state.limits;
679 let body = {
680 let mut reader = channel.make_reader();
681 let dir = state.layout.scratch_dir().to_path_buf();
682 match tokio::time::timeout(
683 RECEIVE_BODY_DEADLINE,
684 read_receive(
685 &mut reader,
686 dir,
687 state.max_pack_bytes,
688 limits,
689 object_format,
690 ),
691 )
692 .await
693 {
694 Ok(result) => result,
695 Err(_) => Err(ReadError::Deadline),
696 }
697 };
698 let body = match body {
699 Ok(body) => body,
700 Err(ReadError::TooLarge) => {
701 return fail(channel, &state.catalog.ssh.push_too_large.text()).await;
702 }
703 Err(ReadError::Deadline) => {
704 return fail(channel, &state.catalog.ssh.receive_deadline.text()).await;
705 }
706 Err(ReadError::Pack(error)) => {
707 tracing::warn!(repo = repo_did.as_str(), %error, "receive framing failed");
708 return fail(channel, &state.catalog.ssh.malformed_pack.text()).await;
709 }
710 Err(ReadError::Io(error)) => {
711 tracing::warn!(repo = repo_did.as_str(), %error, "receive read error");
712 return fail(channel, &state.catalog.ssh.receive_read_error.text()).await;
713 }
714 Err(ReadError::Truncated) => {
715 return fail(channel, &state.catalog.ssh.receive_ended_early.text()).await;
716 }
717 };
718 if body.is_empty() {
719 return finish(channel, 0).await;
720 }
721
722 let _pack_permit = state.slots.pack.acquire().await;
723 let landed = knot_receive::land(knot_receive::Push {
724 layout: &state.layout,
725 repo_did: &repo_did,
726 received: body,
727 limits: state.limits,
728 knot_actor: state.knot_actor.clone(),
729 committer,
730 events: Arc::clone(&state.events),
731 index: &state.index,
732 atproto: &state.atproto,
733 resolve_slots: &state.slots.resolve,
734 appview: &state.appview,
735 maintenance: &state.maintenance,
736 hostname: &state.hostname,
737 languages_push_budget: state.languages_push_budget,
738 catalog: Arc::clone(&state.catalog),
739 ci_logs: state.ci_logs.clone(),
740 })
741 .await;
742 match landed {
743 Ok(framed) => {
744 let _ = writer.write_all(&framed).await;
745 let _ = writer.flush().await;
746 finish(channel, 0).await;
747 }
748 Err(error) => {
749 tracing::warn!(repo = repo_did.as_str(), %error, "receive-pack failed");
750 fail(channel, &state.catalog.ssh.receive_failed.text()).await;
751 }
752 }
753}
754
755pub(crate) async fn run_greeting<H: HttpTransport, C: Clock>(
756 state: Arc<SshState<H, C>>,
757 key: Option<OfferedKey>,
758 channel: Channel<Msg>,
759) {
760 let who = greeting_identity(&state, key.as_ref()).await;
761 let greeting = state.catalog.ssh.greeting.lines(|key| match key {
762 knot_messages::GreetingKey::User => who.clone(),
763 knot_messages::GreetingKey::Knot => state.hostname.as_str().to_string(),
764 });
765 if greeting.is_empty() {
766 return finish(channel, 0).await;
767 }
768 let body = greeting.join("\r\n");
769 let _ = channel
770 .extended_data_bytes(1, format!("{body}\r\n").into_bytes())
771 .await;
772 finish(channel, 0).await;
773}
774
775async fn greeting_identity<H: HttpTransport, C: Clock>(
776 state: &Arc<SshState<H, C>>,
777 key: Option<&OfferedKey>,
778) -> String {
779 let Some(did) = key.and_then(|key| state.roster.did_for(key)) else {
780 return "there".to_string();
781 };
782 match knot_receive::resolve_handle(&state.atproto, &state.slots.resolve, &did).await {
783 Some(handle) => format!("@{}", handle.as_str()),
784 None => did.as_str().to_string(),
785 }
786}
787
788async fn resolve_pusher<H: HttpTransport, C: Clock>(
789 state: &Arc<SshState<H, C>>,
790 key: Option<&OfferedKey>,
791 repo: &RepoDid,
792) -> Option<AccountDid> {
793 let key = key?;
794 let owner = match state.index.owner_of(repo) {
795 Resolved::Ready(Some(owner)) => Some(AccountDid::from(owner)),
796 _ => None,
797 };
798 {
799 let index = Arc::clone(&state.index);
800 let target = repo.clone();
801 let _ = tokio::task::spawn_blocking(move || index.ensure_collaborators(&target)).await;
802 }
803 let collaborators = match state.index.collaborators_of(repo) {
804 Resolved::Ready(collaborators) => collaborators,
805 _ => Vec::new(),
806 };
807 let candidates: Vec<AccountDid> = owner.into_iter().chain(collaborators).collect();
808 if let Resolved::Ready(Some(cached)) = state.index.owner_of_key(key)
809 && candidates.contains(&cached)
810 {
811 return Some(cached);
812 }
813 let _permit = state.slots.resolve.acquire().await;
814 let matches = futures::stream::iter(candidates).filter_map(|did| async move {
815 let keys = state.atproto.resolve_pubkeys(&did).await.ok()?;
816 keys.iter()
817 .for_each(|resolved| state.index.cache_key(resolved.clone(), &did));
818 keys.iter().any(|resolved| resolved == key).then_some(did)
819 });
820 futures::pin_mut!(matches);
821 matches.next().await
822}
823
824async fn read_chunk<R: AsyncRead + Unpin>(
825 reader: &mut R,
826 buf: &mut Vec<u8>,
827 limit: usize,
828) -> Result<bool, ReadError> {
829 let mut chunk = [0u8; READ_CHUNK];
830 let read = reader.read(&mut chunk).await.map_err(ReadError::Io)?;
831 if read == 0 {
832 return Ok(false);
833 }
834 buf.extend_from_slice(&chunk[..read]);
835 if buf.len() > limit {
836 return Err(ReadError::TooLarge);
837 }
838 Ok(true)
839}
840
841async fn read_receive<R: AsyncRead + Unpin>(
842 reader: &mut R,
843 dir: PathBuf,
844 limit: knot_pack::MaxWireBytes,
845 limits: PackLimits,
846 format: ObjectFormat,
847) -> Result<knot_pack::ReceivedPack, ReadError> {
848 let (tx, rx) = mpsc::channel::<Vec<u8>>(8);
849 let mut framer =
850 tokio::task::spawn_blocking(move || frame_receive(rx, dir, limit, limits, format));
851 let mut chunk = [0u8; READ_CHUNK];
852 let mut io_error = None;
853 loop {
854 tokio::select! {
855 biased;
856 framed = &mut framer => return join_framed(framed, io_error),
857 read = reader.read(&mut chunk) => match read {
858 Ok(0) => break,
859 Ok(read) => {
860 if tx.send(chunk[..read].to_vec()).await.is_err() {
861 break;
862 }
863 }
864 Err(error) => {
865 io_error = Some(error);
866 break;
867 }
868 },
869 }
870 }
871 drop(tx);
872 join_framed(framer.await, io_error)
873}
874
875fn join_framed(
876 framed: Result<Result<knot_pack::ReceivedPack, ReadError>, tokio::task::JoinError>,
877 io_error: Option<std::io::Error>,
878) -> Result<knot_pack::ReceivedPack, ReadError> {
879 match framed {
880 Ok(Ok(body)) => Ok(body),
881 Ok(Err(ReadError::Truncated)) => {
882 Err(io_error.map(ReadError::Io).unwrap_or(ReadError::Truncated))
883 }
884 Ok(Err(other)) => Err(other),
885 Err(_) => Err(ReadError::Truncated),
886 }
887}
888
889fn read_error(error: knot_pack::ReceiveReadError) -> ReadError {
890 match error {
891 knot_pack::ReceiveReadError::Io(error) => ReadError::Io(error),
892 knot_pack::ReceiveReadError::Pack(error) => ReadError::Pack(error),
893 knot_pack::ReceiveReadError::TooLarge => ReadError::TooLarge,
894 knot_pack::ReceiveReadError::Truncated => ReadError::Truncated,
895 }
896}
897
898fn frame_receive(
899 mut rx: mpsc::Receiver<Vec<u8>>,
900 dir: PathBuf,
901 limit: knot_pack::MaxWireBytes,
902 limits: PackLimits,
903 format: ObjectFormat,
904) -> Result<knot_pack::ReceivedPack, ReadError> {
905 let mut receiver =
906 knot_pack::PackReceiver::new(&dir, limit, limits, format.kind()).map_err(ReadError::Io)?;
907 loop {
908 match rx.blocking_recv() {
909 Some(chunk) => {
910 if receiver.write(&chunk).map_err(read_error)? {
911 return receiver.finish().map_err(read_error);
912 }
913 }
914 None => return receiver.finish().map_err(read_error),
915 }
916 }
917}
918
919async fn fail(channel: Channel<Msg>, message: &str) {
920 let _ = channel
921 .extended_data_bytes(1, format!("{message}\n").into_bytes())
922 .await;
923 finish(channel, 1).await;
924}
925
926async fn finish(channel: Channel<Msg>, status: u32) {
927 let _ = channel.exit_status(status).await;
928 let _ = channel.eof().await;
929 let _ = channel.close().await;
930}
931
932#[cfg(test)]
933mod tests {
934 use super::*;
935
936 #[test]
937 fn the_lfs_progress_budget_spares_slow_links_and_cuts_trickles() {
938 assert!(lfs_within_progress_budget(Duration::from_secs(59), 0));
939 assert!(!lfs_within_progress_budget(Duration::from_secs(61), 0));
940 assert!(lfs_within_progress_budget(
941 Duration::from_secs(50_000),
942 5 * 1024 * 1024 * 1024
943 ));
944 assert!(!lfs_within_progress_budget(Duration::from_secs(1_000), 10));
945 }
946
947 #[test]
948 fn the_repo_path_parser_separates_dids_from_handles() {
949 assert!(matches!(
950 parse_repo_path("did:plc:nel/squid"),
951 Some(RepoRef::OwnerPath(..))
952 ));
953 assert!(matches!(
954 parse_repo_path("nel.pet/squid"),
955 Some(RepoRef::HandlePath(..))
956 ));
957 assert!(matches!(
958 parse_repo_path("did:plc:barnacle"),
959 Some(RepoRef::Did(_))
960 ));
961 assert!(parse_repo_path("did:nonsense/squid").is_none());
962 assert!(parse_repo_path("nel.pet").is_none());
963 }
964}