use std::collections::HashSet;
use std::path::PathBuf;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use crate::transport_iroh::ratelimit::FailureLimiter;
use anyhow::Context;
use clap::Args as ClapArgs;
use iroh::EndpointId;
use tokio::signal::unix::{signal, SignalKind};
use tokio_util::sync::CancellationToken;
use crate::server::{run_attached, session, SessionExit};
use crate::transport_iroh::{
bind_endpoint, bind_endpoint_local, bind_endpoint_with_relay, format_endpoint_id,
load_or_create_secret_key, parse_endpoint_id, parse_relay_url, MonoClock, ALPN,
};
use secrecy::{ExposeSecret, SecretString};
use tracing::{error, info, warn};
const AUTH_FAIL_WINDOW_MS: u64 = 60_000;
const AUTH_MAX_FAILURES: usize = 5;
#[derive(ClapArgs, Debug)]
pub struct ServeArgs {
#[arg(long)]
key_file: Option<PathBuf>,
#[arg(long = "allow", value_name = "ENDPOINT_ID")]
allow: Vec<String>,
#[arg(long)]
allow_any: bool,
#[arg(long)]
shell: Option<String>,
#[arg(long, default_value_t = 1000)]
scrollback: usize,
#[arg(long, default_value_t = 86_400)]
session_ttl_secs: u64,
#[arg(long, env = "KOH_SERVER_NETWORK_TMOUT", default_value_t = 0)]
network_timeout_secs: u64,
#[arg(long, value_name = "URL")]
relay_url: Option<String>,
#[arg(long, conflicts_with = "relay_url")]
local: bool,
#[arg(long)]
passphrase: Option<String>,
#[arg(long, default_value_t = 64, value_parser = clap::value_parser!(u32).range(1..))]
max_connections: u32,
#[arg(long, default_value_t = 64, value_parser = clap::value_parser!(u32).range(1..))]
max_sessions: u32,
}
fn connect_qr(data: &str) -> Option<String> {
use qrcode::render::unicode::Dense1x2;
let code = qrcode::QrCode::new(data).ok()?;
Some(
code.render::<Dense1x2>()
.dark_color(Dense1x2::Light)
.light_color(Dense1x2::Dark)
.quiet_zone(true)
.build(),
)
}
fn default_key_file() -> PathBuf {
crate::transport_iroh::default_key_path("server")
}
#[expect(
clippy::expect_used,
reason = "a poisoned auth-limiter mutex is a bug, not peer input"
)]
fn lock_limiter(
limiter: &session::AuthLimiter,
) -> std::sync::MutexGuard<'_, FailureLimiter<EndpointId>> {
limiter.lock().expect("auth limiter mutex poisoned")
}
pub async fn serve(args: ServeArgs) -> anyhow::Result<()> {
tracing_subscriber::fmt()
.with_env_filter(
tracing_subscriber::EnvFilter::try_from_default_env()
.unwrap_or_else(|_| "koh_server=info,koh=info".into()),
)
.with_writer(std::io::stderr)
.init();
let mut allow: HashSet<EndpointId> = HashSet::new();
for s in &args.allow {
let id = parse_endpoint_id(s).with_context(|| format!("bad --allow id: {s}"))?;
allow.insert(id);
}
if allow.is_empty() && !args.allow_any {
anyhow::bail!(
"no clients authorized: pass --allow <endpoint-id> (repeatable), or --allow-any for testing"
);
}
let key_file = args.key_file.clone().unwrap_or_else(default_key_file);
let secret = load_or_create_secret_key(&key_file).with_context(|| {
format!(
"loading server key from {} (pass --key-file to use a writable path)",
key_file.display()
)
})?;
let endpoint = if let Some(url) = &args.relay_url {
let relay = parse_relay_url(url)?;
bind_endpoint_with_relay(secret, true, relay)
.await
.context("binding endpoint")?
} else if args.local {
bind_endpoint_local(secret, true)
.await
.context("binding endpoint")?
} else {
bind_endpoint(secret, true)
.await
.context("binding endpoint")?
};
let my_id = endpoint.id();
let id_str = format_endpoint_id(&my_id);
let connect_hint = if let Some(url) = &args.relay_url {
format!("koh connect {id_str} --relay-url {url}")
} else if args.local {
let port = endpoint
.bound_sockets()
.iter()
.find(|s| s.is_ipv4())
.map_or(0, std::net::SocketAddr::port);
format!("koh connect {id_str} --direct <this-host-ip>:{port}")
} else {
format!("koh connect {id_str}")
};
eprintln!("┌─ koh server ready ──────────────────────────────────────");
eprintln!("│ endpoint id : {id_str}");
eprintln!("│ key file : {}", key_file.display());
eprintln!("│ alpn : {}", String::from_utf8_lossy(ALPN));
if args.allow_any {
eprintln!("│ auth : ⚠ ALLOW-ANY (insecure)");
} else {
eprintln!("│ auth : allowlist ({} client(s))", allow.len());
}
if args.passphrase.is_some() || std::env::var("KOH_PASSPHRASE").is_ok() {
eprintln!("│ 2nd factor : passphrase required");
}
eprintln!("│ connect : {connect_hint}");
eprintln!("└───────────────────────────────────────────────────────────");
if let Some(qr) = connect_qr(&id_str) {
eprintln!("\nScan for the endpoint id (point a phone camera at it):\n");
eprintln!("{qr}");
} else {
warn!("could not render the connect QR (endpoint id too large to encode)");
}
let shell = args.shell.clone();
let scrollback = args.scrollback;
let allow = std::sync::Arc::new(allow);
let allow_any = args.allow_any;
let passphrase: std::sync::Arc<Option<SecretString>> = std::sync::Arc::new(
args.passphrase
.clone()
.or_else(|| std::env::var("KOH_PASSPHRASE").ok())
.map(SecretString::from),
);
let clock = MonoClock::new();
let limiter: session::AuthLimiter = std::sync::Arc::new(std::sync::Mutex::new(
FailureLimiter::new(AUTH_FAIL_WINDOW_MS, AUTH_MAX_FAILURES),
));
let store = session::SessionStore::default();
let session_ttl = std::time::Duration::from_secs(args.session_ttl_secs);
let reaper_shutdown = tokio_util::sync::CancellationToken::new();
let reaper = tokio::spawn(session::run_reaper(
store.clone(),
session_ttl,
limiter.clone(),
clock,
session::REAP_INTERVAL,
reaper_shutdown.clone(),
));
let shutdown = CancellationToken::new();
spawn_signal_drain(shutdown.clone())?;
let conn_limit = Arc::new(tokio::sync::Semaphore::new(args.max_connections as usize));
let max_sessions = args.max_sessions as usize;
let active = Arc::new(AtomicUsize::new(0));
if args.network_timeout_secs > 0 {
spawn_idle_watchdog(
active.clone(),
Duration::from_secs(args.network_timeout_secs),
shutdown.clone(),
);
}
loop {
let incoming = tokio::select! {
biased;
() = shutdown.cancelled() => break,
inc = endpoint.accept() => match inc {
Some(i) => i,
None => break, },
};
let Ok(permit) = conn_limit.clone().try_acquire_owned() else {
warn!("refusing connection: at max-connections capacity");
incoming.refuse();
continue;
};
let allow = allow.clone();
let shell = shell.clone();
let passphrase = passphrase.clone();
let store = store.clone();
let limiter = limiter.clone();
let active_guard = ConnGuard::new(active.clone());
tokio::spawn(async move {
let _permit = permit;
let _active_guard = active_guard;
let conn = match incoming.await {
Ok(c) => c,
Err(e) => {
warn!(error = %e, "incoming handshake failed");
return;
}
};
let peer = conn.remote_id();
if !allow_any && !allow.contains(&peer) {
warn!(peer = %format_endpoint_id(&peer), "rejected: not on allowlist");
conn.close(1u32.into(), b"not authorized");
return;
}
if !lock_limiter(&limiter).check(&peer, clock.now_ms()) {
warn!(peer = %format_endpoint_id(&peer), "rejected: too many failed auth attempts");
conn.close(1u32.into(), b"rate limited");
return;
}
match tokio::time::timeout(
std::time::Duration::from_secs(10),
crate::transport_iroh::auth::handshake_server(
&conn,
passphrase
.as_ref()
.as_ref()
.map(ExposeSecret::expose_secret),
),
)
.await
{
Ok(Ok(())) => {
lock_limiter(&limiter).record_success(&peer);
}
Ok(Err(e)) => {
warn!(peer = %format_endpoint_id(&peer), error = %e, "passphrase handshake rejected");
lock_limiter(&limiter).record_failure(peer, clock.now_ms());
conn.close(1u32.into(), b"auth failed");
return;
}
Err(_) => {
warn!(peer = %format_endpoint_id(&peer), "passphrase handshake timed out");
lock_limiter(&limiter).record_failure(peer, clock.now_ms());
conn.close(1u32.into(), b"auth timeout");
return;
}
}
info!(peer = %format_endpoint_id(&peer), "client authorized; attaching session");
let (handle, attach_kind) = match session::attach(
&store,
peer,
shell.as_deref(),
scrollback,
max_sessions,
)
.await
{
Ok(Some(pair)) => pair,
Ok(None) => {
warn!(peer = %format_endpoint_id(&peer), "refusing session: at max-sessions capacity");
conn.close(1u32.into(), b"server at session capacity");
return;
}
Err(e) => {
error!(error = %e, "failed to start session");
conn.close(1u32.into(), b"session error");
return;
}
};
match attach_kind {
session::AttachKind::Created => {
info!(peer = %format_endpoint_id(&peer), "started a new session");
}
session::AttachKind::Reattached { detached_for } => {
info!(
peer = %format_endpoint_id(&peer),
detached_secs = detached_for.map(|d| d.as_secs()),
"reattaching to this peer's existing session"
);
}
}
match run_attached(conn, handle).await {
Ok(SessionExit::Detached) => {
session::detach(&store, peer).await;
info!(peer = %format_endpoint_id(&peer), "client detached (session retained)");
}
Ok(SessionExit::ShellExited) => {
session::reap(&store, peer).await;
info!(peer = %format_endpoint_id(&peer), "shell exited; session reaped");
}
Err(e) => {
error!(error = %e, "session loop error");
session::detach(&store, peer).await;
}
}
});
}
info!("draining: stopping reaper and closing endpoint");
shutdown.cancel();
reaper_shutdown.cancel();
let _ = reaper.await;
endpoint.close().await;
Ok(())
}
struct ConnGuard(Arc<AtomicUsize>);
impl ConnGuard {
fn new(active: Arc<AtomicUsize>) -> Self {
active.fetch_add(1, Ordering::SeqCst);
Self(active)
}
}
impl Drop for ConnGuard {
fn drop(&mut self) {
self.0.fetch_sub(1, Ordering::SeqCst);
}
}
fn spawn_signal_drain(shutdown: CancellationToken) -> anyhow::Result<()> {
let mut term = signal(SignalKind::terminate()).context("installing SIGTERM handler")?;
let mut intr = signal(SignalKind::interrupt()).context("installing SIGINT handler")?;
tokio::spawn(async move {
tokio::select! {
_ = term.recv() => {}
_ = intr.recv() => {}
}
info!("received shutdown signal; draining");
shutdown.cancel();
});
Ok(())
}
fn spawn_idle_watchdog(active: Arc<AtomicUsize>, timeout: Duration, shutdown: CancellationToken) {
tokio::spawn(async move {
let tick = Duration::from_secs(1).min(timeout);
let mut idle_since: Option<Instant> = None;
loop {
tokio::select! {
() = shutdown.cancelled() => return,
() = tokio::time::sleep(tick) => {}
}
if active.load(Ordering::SeqCst) == 0 {
let since = *idle_since.get_or_insert_with(Instant::now);
if since.elapsed() >= timeout {
info!(
timeout_secs = timeout.as_secs(),
"network idle timeout; shutting down"
);
shutdown.cancel();
return;
}
} else {
idle_since = None;
}
}
});
}
#[cfg(test)]
mod tests {
use super::connect_qr;
#[test]
fn connect_qr_renders_an_id_and_handles_overlong_input() {
let id = "3f9c".repeat(16);
let qr = connect_qr(&id).expect("an endpoint id must fit in a QR");
assert!(qr.lines().count() > 5, "a QR should be a multi-row block");
assert!(
qr.contains('█') || qr.contains('▀') || qr.contains('▄'),
"the unicode renderer uses half-block glyphs"
);
assert!(
connect_qr(&"a".repeat(10_000)).is_none(),
"overlong input must return None, not panic"
);
}
}