This repository has no description
0

Configure Feed

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

core / knot2 / crates / knot-ssh / src / exec.rs
32 kB 963 lines
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 let now = state.atproto.now().seconds(); 809 if let Resolved::Ready(Some(cached)) = state.index.owner_of_key(key, now) 810 && candidates.contains(&cached) 811 { 812 return Some(cached); 813 } 814 let _permit = state.slots.resolve.acquire().await; 815 let matches = futures::stream::iter(candidates).filter_map(|did| async move { 816 let keys = state.atproto.resolve_pubkeys(&did).await.ok()?; 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}