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 => {} } }