summaryrefslogtreecommitdiff
path: root/server/src/datasets.rs
diff options
context:
space:
mode:
authorSomhairle H. Marisol <[email protected]>2026-09-17 14:32:37 +0800
committerSomhairle H. Marisol <[email protected]>2026-09-17 14:32:37 +0800
commit5c0ba37eda80d39e6ceca59bb1d5f4942f858995 (patch)
tree948723f9cedf7ccb0707fa6ee516bd30fe20fd10 /server/src/datasets.rs
downloadstrategy-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.rs244
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")));
+ }
+}