use std::os::fd::{AsRawFd, RawFd};
use std::os::unix::fs::PermissionsExt;
use std::path::Path;
use std::time::Duration;
use tokio::io::{AsyncWriteExt, copy_bidirectional};
use tokio::net::{UnixListener, UnixStream};
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use tokio_vsock::{VMADDR_CID_ANY, VsockAddr, VsockListener, VsockStream};
use arcbox_constants::paths::HOST_SERVICES_SSH_AUTH_SOCK;
use arcbox_constants::ports::SSH_AUTH_RELAY_PORT;
const PAIR_MARKER: u8 = 0x01;
const PAIR_TIMEOUT: Duration = Duration::from_secs(10);
pub async fn run_ssh_auth_relay(cancel: CancellationToken) {
let vsock_listener = match VsockListener::bind(VsockAddr::new(
VMADDR_CID_ANY,
SSH_AUTH_RELAY_PORT,
)) {
Ok(l) => l,
Err(e) => {
tracing::error!(port = SSH_AUTH_RELAY_PORT, error = %e, "failed to bind ssh-auth vsock relay");
return;
}
};
let unix_listener = match bind_unix_socket() {
Ok(l) => l,
Err(e) => {
tracing::error!(path = HOST_SERVICES_SSH_AUTH_SOCK, error = %e, "failed to bind ssh-auth unix socket");
return;
}
};
tracing::info!(
vsock_port = SSH_AUTH_RELAY_PORT,
path = HOST_SERVICES_SSH_AUTH_SOCK,
"ssh-auth relay listening"
);
let (slot_tx, mut slot_rx) = mpsc::unbounded_channel::<VsockStream>();
let acceptor = tokio::spawn(accept_parked_slots(vsock_listener, slot_tx, cancel.clone()));
loop {
let (unix_stream, _) = tokio::select! {
biased;
() = cancel.cancelled() => break,
result = unix_listener.accept() => match result {
Ok(pair) => pair,
Err(e) => {
tracing::warn!(error = %e, "ssh-auth unix accept failed");
continue;
}
}
};
match acquire_slot(&mut slot_rx, &cancel).await {
Some(slot) => {
tokio::spawn(pump(unix_stream, slot));
}
None => {
tracing::warn!("ssh-auth: no relay slot available; is the ArcBox daemon running?");
drop(unix_stream);
}
}
}
let _ = acceptor.await;
}
async fn accept_parked_slots(
mut listener: VsockListener,
slot_tx: mpsc::UnboundedSender<VsockStream>,
cancel: CancellationToken,
) {
loop {
let stream = tokio::select! {
biased;
() = cancel.cancelled() => return,
result = listener.accept() => match result {
Ok((stream, _)) => stream,
Err(e) => {
tracing::warn!(error = %e, "ssh-auth vsock accept failed");
continue;
}
}
};
if slot_tx.send(stream).is_err() {
return;
}
}
}
async fn acquire_slot(
slot_rx: &mut mpsc::UnboundedReceiver<VsockStream>,
cancel: &CancellationToken,
) -> Option<VsockStream> {
let deadline = tokio::time::Instant::now() + PAIR_TIMEOUT;
loop {
let slot = tokio::select! {
biased;
() = cancel.cancelled() => return None,
() = tokio::time::sleep_until(deadline) => return None,
slot = slot_rx.recv() => slot?,
};
if slot_is_parked(&slot) {
return Some(slot);
}
tracing::debug!("ssh-auth: discarded a stale parked slot");
}
}
async fn pump(mut unix: UnixStream, mut slot: VsockStream) {
if let Err(e) = slot.write_all(&[PAIR_MARKER]).await {
tracing::debug!(error = %e, "ssh-auth: failed to signal a parked slot");
return;
}
if let Err(e) = copy_bidirectional(&mut unix, &mut slot).await {
tracing::debug!(error = %e, "ssh-auth: relay copy ended with error");
}
}
fn slot_is_parked(slot: &VsockStream) -> bool {
fd_has_no_pending_input(slot.as_raw_fd())
}
fn fd_has_no_pending_input(fd: RawFd) -> bool {
let mut byte = [0u8; 1];
let n = unsafe {
libc::recv(
fd,
byte.as_mut_ptr().cast::<libc::c_void>(),
1,
libc::MSG_PEEK | libc::MSG_DONTWAIT,
)
};
if n < 0 {
std::io::Error::last_os_error().raw_os_error() == Some(libc::EAGAIN)
} else {
false
}
}
fn bind_unix_socket() -> std::io::Result<UnixListener> {
let path = Path::new(HOST_SERVICES_SSH_AUTH_SOCK);
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
match std::fs::remove_file(path) {
Ok(()) => {}
Err(e) if e.kind() == std::io::ErrorKind::NotFound => {}
Err(e) => return Err(e),
}
let listener = UnixListener::bind(path)?;
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o666))?;
Ok(listener)
}
#[cfg(test)]
mod tests {
use super::fd_has_no_pending_input;
use std::io::Write;
use std::os::fd::AsRawFd;
use std::os::unix::net::UnixStream;
#[test]
fn parked_slot_with_no_pending_data_reads_as_live() {
let (a, _b) = UnixStream::pair().expect("socketpair");
assert!(fd_has_no_pending_input(a.as_raw_fd()));
}
#[test]
fn slot_with_pending_data_is_not_a_fresh_slot() {
let (a, mut b) = UnixStream::pair().expect("socketpair");
b.write_all(b"x").expect("write");
assert!(!fd_has_no_pending_input(a.as_raw_fd()));
}
#[test]
fn closed_peer_reads_as_dead() {
let (a, b) = UnixStream::pair().expect("socketpair");
drop(b);
assert!(!fd_has_no_pending_input(a.as_raw_fd()));
}
}