diff options
| author | Somhairle H. Marisol <[email protected]> | 2026-09-17 14:32:37 +0800 |
|---|---|---|
| committer | Somhairle H. Marisol <[email protected]> | 2026-09-17 14:32:37 +0800 |
| commit | 5c0ba37eda80d39e6ceca59bb1d5f4942f858995 (patch) | |
| tree | 948723f9cedf7ccb0707fa6ee516bd30fe20fd10 /server/src/datasets.rs | |
| download | strategy-lab-5c0ba37eda80d39e6ceca59bb1d5f4942f858995.tar.gz | |
chore: establish Strategy Lab source baseline (development, not release)
Diffstat (limited to 'server/src/datasets.rs')
| -rw-r--r-- | server/src/datasets.rs | 244 |
1 files changed, 244 insertions, 0 deletions
diff --git a/server/src/datasets.rs b/server/src/datasets.rs new file mode 100644 index 0000000..5b2428f --- /dev/null +++ b/server/src/datasets.rs @@ -0,0 +1,244 @@ +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<S>; +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<String>, +} + +#[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<String>, + pub instruments: Vec<InstrumentReq>, + pub start_date: String, + pub end_date: String, + pub frequency: String, + pub adjustment: String, + pub fields: Vec<String>, +} + +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::<Vec<_>>(), + "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<String> { + 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<Value> { + let manifest: Option<String> = r.get(6)?; + let manifest_v: Value = manifest.and_then(|m| serde_json::from_str::<Value>(&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::<Value>(&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<String>>(4)?, + "cache_hit": r.get::<_, Option<i64>>(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<String> = 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<Value> { + 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<Json<Value>> { + let items: Vec<Value> = 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<String>) -> AppResult<Json<Value>> { + 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<Value> { + 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<Json<DatasetRequest>>) -> AppResult<(axum::http::StatusCode, Json<Value>)> { + 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<String> = 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<String>) -> AppResult<Json<Value>> { + 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"))); + } +} |
