use axum::{extract::{Path, State}, Json}; use chrono::NaiveDate; use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; use crate::auth::{audit, AuthUser}; use crate::config::{MAX_RANGE_YEARS, MAX_SYMBOLS}; use crate::error::{AppError, AppResult}; use crate::state::S; use crate::util::{canonical_json, new_id, now_iso, sha256_hex}; /// Matches the fixed backend contract `Cx`; wired by State extractor in main.rs. pub type Cx = State; const FIELDS: [&str; 6] = ["open", "high", "low", "close", "volume", "adj_factor"]; #[derive(Deserialize, Serialize, Clone)] pub struct InstrumentReq { pub symbol: String, pub market: String, pub asset_type: String, pub name: Option, } #[derive(Deserialize)] pub struct DatasetRequest { /// Optional: SPEC/UI permit an absent or blank name; the backend then /// auto-generates a descriptive name and persists it. /// Optional: SPEC/UI permit an absent or blank name; the backend then /// auto-generates a descriptive name and persists it. #[serde(default)] pub name: Option, pub instruments: Vec, pub start_date: String, pub end_date: String, pub frequency: String, pub adjustment: String, pub fields: Vec, } pub fn canonical_request_value(req: &DatasetRequest) -> Value { json!({ "instruments": req.instruments.iter().map(|i| json!({ "symbol": i.symbol, "market": i.market, "asset_type": i.asset_type, "name": i.name.clone(), })).collect::>(), "start_date": req.start_date, "end_date": req.end_date, "frequency": req.frequency, "adjustment": req.adjustment, "fields": req.fields, }) } pub fn validate_request(req: &DatasetRequest) -> AppResult { if req.instruments.is_empty() || req.instruments.len() > MAX_SYMBOLS { return Err(AppError::bad("validation", format!("instruments must be 1-{} items", MAX_SYMBOLS))); } if req.frequency != "daily" { return Err(AppError::bad("validation", "only daily frequency is supported")); } if !matches!(req.adjustment.as_str(), "none" | "qfq" | "hfq") { return Err(AppError::bad("validation", "adjustment must be none|qfq|hfq")); } if req.fields.is_empty() { return Err(AppError::bad("validation", "fields must not be empty")); } for f in &req.fields { if !FIELDS.contains(&f.as_str()) { return Err(AppError::bad("validation", format!("unsupported field: {f}"))); } } let mut seen = std::collections::HashSet::new(); for i in &req.instruments { if i.symbol.trim().is_empty() || i.market.trim().is_empty() { return Err(AppError::bad("validation", "each instrument needs a symbol and a market")); } if !matches!(i.asset_type.as_str(), "stock" | "etf" | "index") { return Err(AppError::bad("validation", "asset_type must be stock|etf|index")); } if i.asset_type == "index" && req.adjustment != "none" { return Err(AppError::bad("validation", "index instruments support adjustment 'none' only (explicit restriction, no factors)")); } if !seen.insert(format!("{}|{}|{}", i.market, i.asset_type, i.symbol)) { return Err(AppError::bad("validation", "duplicate instrument in request")); } } let sd = NaiveDate::parse_from_str(&req.start_date, "%Y-%m-%d").map_err(|_| AppError::bad("validation", "start_date must be YYYY-MM-DD"))?; let ed = NaiveDate::parse_from_str(&req.end_date, "%Y-%m-%d").map_err(|_| AppError::bad("validation", "end_date must be YYYY-MM-DD"))?; if ed < sd { return Err(AppError::bad("validation", "end_date must not precede start_date")); } if (ed - sd).num_days() > MAX_RANGE_YEARS * 366 { return Err(AppError::bad("validation", "range exceeds maximum of 15 years")); } Ok(sha256_hex(canonical_json(&canonical_request_value(req)).as_bytes())) } /// Full stored manifest JSON -> client copy without host/internal paths. pub fn client_manifest(m: &Value) -> Value { let mut o = m.clone(); o.as_object_mut().map(|m| m.remove("preview")); if let Some(objs) = o.get_mut("objects").and_then(|v| v.as_array_mut()) { for obj in objs.iter_mut() { if let Some(map) = obj.as_object_mut() { map.remove("path"); } } } o } fn row_dataset(r: &rusqlite::Row) -> rusqlite::Result { let manifest: Option = r.get(6)?; let manifest_v: Value = manifest.and_then(|m| serde_json::from_str::(&m).ok()).unwrap_or(Value::Null); // The stored canonical request is persisted as a JSON string; clients get an object. let request_s: String = r.get(2)?; let request_v: Value = serde_json::from_str::(&request_s) .map_err(|_| rusqlite::Error::InvalidColumnType(2, "dataset request".into(), rusqlite::types::Type::Text))?; Ok(json!({ "id": r.get::<_, String>(0)?, "name": r.get::<_, String>(1)?, "request": request_v, "status": r.get::<_, String>(3)?, "error": r.get::<_, Option>(4)?, "cache_hit": r.get::<_, Option>(5)?.map(|v| v != 0), "manifest": if manifest_v.is_null() { Value::Null } else { client_manifest(&manifest_v) }, "warnings": manifest_v.get("warnings").cloned().unwrap_or(json!([])), "manifest_hash": manifest_v.get("hash").cloned().unwrap_or(Value::Null), "created_at": r.get::<_, String>(7)?, "updated_at": r.get::<_, String>(8)?, })) } pub async fn assert_owned(cx: &Cx, user_id: &str, dataset_id: &str) -> AppResult<()> { let found: Option = cx.with_db(|db| { db.query_row("SELECT user_id FROM datasets WHERE id=?1", [dataset_id], |r| r.get(0)).ok() }).await; match found { Some(o) if o == user_id => Ok(()), _ => Err(AppError::not_found("dataset not found")), } } async fn query_dataset(cx: &Cx, sql: &str, dataset_id: &str) -> AppResult { cx.with_db(move |db| { db.query_row(sql, [dataset_id], row_dataset) .map_err(|_| AppError::not_found("dataset not found")) }).await } const LIST_SQL: &str = "SELECT id,name,request,status,error,cache_hit,manifest,created_at,updated_at FROM datasets WHERE user_id=?1 ORDER BY created_at DESC LIMIT 200"; pub async fn list(cx: Cx, auth: AuthUser) -> AppResult> { let items: Vec = cx.with_db(|db| { let mut st = db.prepare(LIST_SQL)?; let mut rows = st.query([auth.id.clone()])?; let mut out = Vec::new(); while let Some(r) = rows.next()? { out.push(row_dataset(r)?); } Ok::<_, AppError>(out) }).await?; Ok(Json(json!({ "items": items }))) } pub async fn get(cx: Cx, auth: AuthUser, Path(id): Path) -> AppResult> { assert_owned(&cx, &auth.id, &id).await?; let v = query_dataset(&cx, "SELECT id,name,request,status,error,cache_hit,manifest,created_at,updated_at FROM datasets WHERE id=?1", &id).await?; Ok(Json(v)) } /// Load the stored worker manifest (with internal paths) of a ready dataset. pub async fn load_manifest(cx: &Cx, dataset_id: &str) -> AppResult { let m: String = cx.with_db(|db| { db.query_row("SELECT manifest FROM datasets WHERE id=?1 AND status='ready'", [dataset_id], |r| r.get(0)) .map_err(|_| AppError::conflict("dataset_not_ready", "dataset not ready")) }).await?; serde_json::from_str(&m).map_err(|e| AppError::internal(format!("manifest corrupt: {e}"))) } pub async fn create(cx: Cx, auth: AuthUser, body: Option>) -> AppResult<(axum::http::StatusCode, Json)> { let Json(r) = body.ok_or_else(|| AppError::bad("invalid_body", "JSON body required"))?; // Frontend sends name: optional/empty; default is an auto-generated descriptive name. let name = if r.name.as_deref().map(|n| n.trim().is_empty()).unwrap_or(true) { // Frontend allows an empty name to auto-generate: instruments + range. let syms: Vec = r.instruments.iter().map(|i| i.symbol.clone()).collect(); format!("{} · {} ~ {}", syms.join(","), r.start_date, r.end_date) } else { r.name.as_deref().unwrap_or_default().trim().to_string() }; if name.len() > 200 { return Err(AppError::bad("validation", "name required (max 200)")); } let _key = validate_request(&r)?; let stored_request = canonical_request_value(&r); let id = new_id(); let ts = now_iso(); let uid = auth.id.clone(); cx.with_db(|db| -> AppResult<()> { db.execute("INSERT INTO datasets (id,user_id,name,request,status,cache_hit,created_at,updated_at) VALUES (?1,?2,?3,?4,'pending',0,?5,?5)", rusqlite::params![&id, &uid, &name, stored_request.to_string(), &ts])?; Ok(()) }).await?; audit(&cx, Some(&auth.id), "dataset_create", &id, "ok").await; let v = query_dataset(&cx, "SELECT id,name,request,status,error,cache_hit,manifest,created_at,updated_at FROM datasets WHERE id=?1", &id).await?; Ok((axum::http::StatusCode::ACCEPTED, Json(v))) } pub async fn preview(cx: Cx, auth: AuthUser, Path(id): Path) -> AppResult> { assert_owned(&cx, &auth.id, &id).await?; let m = load_manifest(&cx, &id).await?; let p = m.get("preview").cloned().unwrap_or(Value::Null); if p.is_null() { return Err(AppError::not_found("preview not available yet")); } Ok(Json(p)) } #[cfg(test)] mod tests { use super::*; fn req() -> DatasetRequest { serde_json::from_value(json!({ "name": "t", "instruments": [{"symbol": "600000", "market": "cn", "asset_type": "stock", "name": "浦发银行"}], "start_date": "2024-01-01", "end_date": "2024-06-30", "frequency": "daily", "adjustment": "none", "fields": ["open","high","low","close","volume"] })).unwrap() } #[test] fn validate_rejects_unsupported() { assert_eq!(validate_request(&req()).unwrap().len(), 64); let mut r = req(); r.frequency = "hourly".into(); assert_eq!(validate_request(&r).unwrap_err().code, "validation"); let mut r = req(); r.adjustment = "qfq".into(); r.instruments[0].asset_type = "index".into(); assert_eq!(validate_request(&r).unwrap_err().code, "validation", "index+adjustment must be explicit rejections"); let mut r = req(); r.instruments.push(r.instruments[0].clone()); assert_eq!(validate_request(&r).unwrap_err().code, "validation", "duplicate instruments rejected"); } #[test] fn cache_key_is_stable_and_shared_regardless_of_display_name() { let other = { let mut o = req(); o.name = Some("别的名字".to_string()); o }; assert_eq!(canonical_request_value(&req()), canonical_request_value(&other)); assert_eq!(validate_request(&req()).unwrap(), validate_request(&other).unwrap()); } #[test] fn client_manifest_strips_internal_paths_and_preview() { let m = json!({ "hash": "h", "warnings": [], "preview": {"rows": [1]}, "objects": [{"instrument": {"symbol": "SH#600000"}, "path": "objects/ab/ab12.csv"}] }); let c = client_manifest(&m); assert!(serde_json::to_string(&c).unwrap().find("objects/ab").is_none(), "host paths must not leak"); assert_eq!(c.get("hash"), Some(&json!("h"))); } }