From 5c0ba37eda80d39e6ceca59bb1d5f4942f858995 Mon Sep 17 00:00:00 2001 From: "Somhairle H. Marisol" Date: Thu, 17 Sep 2026 14:32:37 +0800 Subject: chore: establish Strategy Lab source baseline (development, not release) --- server/src/main.rs | 667 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 667 insertions(+) create mode 100644 server/src/main.rs (limited to 'server/src/main.rs') diff --git a/server/src/main.rs b/server/src/main.rs new file mode 100644 index 0000000..87f5467 --- /dev/null +++ b/server/src/main.rs @@ -0,0 +1,667 @@ +mod admin; +mod ai; +mod auth; +mod config; +mod db; +mod datasets; +mod error; +mod jobs; +mod projects; +mod runs; +mod state; +mod store; +mod util; +mod worker; + +use std::collections::HashMap; +use std::path::PathBuf; +use std::sync::Arc; + +use axum::{ + body::Body, + extract::Request, + http::{header, HeaderValue, Method, StatusCode}, + middleware::{self, Next}, + response::{IntoResponse, Response}, + routing::{delete, get, patch, post, put}, + Json, Router, +}; +use tokio::sync::Mutex as AsyncMutex; + +use crate::state::{AppState, Cx}; + +const MIME_FALLBACK: &str = "application/octet-stream"; + +/// CSRF / content protections for writes. +async fn csrf_middleware(req: Request, next: Next) -> Response { + if !is_write(req.method()) { + return next.run(req).await; + } + let headers = req.headers().clone(); + if let Some(origin) = headers.get(header::ORIGIN).and_then(|v| v.to_str().ok()) { + let host = headers.get(header::HOST).and_then(|v| v.to_str().ok()).unwrap_or(""); + // No trust of X-Forwarded-* hosts; only exact canonical/host origin matches. + if !origin_allowed(origin, host, canonical_origin()) { + return AppErr::forbidden("cross-origin write rejected").into_response(); + } + } + if let Some(site) = headers.get("sec-fetch-site").and_then(|v| v.to_str().ok()) { + if site == "cross-site" { + return AppErr::forbidden("cross-site request rejected").into_response(); + } + } + // JSON writes only. DELETE carries no body: allow an absent content type. + let ctype = headers + .get(header::CONTENT_TYPE) + .and_then(|v| v.to_str().ok()) + .unwrap_or("") + .to_string(); + let missing_ok = req.method() == Method::DELETE && ctype.is_empty(); + if !(missing_ok || ctype.starts_with("application/json")) { + return AppErr::forbidden("JSON Content-Type required").into_response(); + } + next.run(req).await +} + +static CANONICAL_ORIGIN: std::sync::OnceLock = std::sync::OnceLock::new(); + +fn canonical_origin() -> &'static str { + CANONICAL_ORIGIN.get().map(|s| s.as_str()).unwrap_or("") +} + +/// Exact-origin check. Substring matching is exploitable: +/// an attacker-supplied Origin `https://allowed.example.evil.invalid` must fail. +/// Empty canonical origin: exact local (Host-based) origin only. +pub fn origin_allowed(origin: &str, host: &str, canonical: &str) -> bool { + let origin = origin.trim(); + if !(origin.starts_with("http://") || origin.starts_with("https://")) { + return false; + } + if !canonical.is_empty() { + // Trust only the configured canonical https origin, compared exactly. + return origin.eq_ignore_ascii_case(canonical.trim()); + } + if host.is_empty() { + // No Host and no canonical: cannot establish trust; reject. + return false; + } + origin == format!("http://{host}") || origin == format!("https://{host}") +} + +fn set_canonical(origin: &str) { + let _ = CANONICAL_ORIGIN.set(origin.trim().to_string()); +} + +fn is_write(m: &Method) -> bool { + matches!(*m, Method::POST | Method::PUT | Method::PATCH | Method::DELETE) +} + +type AppErr = error::AppError; + +async fn health(cx: Cx) -> Json { + let docker = worker::docker_available().await; + Json(serde_json::json!({ + "status": "ok", + "version": cx.cfg.version, + "worker_available": docker, + "ai_configured": std::env::var("OPENCODE_GO_API_KEY").map(|_| true).unwrap_or(false), + })) +} + +async fn capabilities() -> Json { + Json(serde_json::json!({ + "frequencies": ["daily"], + "asset_types": ["stock", "etf", "index"], + "adjustments": [ + {"code": "none", "label": "不复权"}, + {"code": "qfq", "label": "前复权"}, + {"code": "hfq", "label": "后复权"} + ], + "fields": [ + {"code": "open", "label": "开盘", "raw": false}, + {"code": "high", "label": "最高", "raw": false}, + {"code": "low", "label": "最低", "raw": false}, + {"code": "close", "label": "收盘", "raw": false}, + {"code": "volume", "label": "成交量", "raw": false}, + {"code": "adj_factor", "label": "复权因子", "raw": true} + ], + "limits": {"max_symbols": 5, "max_years": 15, "internal_only": true}, + })) +} + +/// Instrument search through the real worker container catalog command. +/// Failures are surfaced honestly in `status`; no fake empty success. +/// NOTE: Query deserialization: a plain `HashMap` accepts both absent and +/// present query params. `Query>` rejects any non-empty query with +/// `invalid type: map, expected option` (HTTP400 observed in live browser QA). +async fn instruments( + cx: Cx, + q: axum::extract::Query>, +) -> Json { + let qm = q.0; + let query = qm.get("q").cloned().unwrap_or_default(); + let limit = qm + .get("limit") + .and_then(|v| v.parse::().ok()) + .unwrap_or(50); + match worker::search_instruments(&cx, &query, limit).await { + Ok(items) => { + let source = if items.is_empty() { "none" } else { "provider_suggest" }; + Json(serde_json::json!({"items": items, "source": source, "status": "ok"})) + } + Err(e) => Json(serde_json::json!({ + "items": [], + "source": "none", + "status": format!("unavailable: {}", e.message) + })), + } +} + +async fn api_fallback() -> Response { + ( + StatusCode::NOT_FOUND, + Json(serde_json::json!({"error": {"code": "not_found", "message": "unknown API route"}})), + ) + .into_response() +} + +fn mime_of(path: &std::path::Path) -> &'static str { + match path.extension().and_then(|e| e.to_str()).unwrap_or("") { + "html" => "text/html; charset=utf-8", + "css" => "text/css; charset=utf-8", + "js" | "mjs" => "text/javascript; charset=utf-8", + "json" => "application/json", + "svg" => "image/svg+xml", + "png" => "image/png", + "webp" => "image/webp", + "woff2" => "font/woff2", + "woff" => "font/woff", + "ico" => "image/x-icon", + "map" => "application/json", + "txt" => "text/plain; charset=utf-8", + _ => MIME_FALLBACK, + } +} + +/// Serve the built SPA from frontend/dist. Exact files when they exist, +/// otherwise /index.html so client routes work; unknown /api is handled by the +/// inner fallback above and never falls back to the SPA. + +/// Static file resolution for the SPA (GET/HEAD only). +async fn serve_static(root: PathBuf, path: &str) -> Response { + let rel = path.trim_start_matches('/'); + if rel.contains("..") || rel.contains('\\') { + return StatusCode::NOT_FOUND.into_response(); + } + let base = root.join("index.html"); + let target = if rel.is_empty() { + base + } else { + let p = root.join(rel); + if p.is_file() { + p + } else if p.is_dir() || !p.exists() { + base + } else { + return StatusCode::NOT_FOUND.into_response(); + } + }; + match tokio::fs::read(&target).await { + Ok(bytes) => { + let mut resp = ( + StatusCode::OK, + [(header::CONTENT_TYPE, mime_of(&target))], + bytes, + ).into_response(); + if mime_of(&target) != "text/html" { + resp.headers_mut().insert( + header::CACHE_CONTROL, + HeaderValue::from_static("no-cache"), + ); + } + resp + } + Err(_) => StatusCode::NOT_FOUND.into_response(), + } +} + +pub fn build_app(cx: Arc, frontend_dir: String) -> Router { + let api = Router::new() + .route("/health", get(health)) + .route("/capabilities", get(capabilities)) + .route("/instruments", get(instruments)) + .route("/auth/login", post(auth::login)) + .route("/auth/register", post(auth::register)) + .route("/auth/logout", post(auth::logout)) + .route("/auth/me", get(auth::me)) + .route("/auth/profile", patch(auth::patch_profile)) + .route("/auth/password", post(auth::change_password)) + .route("/auth/sessions", get(auth::list_sessions)) + .route("/auth/sessions/{id}", delete(auth::delete_session)) + .route("/auth/reset-password", post(auth::reset_password)) + .route("/projects", get(projects::list).post(projects::create)) + .route("/projects/{id}", get(projects::get).patch(projects::patch)) + .route("/projects/{id}/draft", put(projects::put_draft)) + .route("/projects/{id}/versions", get(projects::list_versions).post(projects::create_version)) + .route("/projects/{id}/versions/{vid}", get(projects::get_version)) + .route("/projects/{id}/versions/{vid}/diff", post(projects::diff_versions)) + .route("/projects/{id}/restore", post(projects::restore)) + .route("/datasets", get(datasets::list).post(datasets::create)) + .route("/datasets/{id}", get(datasets::get)) + .route("/datasets/{id}/preview", get(datasets::preview)) + .route("/runs", get(runs::list).post(runs::enqueue)) + .route("/runs/{id}", get(runs::get)) + .route("/runs/{id}/cancel", post(runs::cancel)) + .route("/runs/{id}/rerun", post(runs::rerun)) + .route("/ai/assist", post(ai::assist)) + .route("/ai/usage", get(ai::list_ai_usage)) + .route("/ai/{id}/accept", post(ai::accept)) + .route("/admin/users", get(admin::users)) + .route("/admin/users/{id}", patch(admin::patch_user)) + .route("/admin/invitations", get(admin::list_invitations).post(admin::create_invitation)) + .route("/admin/invitations/{id}", delete(admin::delete_invitation)) + .route("/admin/users/{id}/reset-password", post(admin::create_reset)) + .route("/admin/audit", get(admin::audit_list)) + .fallback(api_fallback) + .layer(middleware::from_fn(csrf_middleware)) + .with_state(cx.clone()); + + let spa = move |req: Request| { + let root = frontend_dir.clone(); + async move { + let method = req.method().clone(); + let path = req.uri().path().to_string(); + if !matches!(method, Method::GET | Method::HEAD) { + return StatusCode::METHOD_NOT_ALLOWED.into_response(); + } + serve_static(PathBuf::from(root), &path).await + } + }; + + Router::new().nest("/api", api).fallback(spa) +} + +#[cfg(test)] +mod probe { + use super::*; + + fn probe_static(root: &str, path: &str) -> Response { + let rt = tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .unwrap(); + rt.block_on(serve_static(PathBuf::from(root), path)) + } + + #[axum::debug_handler(state = Arc)] + async fn probe_auth( + _cx: Cx, + _auth: auth::AuthUser, + _path: axum::extract::Path, + ) -> Result, error::AppError> { + Ok(Json(serde_json::json!({}))) + } + + #[axum::debug_handler(state = Arc)] + async fn probe_query( + _cx: Cx, + _q: axum::extract::Query>, + ) -> Json { + Json(serde_json::json!({})) + } + + #[test] + fn probe_router() { + let state = Arc::new(AppState { + cfg: config::Config::from_env(), + db: AsyncMutex::new(rusqlite::Connection::open_in_memory().unwrap()), + run_sem: Arc::new(tokio::sync::Semaphore::new(1)), + fetch_sem: Arc::new(tokio::sync::Semaphore::new(1)), + }); + let _r: Router = Router::new() + .route("/x", get(probe_auth)) + .route("/y", get(probe_query)) + .with_state(state); + } + + /// Exact-origin CSRF: substring attacks must be rejected. + #[test] + fn csrf_rejects_malicious_substring_origin() { + const CANON: &str = "https://fin.somhairle.bid"; + assert!(origin_allowed(CANON, "fin.somhairle.bid", CANON)); + // attacker-controlled suffix + assert!(!origin_allowed("https://fin.somhairle.bid.evil.invalid", "fin.somhairle.bid", CANON)); + // attacker-controlled prefix host + assert!(!origin_allowed("https://evil.fin.somhairle.bid", "fin.somhairle.bid", CANON)); + // different scheme + assert!(!origin_allowed("http://fin.somhairle.bid", "fin.somhairle.bid", CANON)); + // different port + assert!(!origin_allowed("https://fin.somhairle.bid:8443", "fin.somhairle.bid", CANON)); + // no canonical: exact local origin only + assert!(origin_allowed("http://127.0.0.1:8787", "127.0.0.1:8787", "")); + assert!(!origin_allowed("http://127.0.0.1:8787.evil.invalid", "127.0.0.1:8787", "")); + assert!(!origin_allowed("http://127.0.0.1:8787x", "127.0.0.1:8787", "")); + // no host, no canonical: reject + assert!(!origin_allowed("http://whatever", "", "")); + assert!(!origin_allowed("javascript:alert(1)", "127.0.0.1:8787", "")); + assert!(!origin_allowed("", "127.0.0.1:8787", "")); + } + + #[test] + fn spa_serves_index_for_unknown_paths_and_sanitizes_traversal() { + let td = tempfile::tempdir_in("/tmp/opencode").unwrap(); + std::fs::write(td.path().join("index.html"), b"ok").unwrap(); + std::fs::write(td.path().join("assets.js"), b"console.log(1)").unwrap(); + let resp = probe_static(td.path().to_str().unwrap(), "/"); + assert_eq!(resp.status(), 200); + assert!(resp.headers().get(header::CONTENT_TYPE).unwrap().to_str().unwrap().starts_with("text/html")); + let resp = probe_static(td.path().to_str().unwrap(), "/assets.js"); + assert_eq!(resp.status(), 200); + let resp = probe_static(td.path().to_str().unwrap(), "/unknown/route"); + // SPA fallback serves index.html for client routes + assert_eq!(resp.status(), 200); + let resp = probe_static(td.path().to_str().unwrap(), "/../../etc/passwd"); + assert_ne!(resp.status(), 200); + } + + /// Executable-level regression: boot the real binary on an isolated + /// DB/port, assert it stays alive (>10s) WITHOUT any shutdown signal, then + /// SIGTERM must produce a bounded timely exit. This formerly caught a bug + /// where the drain deadline incorrectly started at startup and killed the + /// server at ~10s of healthy uptime (premature-exit regression). + #[tokio::test] + async fn server_survives_past_10s_then_bounds_sigterm_exit() { + let td = tempfile::tempdir_in("/tmp/opencode").unwrap(); + let db_path = td.path().join("db.sqlite3"); + let data_dir = td.path().join("data"); + std::fs::create_dir_all(&data_dir).unwrap(); + // reserve a port using the OS + let probe_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = probe_listener.local_addr().unwrap().port(); + std::mem::drop(probe_listener); + + // a known-empty frontend dir keeps static serving honest for this probe + let fe = td.path().join("fe"); + std::fs::create_dir_all(&fe).unwrap(); + + let bin = std::env::var("CARGO_BIN_EXE_strategy-lab-server") + .unwrap_or_else(|_| "target/debug/strategy-lab-server".to_string()); + let bin_path = std::path::PathBuf::from(bin); + assert!(bin_path.is_file(), "isolated server binary not built at {}", bin_path.display()); + let mut child = tokio::process::Command::new(bin_path) + .env("BIND", format!("127.0.0.1:{port}")) + .env("DB_PATH", db_path.display().to_string()) + .env("DATA_DIR", data_dir.display().to_string()) + .env("FRONTEND_DIR", fe.display().to_string()) + .env("WORKER_IMAGE", "strategy-lab-worker-unset") + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()) + .spawn() + .unwrap(); + + async fn health_ok(port: u16) -> bool { + let url = format!("http://127.0.0.1:{port}/api/health"); + match reqwest::get(&url).await { + Ok(r) => r.status().is_success(), + Err(_) => false, + } + } + // server must come up + let mut up = false; + for _ in 0..100 { + if health_ok(port).await { up = true; break; } + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + } + assert!(up, "server did not come up in 20s"); + // CRITICAL: still alive WELL PAST 10 seconds with NO shutdown signal + tokio::time::sleep(std::time::Duration::from_secs(12)).await; + assert!(health_ok(port).await, "premature-exit regression: server died ~10s after startup without any signal"); + // now signal and require bounded exit + let pid = child.id().unwrap(); + nix_pid_kill_term(pid); + let deadline = std::time::Instant::now() + std::time::Duration::from_secs(20); + let mut exited = false; + while std::time::Instant::now() < deadline { + if child.try_wait().unwrap().is_some() { + exited = true; + break; + } + tokio::time::sleep(std::time::Duration::from_millis(200)).await; + } + assert!(exited, "server did not exit in bounded time after SIGTERM"); + // give the runtime a moment to reap + let _ = child.wait().await; + let pid_gone = !std::path::Path::new(format!("/proc/{pid}").as_str()).exists(); + assert!(pid_gone, "server process still alive after bounded exit check"); + } + + fn nix_pid_kill_term(pid: u32) { + // POSIX kill of the exact PID only; no process-name scans, no group ops + std::process::Command::new("kill").args(["-TERM", &pid.to_string()]) + .stdout(std::process::Stdio::null()).stderr(std::process::Stdio::null()) + .status().ok(); + } + + /// Route-level regression: `GET /api/instruments?q=600000` formerly failed + /// with HTTP400 "invalid type: map, expected option" because the query + /// extractor was `Query>`. Must be JSON 200 with `items`, + /// `source`, `status` (parent browser workflow live failure). + #[tokio::test] + async fn instruments_query_with_q_is_json_200() { + let conn = rusqlite::Connection::open_in_memory().unwrap(); + crate::db::init_db(&conn).unwrap(); + let mut cfg = crate::config::Config { + db_path: format!("{}/nonexistent.sqlite3", std::env::temp_dir().display()), + data_dir: std::env::temp_dir().display().to_string(), + worker_image: "strategy-lab-worker-unset".into(), + ..config::Config::from_env() + }; + cfg.bind_addr = "127.0.0.1:0".into(); + let st = std::sync::Arc::new(AppState { + cfg, + db: AsyncMutex::new(conn), + run_sem: Arc::new(tokio::sync::Semaphore::new(1)), + fetch_sem: Arc::new(tokio::sync::Semaphore::new(1)), + }); + let app = build_app(st, "frontend-dist-missing-for-probe".to_string()); + let req = axum::http::Request::builder() + .method("GET") + .uri("/api/instruments?q=600000") + .body(axum::body::Body::empty()) + .unwrap(); + let resp = tower::ServiceExt::oneshot(app, req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK, + "q= query must deserialize (was HTTP400 invalid type: map, expected option)"); + let body = axum::body::to_bytes(resp.into_body(), 64_000).await.unwrap(); + let v: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert!(v.is_object(), "structured JSON error/shape, never plain text: {v}"); + assert!(v.get("items").and_then(|i| i.as_array()).is_some(), "items array present"); + assert!(v.get("source").and_then(|s| s.as_str()).is_some(), "source present"); + // absent q must not 400 either + let app2 = { + let conn = rusqlite::Connection::open_in_memory().unwrap(); + crate::db::init_db(&conn).unwrap(); + let mut cfg = crate::config::Config::from_env(); + cfg.db_path = format!("{}/x.sqlite3", std::env::temp_dir().display()); + cfg.data_dir = std::env::temp_dir().display().to_string(); + cfg.bind_addr = "127.0.0.1:0".into(); + std::sync::Arc::new(AppState { + cfg, + db: AsyncMutex::new(conn), + run_sem: Arc::new(tokio::sync::Semaphore::new(1)), + fetch_sem: Arc::new(tokio::sync::Semaphore::new(1)), + }) + }; + let app2 = build_app(app2, "missing-frontend-probe".to_string()); + let req = axum::http::Request::builder().method("GET").uri("/api/instruments?") + .body(axum::body::Body::empty()).unwrap(); + let resp = tower::ServiceExt::oneshot(app2, req).await.unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + } + + /// Dataset creation shape: absent `name` and blank `name` must both + /// deserialize (SPEC/UI permit; backend auto-generates and persists). + #[test] + fn dataset_request_allows_missing_or_blank_name() { + let base = serde_json::json!({ + "instruments": [{"symbol": "600000", "market": "SH", "asset_type": "stock"}], + "start_date": "2024-01-01", "end_date": "2024-06-30", + "frequency": "daily", "adjustment": "none", + "fields": ["open", "high", "low", "close", "volume"] + }); + // missing `name` entirely (previously HTTP422 missing field `name`) + let r: datasets::DatasetRequest = serde_json::from_value(base.clone()).unwrap(); + assert!(r.name.is_none(), "absent name must deserialize as None"); + // blank name + let mut with_blank = base.clone(); + with_blank["name"] = serde_json::json!(" "); + let r = serde_json::from_value::(with_blank).unwrap(); + assert!(r.name.unwrap().trim().is_empty()); + // normal name still works + let mut with_name = base; + with_name["name"] = serde_json::json!("真实行情验收:浦发银行"); + let r = serde_json::from_value::(with_name).unwrap(); + assert_eq!(r.name.as_deref(), Some("真实行情验收:浦发银行")); + } +} + +fn make_state(cfg: config::Config) -> Arc { + if let Some(parent) = std::path::Path::new(&cfg.db_path).parent() { + std::fs::create_dir_all(parent).expect("create db directory"); + } + let conn = rusqlite::Connection::open(&cfg.db_path).expect("open db"); + db::init_db(&conn).expect("init schema"); + Arc::new(AppState { + run_sem: Arc::new(tokio::sync::Semaphore::new(cfg.run_concurrency)), + fetch_sem: Arc::new(tokio::sync::Semaphore::new(cfg.fetch_concurrency)), + cfg: cfg.clone(), + db: AsyncMutex::new(conn), + }) +} + +/// Create required runtime directories (databases, object store, job dirs). +fn ensure_directories(cfg: &config::Config) { + let data = std::path::Path::new(&cfg.data_dir); + for d in [data.join("objects"), data.join("jobs")] { + if let Err(e) = std::fs::create_dir_all(&d) { + tracing::error!("cannot create {d:?}: {e}"); + std::process::exit(1); + } + } +} + +/// Remove leftover containers from any earlier (crashed) server instance. +async fn cleanup_orphan_containers(cx: &Arc) { + let names: Vec = cx + .with_db(|db| { + let mut st = db.prepare("SELECT container_id FROM runs WHERE container_id IS NOT NULL")?; + let rows = st.query_map([], |r| r.get::<_, Option>(0))?; + let t: Vec> = + rows.collect::>, rusqlite::Error>>()?; + let out: Vec = t + .into_iter() + .flatten() + .filter(|s| !s.is_empty()) + .collect(); + Ok::, rusqlite::Error>(out) + }) + .await + .unwrap_or_default(); + let mut removed = 0usize; + for name in names { + if worker::cancel_container(&name).await { + removed += 1; + } + } + if removed > 0 { + tracing::info!("cleaned {removed} leftover worker container(s)"); + } +} + +#[tokio::main] +async fn main() { + tracing_subscriber::fmt() + .with_env_filter( + tracing_subscriber::EnvFilter::try_from_default_env() + .unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")), + ) + .init(); + let cfg = config::Config::from_env(); + set_canonical(&cfg.canonical_origin); + let mode = std::env::args().nth(1).unwrap_or_else(|| "serve".into()); + + if mode == "bootstrap-admin" { + let state = make_state(cfg.clone()); + ensure_directories(&cfg); + let cx = axum::extract::State(state.clone()); + let _ = admin::bootstrap_admin(&cx) + .await + .map_err(|e| tracing::error!("bootstrap-admin failed: {e}")); + println!("bootstrap-admin done (if ADMIN_BOOTSTRAP env configured)"); + return; + } + + ensure_directories(&cfg); + let state = make_state(cfg.clone()); + let cx = axum::extract::State(state.clone()); + match admin::bootstrap_admin(&cx).await { + Ok(_) => {} + Err(e) => tracing::error!("bootstrap_admin: {e}"), + } + + // Restart safety: runs stuck as running become failed; stale containers are + // removed by name (never a global docker prune). + runs::cleanup_interrupted(&cx).await; + cleanup_orphan_containers(&state).await; + + let app = build_app(state.clone(), cfg.frontend_dir.clone()); + + tokio::spawn(async move { + jobs::main_loop(state, jobs::Signals::new()).await; + }); + + let addr = cfg.bind_addr.clone(); + let listener = tokio::net::TcpListener::bind(&addr).await.expect("bind"); + tracing::info!("listening on {addr}"); + // NO startup deadline. The drain budget starts only after the shutdown + // signal: select between + // (a) the serve future completing normally / after graceful shutdown, and + // (b) signal-received AFTER which a 10s post-signal ceiling passes — + // the (b) arm cannot fire before the signal is delivered, so a healthy + // server with no signal stays up indefinitely. + tokio::select! { + drained = axum::serve(listener, app.into_make_service()) + .with_graceful_shutdown(wait_shutdown_signal()) => + { + match drained { + Ok(()) => tracing::info!("http drained, exiting"), + Err(e) => tracing::error!("http serve failed: {e}"), + } + } + _ = async { + wait_shutdown_signal().await; + tracing::info!("shutdown signal received; http draining (<=10s)"); + tokio::time::sleep(std::time::Duration::from_secs(10)).await; + } => { + tracing::warn!("post-signal drain budget (10s) exceeded; exiting now"); + } + } +} + +/// Waits for SIGTERM or SIGINT (whichever arrives first). +async fn wait_shutdown_signal() { + use tokio::signal::unix::{signal, SignalKind}; + let term_fut = async { + match signal(SignalKind::terminate()) { + Ok(mut s) => { s.recv().await; } + Err(_) => std::future::pending::<()>().await, + } + }; + let int_fut = async { + match signal(SignalKind::interrupt()) { + Ok(mut s) => { s.recv().await; } + Err(_) => std::future::pending::<()>().await, + } + }; + tokio::select! { _ = term_fut => {}, _ = int_fut => {} } +} -- cgit v1.2.3