use std::future::Future;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicBool, AtomicI64, Ordering};
use std::sync::{Arc, PoisonError, RwLock};
use std::time::{Instant, SystemTime, UNIX_EPOCH};
use anyhow::{Context, Result};
use axum::extract::DefaultBodyLimit;
use axum::middleware::{from_fn, from_fn_with_state};
use axum::routing::{get, post};
use axum::Router;
use recall_wire::devices as paths;
use recall_wire::{ClaudeCliStatus, MergeError};
use tokio::net::TcpListener;
use tokio::sync::Notify;
use tokio::task::JoinHandle;
use crate::config::TlsMode;
use crate::merge::{Merger, Status};
use crate::{format_timestamp, now, Config, Store};
#[cfg_attr(not(feature = "passkeys"), allow(dead_code))]
mod admin;
mod audit;
mod auth;
mod devices;
mod evaluations;
mod handlers;
mod jobs;
mod limit;
mod middleware;
#[cfg(feature = "passkeys")]
mod passkeys;
mod respond;
mod tls;
use audit::{handle_checkpoint, handle_consistency, handle_entries};
#[cfg(feature = "passkeys")]
use passkeys::Passkeys;
#[cfg(not(feature = "passkeys"))]
use without_passkeys::Passkeys;
#[cfg(not(feature = "passkeys"))]
mod without_passkeys {
pub(super) struct Passkeys;
impl Passkeys {
pub(super) fn new(_public_url: &str) -> Self {
Self
}
pub(super) fn status(&self) -> super::admin::PasskeyStatus {
super::admin::PasskeyStatus {
enabled: false,
origin: None,
reason: Some("this recall-server was built without passkey support".to_string()),
}
}
pub(super) fn prune(&self, _now: i64) -> usize {
0
}
}
}
use auth::ReplayCache;
use devices::{
handle_approve, handle_create_authkey, handle_deny, handle_enroll, handle_list_authkeys,
handle_list_devices, handle_me, handle_pending, handle_poll, handle_revoke_authkey,
handle_revoke_device,
};
use handlers::{
handle_admin_stats, handle_discovery, handle_health, handle_pull, handle_push, not_found,
};
use limit::RateLimiter;
use middleware::{admin_guard, admin_only, guard, limited, limited_sign_in, not_worker};
const SWEEP_EVERY: std::time::Duration = std::time::Duration::from_secs(10 * 60);
const MAX_BODY_BYTES: usize = 5 << 20;
const ENROLL_BODY_BYTES: usize = 8 << 10;
const ADMIN_BODY_BYTES: usize = 8 << 10;
const SIGN_IN_BODY_BYTES: usize = 64 << 10;
struct Runtime {
last_backup_at: String,
last_merge_at: String,
last_merge_error: Option<MergeError>,
claude_status: Status,
worker_last_claim_at: Option<String>,
worker_last_claim: Instant,
worker_agent: String,
worker_cli: Option<ClaudeCliStatus>,
}
struct AppState {
cfg: Config,
store: Arc<Store>,
merger: Merger,
started_at: String,
started_unix: AtomicI64,
clock_offset: AtomicI64,
runtime: RwLock<Runtime>,
limiter: RateLimiter,
replay: ReplayCache,
passkeys: Passkeys,
jobs_ready: Notify,
closing: AtomicBool,
draining: AtomicBool,
}
impl AppState {
fn read(&self) -> std::sync::RwLockReadGuard<'_, Runtime> {
self.runtime.read().unwrap_or_else(PoisonError::into_inner)
}
fn write(&self) -> std::sync::RwLockWriteGuard<'_, Runtime> {
self.runtime.write().unwrap_or_else(PoisonError::into_inner)
}
fn now(&self) -> i64 {
unix_now() + self.clock_offset.load(Ordering::Relaxed)
}
fn clock(&self) -> time::OffsetDateTime {
time::OffsetDateTime::now_utc()
+ time::Duration::seconds(self.clock_offset.load(Ordering::Relaxed))
}
fn started(&self) -> i64 {
self.started_unix.load(Ordering::Relaxed)
}
}
fn unix_now() -> i64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.unwrap_or(0)
}
fn nonces_per_device(cfg: &Config) -> usize {
let live_ms = auth::NONCE_LIFETIME as u128 * 1000;
let window_ms = cfg.rate_limit_window.as_millis().max(1);
let windows = live_ms.div_ceil(window_ms);
(cfg.rate_limit_max as u128 * windows).min(usize::MAX as u128) as usize
}
pub struct Server {
state: Arc<AppState>,
}
impl Server {
pub fn new(cfg: Config, store: Arc<Store>) -> Self {
let limiter = RateLimiter::new(cfg.rate_limit_window, cfg.rate_limit_max);
let merger = Merger::new(cfg.claude_bin.clone(), cfg.merge_timeout);
let replay = ReplayCache::new(auth::WINDOW, nonces_per_device(&cfg));
let passkeys = Passkeys::new(&cfg.public_url);
Self {
state: Arc::new(AppState {
cfg,
store,
merger,
started_at: now(),
started_unix: AtomicI64::new(unix_now()),
clock_offset: AtomicI64::new(0),
runtime: RwLock::new(Runtime {
last_backup_at: String::new(),
last_merge_at: String::new(),
last_merge_error: None,
claude_status: Status::default(),
worker_last_claim_at: None,
worker_last_claim: Instant::now(),
worker_agent: String::new(),
worker_cli: None,
}),
limiter,
replay,
passkeys,
jobs_ready: Notify::new(),
closing: AtomicBool::new(false),
draining: AtomicBool::new(false),
}),
}
}
pub fn router(&self) -> Router {
let state = self.state.clone();
let admin = Router::new()
.route("/admin/stats", get(handle_admin_stats).fallback(not_found))
.route(
paths::DEVICES_PATH,
get(handle_list_devices).fallback(not_found),
)
.route(
paths::APPROVE_PATH,
post(handle_approve).fallback(not_found),
)
.route(paths::DENY_PATH, post(handle_deny).fallback(not_found))
.route(
"/v1/devices/pending/{user_code}",
get(handle_pending).fallback(not_found),
)
.route(
"/v1/devices/{id}/revoke",
post(handle_revoke_device).fallback(not_found),
)
.route(
paths::AUTHKEYS_PATH,
get(handle_list_authkeys)
.post(handle_create_authkey)
.fallback(not_found),
)
.route(
"/v1/authkeys/{id}/revoke",
post(handle_revoke_authkey).fallback(not_found),
)
.route(
recall_wire::audit::ENTRIES_PATH,
get(handle_entries).fallback(not_found),
)
.route(
recall_wire::evaluations::EVALUATIONS_PATH,
get(evaluations::handle_list)
.post(evaluations::handle_request)
.fallback(not_found),
)
.route(
"/v1/evaluations/{id}",
get(evaluations::handle_get).fallback(not_found),
)
.route_layer(DefaultBodyLimit::max(ADMIN_BODY_BYTES))
.route_layer(from_fn(admin_only))
.route_layer(from_fn_with_state(state.clone(), admin_guard));
let audit_routes = Router::new()
.route(
recall_wire::audit::CHECKPOINT_PATH,
get(handle_checkpoint).fallback(not_found),
)
.route(
recall_wire::audit::CONSISTENCY_PATH,
get(handle_consistency).fallback(not_found),
)
.route_layer(from_fn_with_state(state.clone(), guard));
let enrolment = Router::new()
.route(paths::ENROLL_PATH, post(handle_enroll).fallback(not_found))
.route(
paths::ENROLL_POLL_PATH,
post(handle_poll).fallback(not_found),
)
.route_layer(DefaultBodyLimit::max(ENROLL_BODY_BYTES))
.route_layer(from_fn_with_state(state.clone(), limited));
let page = Router::new()
.route(
"/admin/session",
get(admin::handle_session_status).fallback(not_found),
)
.route_layer(from_fn_with_state(state.clone(), limited_sign_in));
Router::new()
.route(
"/sync",
get(handle_pull).post(handle_push).fallback(not_found),
)
.route(paths::DEVICES_ME_PATH, get(handle_me).fallback(not_found))
.route_layer(from_fn(not_worker))
.route_layer(from_fn_with_state(state.clone(), guard))
.merge(admin)
.merge(audit_routes)
.merge(enrolment)
.merge(page)
.merge(sign_in_routes(&state))
.merge(jobs::routes(state.clone()))
.route("/health", get(handle_health).fallback(not_found))
.route(
recall_wire::DISCOVERY_PATH,
get(handle_discovery).fallback(not_found),
)
.route("/admin", get(admin::handle_admin_page).fallback(not_found))
.fallback(not_found)
.layer(DefaultBodyLimit::max(MAX_BODY_BYTES))
.with_state(state)
}
pub async fn refresh_claude_status(&self) {
let status = self.state.merger.check_status().await;
self.state.write().claude_status = status;
}
pub fn claude_status(&self) -> Status {
self.state.read().claude_status.clone()
}
pub fn set_claude_status(&self, status: Status) {
self.state.write().claude_status = status;
}
pub fn set_clock_offset(&self, seconds: i64) {
self.state.clock_offset.store(seconds, Ordering::Relaxed);
}
pub fn backdate_start(&self, seconds: i64) {
self.state
.started_unix
.fetch_sub(seconds, Ordering::Relaxed);
}
pub fn backdate_last_claim(&self, seconds: u64) {
let mut rt = self.state.write();
if let Some(earlier) = rt
.worker_last_claim
.checked_sub(std::time::Duration::from_secs(seconds))
{
rt.worker_last_claim = earlier;
}
}
pub async fn drain_jobs(&self) -> Result<()> {
jobs::drain_without_worker(&self.state).await
}
pub fn run_scheduled_evaluation(&self) -> Result<Option<String>> {
evaluations::run_scheduled(&self.state)
}
pub fn run_backup(&self) {
run_backup(&self.state);
}
pub fn issue_bootstrap_code(&self) -> Result<Option<crate::bootstrap::BootstrapCode>> {
if !self.state.passkeys.status().enabled || self.state.store.has_admin_credentials()? {
return Ok(None);
}
crate::bootstrap::issue(&self.state.store, self.state.clock()).map(Some)
}
pub fn public_url(&self) -> Option<&str> {
Some(self.state.cfg.public_url.as_str()).filter(|u| !u.is_empty())
}
pub fn sweep_devices(&self) -> Result<(usize, usize)> {
sweep_devices(&self.state)
}
pub fn start_background(&self) -> Vec<JoinHandle<()>> {
let mut tasks = Vec::new();
{
let state = self.state.clone();
tasks.push(tokio::spawn(async move {
let mut wal = WalWatch::default();
loop {
let s = state.clone();
match tokio::task::spawn_blocking(move || sweep_devices(&s)).await {
Ok(Ok((0, 0))) => {}
Ok(Ok((devices, enrollments))) => eprintln!(
"removed {devices} idle ephemeral devices and {enrollments} expired enrolments"
),
Ok(Err(e)) => eprintln!("device sweep failed: {e:#}"),
Err(_) => {}
}
let s = state.clone();
match tokio::task::spawn_blocking(move || jobs::prune(&s)).await {
Ok(Ok(0)) | Err(_) => {}
Ok(Ok(n)) => eprintln!("removed {n} finished merge jobs"),
Ok(Err(e)) => eprintln!("job prune failed: {e:#}"),
}
if let Err(e) = jobs::drain_without_worker(&state).await {
eprintln!("draining the merge queue failed: {e:#}");
}
let s = state.clone();
match tokio::task::spawn_blocking(move || s.store.checkpoint()).await {
Ok(Ok(complete)) => {
if let Some(said) = wal.record(complete) {
eprintln!("{said}");
}
}
Ok(Err(e)) => eprintln!("checkpointing the WAL failed: {e:#}"),
Err(_) => {}
}
tokio::time::sleep(SWEEP_EVERY).await;
}
}));
}
if self.state.cfg.merge_enabled {
let state = self.state.clone();
tasks.push(tokio::spawn(async move {
let every = state.cfg.claude_status_interval;
let mut first = true;
loop {
let status = state.merger.check_status().await;
state.write().claude_status = status;
if std::mem::take(&mut first) {
if let Err(e) = jobs::drain_without_worker(&state).await {
eprintln!("draining the merge queue failed: {e:#}");
}
}
tokio::time::sleep(every).await;
}
}));
}
if let Some(every) = self.state.cfg.eval_interval {
let state = self.state.clone();
tasks.push(tokio::spawn(async move {
loop {
tokio::time::sleep(every).await;
let s = state.clone();
match tokio::task::spawn_blocking(move || evaluations::run_scheduled(&s)).await
{
Ok(Ok(Some(id))) => eprintln!("queued the scheduled evaluation {id}"),
Ok(Ok(None)) => eprintln!(
"the scheduled evaluation was skipped: no worker is enrolled, or \
another evaluation is still open"
),
Ok(Err(e)) => eprintln!("queueing the scheduled evaluation failed: {e:#}"),
Err(_) => {}
}
}
}));
}
if !self.state.cfg.backup_dir.is_empty() {
let state = self.state.clone();
tasks.push(tokio::spawn(async move {
let every = state.cfg.backup_interval;
loop {
let s = state.clone();
let _ = tokio::task::spawn_blocking(move || run_backup(&s)).await;
tokio::time::sleep(every).await;
}
}));
}
tasks
}
pub async fn serve(&self) -> Result<()> {
let listener = TcpListener::bind(&self.state.cfg.addr)
.await
.with_context(|| format!("binding {}", self.state.cfg.addr))?;
self.serve_with_shutdown(listener, shutdown_signal()).await
}
pub async fn serve_with_shutdown<F>(&self, listener: TcpListener, shutdown: F) -> Result<()>
where
F: Future<Output = ()> + Send + 'static,
{
let transport = match &self.state.cfg.tls {
TlsMode::Off => None,
mode => Some(tls::prepare(mode).await?),
};
eprintln!(
"recall server listening on {} ({}, db: {})",
listener
.local_addr()
.map_or_else(|_| self.state.cfg.addr.clone(), |a| a.to_string()),
transport
.as_ref()
.map_or("plain http", tls::Prepared::description),
self.state.cfg.db_path
);
if let Err(e) = record_start(&self.state.store) {
eprintln!("recording server start in the audit log: {e:#}");
}
let tasks = self.start_background();
let state = self.state.clone();
let shutdown = async move {
shutdown.await;
state.closing.store(true, Ordering::Relaxed);
state.jobs_ready.notify_waiters();
};
let result = match transport {
None => axum::serve(
listener,
self.router()
.into_make_service_with_connect_info::<SocketAddr>(),
)
.with_graceful_shutdown(shutdown)
.await
.map_err(Into::into),
Some(prepared) => {
let listener = listener.into_std().context("preparing the TLS listener")?;
let limits = tls::Limits::from_config(&self.state.cfg);
tls::serve(self.router(), listener, prepared, limits, shutdown).await
}
};
for task in tasks {
task.abort();
}
match self.state.store.checkpoint_all() {
Ok(true) => {}
Ok(false) => eprintln!(
"stopping with commits still in the WAL: a reader held it past the busy \
timeout. Nothing is lost; the next start reads them"
),
Err(e) => eprintln!("checkpointing the WAL at shutdown failed: {e:#}"),
}
result
}
}
fn record_start(store: &Store) -> Result<()> {
let version = recall_wire::discovery::version();
store.audit_append(|seq, at| {
crate::audit::leaf::encode(
seq,
at,
crate::audit::leaf::action::START,
&crate::audit::leaf::Actor::Server,
crate::audit::leaf::subject_start(&version),
None,
)
})?;
Ok(())
}
#[cfg(feature = "passkeys")]
fn sign_in_routes(state: &Arc<AppState>) -> Router<Arc<AppState>> {
use passkeys::{
handle_add_finish, handle_add_start, handle_bootstrap_finish, handle_bootstrap_start,
handle_list_passkeys, handle_remove, handle_sign_in_finish, handle_sign_in_start,
handle_sign_out, handle_sign_out_others, json_only,
};
let sign_in = Router::new()
.route(
"/admin/login/start",
post(handle_sign_in_start).fallback(not_found),
)
.route(
"/admin/login/finish",
post(handle_sign_in_finish).fallback(not_found),
)
.route_layer(from_fn(json_only))
.route_layer(DefaultBodyLimit::max(SIGN_IN_BODY_BYTES))
.route_layer(from_fn_with_state(state.clone(), limited_sign_in));
let bootstrap = Router::new()
.route(
"/admin/bootstrap/register",
post(handle_bootstrap_start).fallback(not_found),
)
.route(
"/admin/bootstrap/register/finish",
post(handle_bootstrap_finish).fallback(not_found),
)
.route_layer(from_fn(json_only))
.route_layer(DefaultBodyLimit::max(SIGN_IN_BODY_BYTES))
.route_layer(from_fn_with_state(state.clone(), guard));
let adding = Router::new()
.route(
"/admin/passkeys/register",
post(handle_add_start).fallback(not_found),
)
.route(
"/admin/passkeys/register/finish",
post(handle_add_finish).fallback(not_found),
)
.route_layer(from_fn(json_only))
.route_layer(DefaultBodyLimit::max(SIGN_IN_BODY_BYTES))
.route_layer(from_fn_with_state(state.clone(), admin::owner_only));
let owner = Router::new()
.route(
"/admin/passkeys",
get(handle_list_passkeys).fallback(not_found),
)
.route(
"/admin/passkeys/{id}/remove",
post(handle_remove).fallback(not_found),
)
.route("/admin/logout", post(handle_sign_out).fallback(not_found))
.route(
"/admin/logout/others",
post(handle_sign_out_others).fallback(not_found),
)
.route_layer(DefaultBodyLimit::max(SIGN_IN_BODY_BYTES))
.route_layer(from_fn_with_state(state.clone(), admin::owner_only));
sign_in.merge(bootstrap).merge(adding).merge(owner)
}
#[cfg(not(feature = "passkeys"))]
fn sign_in_routes(_state: &Arc<AppState>) -> Router<Arc<AppState>> {
Router::new()
}
fn sweep_devices(state: &AppState) -> Result<(usize, usize)> {
let now = time::OffsetDateTime::now_utc();
state.passkeys.prune(state.now());
let clock = state.clock();
if let Err(e) = state.store.sweep_admin_sessions(
&format_timestamp(clock),
&format_timestamp(clock - admin::SESSION_IDLE),
) {
eprintln!("admin session sweep failed: {e:#}");
}
state.store.sweep_devices_audited(
&format_timestamp(now - state.cfg.ephemeral_device_ttl),
&format_timestamp(now - devices::EXPIRED_ENROLLMENT_KEPT),
)
}
#[derive(Debug, Default)]
struct WalWatch {
behind: u32,
}
impl WalWatch {
const SWEEPS: u32 = 3;
fn record(&mut self, complete: bool) -> Option<String> {
if complete {
let was = std::mem::take(&mut self.behind);
return (was >= Self::SWEEPS)
.then(|| "the WAL is fully checkpointed into recall.db again".to_string());
}
self.behind += 1;
self.behind.is_multiple_of(Self::SWEEPS).then(|| {
format!(
"the WAL has not been fully checkpointed into recall.db for {} sweeps in a \
row: a reader is keeping a transaction open (sqlite-web on a page, a \
`sqlite3` shell), and recall.db-wal grows until it ends. Nothing is lost",
self.behind
)
})
}
}
fn run_backup(state: &AppState) {
if state.cfg.backup_dir.is_empty() {
return;
}
match state
.store
.backup(&state.cfg.backup_dir, state.cfg.backup_keep)
{
Ok(dest) => {
state.write().last_backup_at = now();
eprintln!("backup written: {}", dest.display());
}
Err(e) => eprintln!("backup failed: {e:#}"),
}
}
async fn shutdown_signal() {
let ctrl_c = async {
let _ = tokio::signal::ctrl_c().await;
};
#[cfg(unix)]
let terminate = async {
match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) {
Ok(mut sig) => {
sig.recv().await;
}
Err(_) => std::future::pending::<()>().await,
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {}
_ = terminate => {}
}
}
#[cfg(test)]
mod tests {
use super::{Server, WalWatch};
use crate::{Config, Store};
use std::sync::Arc;
use std::time::{Duration, Instant};
#[tokio::test]
async fn the_sweep_checkpoints_the_wal_into_the_file() {
let dir = tempfile::tempdir().unwrap();
let db = dir.path().join("recall.db");
let store = Arc::new(Store::open(&db).unwrap());
let server = Server::new(
Config {
token: "sweep-token".into(),
merge_enabled: false,
..Config::default()
},
store.clone(),
);
assert!(store.checkpoint_all().unwrap());
for i in 0..5 {
store
.upsert_audited(
"acme/app",
&format!("f{i}.md"),
"x",
"",
crate::store::test_leaf,
)
.unwrap();
}
let alone = || -> i64 {
let copy = tempfile::tempdir().unwrap();
let file = copy.path().join("recall.db");
store
.with_raw(|_| {
std::fs::copy(&db, &file).unwrap();
Ok(())
})
.unwrap();
rusqlite::Connection::open(&file)
.unwrap()
.query_row("SELECT count(*) FROM memory_files", [], |r| r.get(0))
.unwrap()
};
assert_eq!(alone(), 0, "the rows are only in the WAL before the sweep");
let tasks = server.start_background();
let started = Instant::now();
while alone() != 5 {
assert!(
started.elapsed() < Duration::from_secs(10),
"the first sweep did not checkpoint: the file alone has {}",
alone()
);
tokio::time::sleep(Duration::from_millis(20)).await;
}
for t in tasks {
t.abort();
}
}
#[test]
fn a_wal_held_back_for_several_sweeps_is_reported_and_so_is_its_end() {
let mut w = WalWatch::default();
assert_eq!(w.record(false), None);
assert_eq!(w.record(true), None, "one sweep behind is nothing");
assert_eq!(w.record(false), None);
assert_eq!(w.record(false), None);
let said = w.record(false).expect("three in a row");
assert!(said.contains("for 3 sweeps"), "{said}");
assert_eq!(w.record(false), None);
assert_eq!(w.record(false), None);
assert!(w.record(false).unwrap().contains("for 6 sweeps"));
assert!(w.record(true).unwrap().contains("again"));
assert_eq!(w.record(true), None);
}
}