1use std::collections::HashMap;
13use std::future::Future;
14use std::net::SocketAddr;
15use std::sync::{Arc, Mutex, PoisonError, RwLock};
16use std::time::{Duration, Instant};
17
18use anyhow::{Context, Result};
19use axum::body::Bytes;
20use axum::extract::{ConnectInfo, DefaultBodyLimit, Query, Request, State};
21use axum::http::{HeaderMap, StatusCode};
22use axum::middleware::{self, Next};
23use axum::response::{IntoResponse, Response};
24use axum::routing::get;
25use axum::{Json, Router};
26use recall_wire::{
27 AdminStats, ClaudeCliStatus, ErrorResponse, Health, MergeError, MergeStatus, PushRequest,
28 PushResponse, SyncResponse,
29};
30use serde::Serialize;
31use tokio::net::TcpListener;
32use tokio::task::JoinHandle;
33
34use crate::merge::{Merger, Status};
35use crate::{now, Config, Store};
36
37const ADMIN_HTML: &str = include_str!("../admin.html");
40
41const ADMIN_CSP: &str = "default-src 'none'; style-src 'unsafe-inline'; script-src 'unsafe-inline'; connect-src 'self'; frame-ancestors 'none'";
45
46const MAX_BODY_BYTES: usize = 5 << 20;
49
50const REQUIRED_FIELDS_MSG: &str =
51 "project_key, file_path, and content (string) are required, unless deleted is true";
52
53struct Runtime {
54 last_backup_at: String,
55 last_merge_at: String,
56 last_merge_error: Option<MergeError>,
57 claude_status: Status,
58}
59
60struct AppState {
61 cfg: Config,
62 store: Arc<Store>,
63 merger: Merger,
64 started_at: String,
65 runtime: RwLock<Runtime>,
66 limiter: RateLimiter,
67}
68
69impl AppState {
70 fn read(&self) -> std::sync::RwLockReadGuard<'_, Runtime> {
71 self.runtime.read().unwrap_or_else(PoisonError::into_inner)
72 }
73 fn write(&self) -> std::sync::RwLockWriteGuard<'_, Runtime> {
74 self.runtime.write().unwrap_or_else(PoisonError::into_inner)
75 }
76}
77
78pub struct Server {
80 state: Arc<AppState>,
81}
82
83impl Server {
84 pub fn new(cfg: Config, store: Arc<Store>) -> Self {
86 let limiter = RateLimiter::new(cfg.rate_limit_window, cfg.rate_limit_max);
87 let merger = Merger::new(cfg.claude_bin.clone(), cfg.merge_timeout);
88 Self {
89 state: Arc::new(AppState {
90 cfg,
91 store,
92 merger,
93 started_at: now(),
94 runtime: RwLock::new(Runtime {
95 last_backup_at: String::new(),
96 last_merge_at: String::new(),
97 last_merge_error: None,
98 claude_status: Status::default(),
99 }),
100 limiter,
101 }),
102 }
103 }
104
105 pub fn router(&self) -> Router {
108 let state = self.state.clone();
109 Router::new()
110 .route(
115 "/sync",
116 get(handle_pull).post(handle_push).fallback(not_found),
117 )
118 .route("/admin/stats", get(handle_admin_stats).fallback(not_found))
119 .route_layer(middleware::from_fn_with_state(state.clone(), guard))
122 .route("/health", get(handle_health).fallback(not_found))
123 .route("/admin", get(handle_admin_page).fallback(not_found))
124 .fallback(not_found)
125 .layer(DefaultBodyLimit::max(MAX_BODY_BYTES))
126 .with_state(state)
127 }
128
129 pub async fn refresh_claude_status(&self) {
131 let status = self.state.merger.check_status().await;
132 self.state.write().claude_status = status;
133 }
134
135 pub fn claude_status(&self) -> Status {
137 self.state.read().claude_status.clone()
138 }
139
140 pub fn set_claude_status(&self, status: Status) {
145 self.state.write().claude_status = status;
146 }
147
148 pub fn run_backup(&self) {
151 run_backup(&self.state);
152 }
153
154 pub fn start_background(&self) -> Vec<JoinHandle<()>> {
158 let mut tasks = Vec::new();
159 if self.state.cfg.merge_enabled {
160 let state = self.state.clone();
161 tasks.push(tokio::spawn(async move {
162 let every = state.cfg.claude_status_interval;
163 loop {
164 let status = state.merger.check_status().await;
165 state.write().claude_status = status;
166 tokio::time::sleep(every).await;
167 }
168 }));
169 }
170 if !self.state.cfg.backup_dir.is_empty() {
171 let state = self.state.clone();
172 tasks.push(tokio::spawn(async move {
173 let every = state.cfg.backup_interval;
174 loop {
175 let s = state.clone();
179 let _ = tokio::task::spawn_blocking(move || run_backup(&s)).await;
180 tokio::time::sleep(every).await;
181 }
182 }));
183 }
184 tasks
185 }
186
187 pub async fn serve(&self) -> Result<()> {
190 let listener = TcpListener::bind(&self.state.cfg.addr)
191 .await
192 .with_context(|| format!("binding {}", self.state.cfg.addr))?;
193 eprintln!(
194 "recall server listening on {} (db: {})",
195 self.state.cfg.addr, self.state.cfg.db_path
196 );
197 self.serve_with_shutdown(listener, shutdown_signal()).await
198 }
199
200 pub async fn serve_with_shutdown<F>(&self, listener: TcpListener, shutdown: F) -> Result<()>
202 where
203 F: Future<Output = ()> + Send + 'static,
204 {
205 let tasks = self.start_background();
206 let result = axum::serve(
207 listener,
208 self.router()
209 .into_make_service_with_connect_info::<SocketAddr>(),
210 )
211 .with_graceful_shutdown(shutdown)
212 .await;
213 for task in tasks {
214 task.abort();
215 }
216 result.map_err(Into::into)
217 }
218}
219
220fn run_backup(state: &AppState) {
221 if state.cfg.backup_dir.is_empty() {
222 return;
223 }
224 match state
225 .store
226 .backup(&state.cfg.backup_dir, state.cfg.backup_keep)
227 {
228 Ok(dest) => {
229 state.write().last_backup_at = now();
230 eprintln!("backup written: {}", dest.display());
231 }
232 Err(e) => eprintln!("backup failed: {e:#}"),
233 }
234}
235
236async fn shutdown_signal() {
237 let ctrl_c = async {
238 let _ = tokio::signal::ctrl_c().await;
239 };
240 #[cfg(unix)]
241 let terminate = async {
242 match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) {
243 Ok(mut sig) => {
244 sig.recv().await;
245 }
246 Err(_) => std::future::pending::<()>().await,
247 }
248 };
249 #[cfg(not(unix))]
250 let terminate = std::future::pending::<()>();
251
252 tokio::select! {
253 _ = ctrl_c => {}
254 _ = terminate => {}
255 }
256}
257
258async fn guard(State(state): State<Arc<AppState>>, req: Request, next: Next) -> Response {
264 if state
265 .limiter
266 .limited(&client_ip(&req, &state.cfg.trusted_ip_header))
267 {
268 let mut resp = error(
269 StatusCode::TOO_MANY_REQUESTS,
270 "rate limit exceeded, try again later",
271 );
272 if let Ok(v) = state
273 .cfg
274 .rate_limit_window
275 .as_secs()
276 .to_string()
277 .parse::<axum::http::HeaderValue>()
278 {
279 resp.headers_mut().insert("retry-after", v);
280 }
281 return resp;
282 }
283 if !authorized(&state.cfg.token, req.headers()) {
284 return error(StatusCode::UNAUTHORIZED, "unauthorized");
285 }
286 next.run(req).await
287}
288
289fn authorized(token: &str, headers: &HeaderMap) -> bool {
290 let Some(value) = headers
291 .get(axum::http::header::AUTHORIZATION)
292 .and_then(|v| v.to_str().ok())
293 .and_then(|v| v.strip_prefix("Bearer "))
294 else {
295 return false;
296 };
297 !value.is_empty() && constant_time_eq(value.as_bytes(), token.as_bytes())
298}
299
300fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
304 if a.len() != b.len() {
305 return false;
306 }
307 let mut diff = 0u8;
308 for (x, y) in a.iter().zip(b) {
309 diff |= x ^ y;
310 }
311 std::hint::black_box(diff) == 0
312}
313
314fn client_ip(req: &Request, trusted_header: &str) -> String {
328 if !trusted_header.is_empty() {
329 if let Some(ip) = header_str(req.headers(), trusted_header) {
330 return ip.to_string();
331 }
332 }
333 req.extensions()
334 .get::<ConnectInfo<SocketAddr>>()
335 .map(|ConnectInfo(addr)| addr.ip().to_string())
336 .unwrap_or_else(|| "unknown".to_string())
337}
338
339fn header_str<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> {
340 headers
341 .get(name)
342 .and_then(|v| v.to_str().ok())
343 .map(str::trim)
344 .filter(|v| !v.is_empty())
345}
346
347async fn handle_push(State(state): State<Arc<AppState>>, body: Bytes) -> Response {
350 let req: PushRequest = match serde_json::from_slice(&body) {
351 Ok(req) => req,
352 Err(_) => {
353 let missing_field = serde_json::from_slice::<serde_json::Value>(&body)
358 .ok()
359 .and_then(|v| {
360 v.as_object()
361 .map(|o| !o.contains_key("project_key") || !o.contains_key("file_path"))
362 })
363 .unwrap_or(false);
364 return error(
365 StatusCode::BAD_REQUEST,
366 if missing_field {
367 REQUIRED_FIELDS_MSG
368 } else {
369 "invalid json body"
370 },
371 );
372 }
373 };
374
375 if req.project_key.is_empty() || req.file_path.is_empty() {
376 return error(StatusCode::BAD_REQUEST, REQUIRED_FIELDS_MSG);
377 }
378 if recall_wire::validate_file_path(&req.file_path).is_err() {
379 return error(
380 StatusCode::BAD_REQUEST,
381 "file_path must be relative, no traversal",
382 );
383 }
384
385 let updated_at = now();
386
387 if req.deleted {
388 if let Err(e) = state.store.tombstone(
389 &req.project_key,
390 &req.file_path,
391 &req.source_env,
392 &updated_at,
393 ) {
394 return internal(e);
395 }
396 return json(
397 StatusCode::OK,
398 &PushResponse {
399 ok: true,
400 project_key: req.project_key,
401 file_path: req.file_path,
402 deleted: true,
403 merged: false,
404 updated_at,
405 },
406 );
407 }
408
409 let existing = match state.store.get(&req.project_key, &req.file_path) {
410 Ok(e) => e,
411 Err(e) => return internal(e),
412 };
413
414 let Some(incoming) = req.content.clone() else {
418 return error(StatusCode::BAD_REQUEST, REQUIRED_FIELDS_MSG);
419 };
420
421 let mut content = incoming.clone();
422 let mut merged = false;
423
424 if let Some(stored) = should_merge(&state, existing.as_ref(), &incoming) {
430 match state.merger.merge(stored, &incoming).await {
431 Ok(out) => {
432 content = out;
433 merged = true;
434 let mut rt = state.write();
435 rt.last_merge_at = now();
436 rt.last_merge_error = None;
437 }
438 Err(e) => {
439 eprintln!(
443 "merge failed for {}/{}, falling back to last-write-wins: {e}",
444 req.project_key, req.file_path
445 );
446 state.write().last_merge_error = Some(MergeError {
447 message: e.to_string(),
448 at: now(),
449 });
450 }
451 }
452 }
453
454 if let Err(e) = state.store.upsert(
455 &req.project_key,
456 &req.file_path,
457 &content,
458 &req.source_env,
459 &updated_at,
460 ) {
461 return internal(e);
462 }
463 json(
464 StatusCode::OK,
465 &PushResponse {
466 ok: true,
467 project_key: req.project_key,
468 file_path: req.file_path,
469 deleted: false,
470 merged,
471 updated_at,
472 },
473 )
474}
475
476fn should_merge<'a>(
479 state: &AppState,
480 existing: Option<&'a crate::store::Existing>,
481 incoming: &str,
482) -> Option<&'a str> {
483 if !state.cfg.merge_enabled {
484 return None;
485 }
486 let stored = existing.filter(|e| !e.deleted && e.content != incoming)?;
487 state
491 .read()
492 .claude_status
493 .logged_in
494 .then_some(stored.content.as_str())
495}
496
497async fn handle_pull(
498 State(state): State<Arc<AppState>>,
499 Query(params): Query<HashMap<String, String>>,
500) -> Response {
501 let Some(project_key) = params.get("project_key").filter(|k| !k.is_empty()) else {
502 return error(
503 StatusCode::BAD_REQUEST,
504 "project_key query param is required",
505 );
506 };
507 match state.store.list(project_key) {
508 Ok(files) => json(
509 StatusCode::OK,
510 &SyncResponse {
511 project_key: project_key.clone(),
512 files,
513 },
514 ),
515 Err(e) => internal(e),
516 }
517}
518
519async fn handle_health(State(state): State<Arc<AppState>>) -> Response {
520 let last_sync_at = match state.store.last_sync_at() {
521 Ok(v) => v,
522 Err(e) => return internal(e),
523 };
524
525 let rt = state.read();
526 let claude_cli = if rt.claude_status.checked_at.is_empty() {
527 ClaudeCliStatus::default()
528 } else {
529 ClaudeCliStatus {
530 checked_at: rt.claude_status.checked_at.clone(),
531 available: Some(rt.claude_status.available),
532 logged_in: Some(rt.claude_status.logged_in),
533 error: rt.claude_status.error.clone(),
534 }
535 };
536 let body = Health {
537 status: "ok".to_string(),
538 git_commit: state.cfg.git_commit.clone(),
539 started_at: state.started_at.clone(),
540 last_sync_at,
541 last_backup_at: rt.last_backup_at.clone(),
542 merge: MergeStatus {
543 enabled: state.cfg.merge_enabled,
544 claude_cli,
545 last_merge_at: rt.last_merge_at.clone(),
546 last_merge_error: rt.last_merge_error.clone(),
547 },
548 };
549 drop(rt);
550 json(StatusCode::OK, &body)
551}
552
553async fn handle_admin_page() -> Response {
556 (
557 StatusCode::OK,
558 [
559 ("content-type", "text/html; charset=utf-8"),
560 ("x-content-type-options", "nosniff"),
561 ("content-security-policy", ADMIN_CSP),
562 ],
563 ADMIN_HTML,
564 )
565 .into_response()
566}
567
568async fn handle_admin_stats(State(state): State<Arc<AppState>>) -> Response {
569 let (projects, totals) = match state.store.admin_stats() {
570 Ok(v) => v,
571 Err(e) => return internal(e),
572 };
573 let last_backup_at = state.read().last_backup_at.clone();
574 json(
575 StatusCode::OK,
576 &AdminStats {
577 projects,
578 totals,
579 git_commit: state.cfg.git_commit.clone(),
580 last_backup_at,
581 },
582 )
583}
584
585async fn not_found() -> Response {
586 error(StatusCode::NOT_FOUND, "not found")
587}
588
589fn json<T: Serialize>(status: StatusCode, body: &T) -> Response {
590 (status, Json(body)).into_response()
591}
592
593fn error(status: StatusCode, message: &str) -> Response {
594 json(
595 status,
596 &ErrorResponse {
597 error: message.to_string(),
598 },
599 )
600}
601
602fn internal(e: anyhow::Error) -> Response {
603 error(StatusCode::INTERNAL_SERVER_ERROR, &e.to_string())
604}
605
606struct RateLimiter {
611 window: Duration,
612 max: u32,
613 state: Mutex<LimiterState>,
614}
615
616struct LimiterState {
617 buckets: HashMap<String, Bucket>,
618 last_sweep: Instant,
619}
620
621struct Bucket {
622 count: u32,
623 window_start: Instant,
624}
625
626impl RateLimiter {
627 fn new(window: Duration, max: u32) -> Self {
628 Self {
629 window,
630 max,
631 state: Mutex::new(LimiterState {
632 buckets: HashMap::new(),
633 last_sweep: Instant::now(),
634 }),
635 }
636 }
637
638 fn limited(&self, ip: &str) -> bool {
639 let now = Instant::now();
640 let mut state = self.state.lock().unwrap_or_else(PoisonError::into_inner);
641
642 if now.duration_since(state.last_sweep) >= self.window {
647 state.last_sweep = now;
648 let window = self.window;
649 state
650 .buckets
651 .retain(|_, b| now.duration_since(b.window_start) < 2 * window);
652 }
653
654 let bucket = state.buckets.entry(ip.to_string()).or_insert(Bucket {
655 count: 0,
656 window_start: now,
657 });
658 if now.duration_since(bucket.window_start) >= self.window {
659 bucket.count = 0;
660 bucket.window_start = now;
661 }
662 bucket.count += 1;
663 bucket.count > self.max
664 }
665}
666
667#[cfg(test)]
668mod tests {
669 use super::*;
670
671 #[test]
672 fn rate_limiter_counts_per_ip_and_resets_after_the_window() {
673 let rl = RateLimiter::new(Duration::from_millis(40), 2);
674 assert!(!rl.limited("a"));
675 assert!(!rl.limited("a"));
676 assert!(rl.limited("a"), "third request in the window is limited");
677 assert!(!rl.limited("b"), "a different IP has its own bucket");
678
679 std::thread::sleep(Duration::from_millis(60));
680 assert!(!rl.limited("a"), "the window resets");
681 }
682
683 #[test]
684 fn rate_limiter_sweeps_stale_buckets() {
685 let rl = RateLimiter::new(Duration::from_millis(10), 100);
686 for i in 0..50 {
687 rl.limited(&format!("10.0.0.{i}"));
688 }
689 std::thread::sleep(Duration::from_millis(30));
690 rl.limited("10.0.1.1");
691 let n = rl
692 .state
693 .lock()
694 .unwrap_or_else(PoisonError::into_inner)
695 .buckets
696 .len();
697 assert_eq!(n, 1, "stale buckets should have been swept, got {n}");
698 }
699
700 #[test]
701 fn bearer_comparison_rejects_everything_but_the_exact_token() {
702 let mut h = HeaderMap::new();
703 assert!(!authorized("secret", &h), "no header");
704 h.insert("authorization", "secret".parse().unwrap());
705 assert!(!authorized("secret", &h), "missing Bearer scheme");
706 h.insert("authorization", "Bearer ".parse().unwrap());
707 assert!(!authorized("secret", &h), "empty token");
708 h.insert("authorization", "Bearer secre".parse().unwrap());
709 assert!(!authorized("secret", &h), "prefix of the token");
710 h.insert("authorization", "Basic secret".parse().unwrap());
711 assert!(!authorized("secret", &h), "wrong scheme");
712 h.insert("authorization", "Bearer secret".parse().unwrap());
713 assert!(authorized("secret", &h));
714 }
715
716 fn request_with(headers: Vec<(&str, &str)>) -> Request {
717 let mut req = Request::new(axum::body::Body::empty());
718 req.extensions_mut()
719 .insert(ConnectInfo(SocketAddr::from(([127, 0, 0, 1], 1234))));
720 for (k, v) in headers {
721 let name = axum::http::HeaderName::from_bytes(k.as_bytes()).unwrap();
722 req.headers_mut().insert(name, v.parse().unwrap());
723 }
724 req
725 }
726
727 #[test]
728 fn client_ip_reads_the_configured_header_then_the_socket() {
729 assert_eq!(
731 client_ip(
732 &request_with(vec![("cf-connecting-ip", "198.51.100.4")]),
733 "cf-connecting-ip"
734 ),
735 "198.51.100.4"
736 );
737 assert_eq!(
739 client_ip(
740 &request_with(vec![("x-real-ip", "198.51.100.7")]),
741 "x-real-ip"
742 ),
743 "198.51.100.7"
744 );
745 assert_eq!(
747 client_ip(&request_with(vec![]), "cf-connecting-ip"),
748 "127.0.0.1"
749 );
750 assert_eq!(
752 client_ip(
753 &request_with(vec![("cf-connecting-ip", "198.51.100.4")]),
754 ""
755 ),
756 "127.0.0.1"
757 );
758 }
759
760 #[test]
768 fn a_header_the_ingress_does_not_set_is_ignored() {
769 let attacker = request_with(vec![
770 ("cf-connecting-ip", "1.1.1.1"),
771 ("x-forwarded-for", "2.2.2.2"),
772 ("true-client-ip", "3.3.3.3"),
773 ("x-real-ip", "198.51.100.7"),
774 ]);
775 assert_eq!(
776 client_ip(&attacker, "x-real-ip"),
777 "198.51.100.7",
778 "only the configured header may decide the bucket"
779 );
780
781 assert_eq!(client_ip(&attacker, "cf-connecting-ip"), "1.1.1.1");
784 }
785
786 #[test]
792 fn forwarded_for_is_no_longer_split_and_trusted() {
793 let req = request_with(vec![("x-forwarded-for", "203.0.113.9, 10.0.0.1")]);
794 assert_ne!(
795 client_ip(&req, "cf-connecting-ip"),
796 "203.0.113.9",
797 "x-forwarded-for must not be consulted when it is not the configured header"
798 );
799 assert_eq!(client_ip(&req, "cf-connecting-ip"), "127.0.0.1");
800 }
801}