summaryrefslogtreecommitdiff
path: root/server/src/datasets.rs
blob: 5b2428fb7637b3f12b16772299bf14046009462e (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
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")));
    }
}