use std::future::Future;
use std::net::SocketAddr;
use std::sync::{Arc, PoisonError, RwLock};
use anyhow::{Context, Result};
use axum::extract::DefaultBodyLimit;
use axum::middleware::from_fn_with_state;
use axum::routing::get;
use axum::Router;
use recall_wire::MergeError;
use tokio::net::TcpListener;
use tokio::task::JoinHandle;
use crate::merge::{Merger, Status};
use crate::{now, Config, Store};
mod handlers;
mod limit;
mod middleware;
mod respond;
use handlers::{
handle_admin_page, handle_admin_stats, handle_health, handle_pull, handle_push, not_found,
};
use limit::RateLimiter;
use middleware::guard;
const MAX_BODY_BYTES: usize = 5 << 20;
struct Runtime {
last_backup_at: String,
last_merge_at: String,
last_merge_error: Option<MergeError>,
claude_status: Status,
}
struct AppState {
cfg: Config,
store: Arc<Store>,
merger: Merger,
started_at: String,
runtime: RwLock<Runtime>,
limiter: RateLimiter,
}
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)
}
}
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);
Self {
state: Arc::new(AppState {
cfg,
store,
merger,
started_at: now(),
runtime: RwLock::new(Runtime {
last_backup_at: String::new(),
last_merge_at: String::new(),
last_merge_error: None,
claude_status: Status::default(),
}),
limiter,
}),
}
}
pub fn router(&self) -> Router {
let state = self.state.clone();
Router::new()
.route(
"/sync",
get(handle_pull).post(handle_push).fallback(not_found),
)
.route("/admin/stats", get(handle_admin_stats).fallback(not_found))
.route_layer(from_fn_with_state(state.clone(), guard))
.route("/health", get(handle_health).fallback(not_found))
.route("/admin", get(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 run_backup(&self) {
run_backup(&self.state);
}
pub fn start_background(&self) -> Vec<JoinHandle<()>> {
let mut tasks = Vec::new();
if self.state.cfg.merge_enabled {
let state = self.state.clone();
tasks.push(tokio::spawn(async move {
let every = state.cfg.claude_status_interval;
loop {
let status = state.merger.check_status().await;
state.write().claude_status = status;
tokio::time::sleep(every).await;
}
}));
}
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))?;
eprintln!(
"recall server listening on {} (db: {})",
self.state.cfg.addr, self.state.cfg.db_path
);
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 tasks = self.start_background();
let result = axum::serve(
listener,
self.router()
.into_make_service_with_connect_info::<SocketAddr>(),
)
.with_graceful_shutdown(shutdown)
.await;
for task in tasks {
task.abort();
}
result.map_err(Into::into)
}
}
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 => {}
}
}