summaryrefslogtreecommitdiff
path: root/server/src/main.rs
diff options
context:
space:
mode:
Diffstat (limited to 'server/src/main.rs')
-rw-r--r--server/src/main.rs667
1 files changed, 667 insertions, 0 deletions
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<Body>, 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<String> = 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<serde_json::Value> {
+ 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<serde_json::Value> {
+ 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<Option<...>>` 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<HashMap<String, String>>,
+) -> Json<serde_json::Value> {
+ let qm = q.0;
+ let query = qm.get("q").cloned().unwrap_or_default();
+ let limit = qm
+ .get("limit")
+ .and_then(|v| v.parse::<i64>().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<AppState>, 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<Body>| {
+ 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<AppState>)]
+ async fn probe_auth(
+ _cx: Cx,
+ _auth: auth::AuthUser,
+ _path: axum::extract::Path<String>,
+ ) -> Result<Json<serde_json::Value>, error::AppError> {
+ Ok(Json(serde_json::json!({})))
+ }
+
+ #[axum::debug_handler(state = Arc<AppState>)]
+ async fn probe_query(
+ _cx: Cx,
+ _q: axum::extract::Query<HashMap<String, String>>,
+ ) -> Json<serde_json::Value> {
+ 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"<html>ok</html>").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<Option<...>>`. 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::<datasets::DatasetRequest>(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::<datasets::DatasetRequest>(with_name).unwrap();
+ assert_eq!(r.name.as_deref(), Some("真实行情验收:浦发银行"));
+ }
+}
+
+fn make_state(cfg: config::Config) -> Arc<AppState> {
+ 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<AppState>) {
+ let names: Vec<String> = 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<String>>(0))?;
+ let t: Vec<Option<String>> =
+ rows.collect::<Result<Vec<Option<String>>, rusqlite::Error>>()?;
+ let out: Vec<String> = t
+ .into_iter()
+ .flatten()
+ .filter(|s| !s.is_empty())
+ .collect();
+ Ok::<Vec<String>, 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 => {} }
+}