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 /worker/main.py | |
| download | strategy-lab-5c0ba37eda80d39e6ceca59bb1d5f4942f858995.tar.gz | |
chore: establish Strategy Lab source baseline (development, not release)
Diffstat (limited to 'worker/main.py')
| -rw-r--r-- | worker/main.py | 267 |
1 files changed, 267 insertions, 0 deletions
diff --git a/worker/main.py b/worker/main.py new file mode 100644 index 0000000..bedeb03 --- /dev/null +++ b/worker/main.py @@ -0,0 +1,267 @@ +"""Strategy Lab Python worker CLI. + +Subcommands: + fetch --request /input/request.json --output /output (network OK) + backtest --request /input/request.json --output /output (network none) + search --query <text> [--limit 50] (network OK) + probe --output <dir> (real AKShare evidence) +""" +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import signal +import sys +import time +from datetime import date, timedelta +from pathlib import Path + +import akshare as ak +import pandas as pd + +from . import backtest as bt_runner +from . import data as data_mod +from .normalize import (NORMALIZATION_VERSION, SCHEMA_VERSION, build_manifest_entry, + compute_object_hash, coverage_summary, normalize_frame) + +MAX_INSTRUMENTS = 5 +MAX_YEARS = 15 +BASE_FIELDS = ["open", "high", "low", "close", "volume"] +ALLOWED_FIELDS = set(BASE_FIELDS) | {"amount", "turnover"} +EXTRA_MAP = {"amount": {"成交额", "amount"}, "turnover": {"换手率", "turnover"}} +ALLOWED_SOURCES = {"eastmoney", "tencent", "sina", "auto"} +FETCH_WALL_SECS = int(os.environ.get("STRATEGY_LAB_FETCH_WALL_SECS", "180")) + + +def _load_json(path: str): + with open(path, "r", encoding="utf-8") as f: + return json.load(f) + + +def _dump_json(path: Path, obj) -> None: + with open(path, "w", encoding="utf-8") as f: + json.dump(obj, f, ensure_ascii=False, indent=1, sort_keys=True, default=str) + + +def _slug(instrument: dict) -> str: + return instrument["symbol"].lower().replace("#", "_") + + +def _validate_request(req: dict) -> list[str]: + errors = [] + instruments = req.get("instruments") or [] + if not instruments or len(instruments) > MAX_INSTRUMENTS: + errors.append(f"1..{MAX_INSTRUMENTS} instruments required") + if req.get("frequency") != "daily": + errors.append("frequency must be 'daily'") + if req.get("adjustment") not in {"none", "qfq", "hfq"}: + errors.append("adjustment must be none|qfq|hfq") + s, e = req.get("start_date"), req.get("end_date") + try: + d0, d1 = date.fromisoformat(s), date.fromisoformat(e) + if d0 > d1: + errors.append("start_date after end_date") + if d1 - d0 > timedelta(days=365 * MAX_YEARS): + errors.append(f"range exceeds {MAX_YEARS} years POC limit") + except (TypeError, ValueError): + errors.append("start_date/end_date must be ISO dates") + for f in req.get("fields", BASE_FIELDS): + if f not in ALLOWED_FIELDS: + errors.append(f"field '{f}' not supported; supported: {sorted(ALLOWED_FIELDS)}") + if req.get("source", "auto") not in ALLOWED_SOURCES: + errors.append(f"source '{req.get('source')}' not supported; supported: {sorted(ALLOWED_SOURCES)}") + return errors + + +def _fetch_timeout(signum, frame): + raise TimeoutError("fetch wall-clock limit exceeded (bounded network call)") + + +def _write_raw(out: Path, slug: str, endpoint: str, params: dict, raw, fetched_at: str) -> tuple[str, str]: + """Immutable raw provider response object. Returns (objects-relative name, content hash).""" + payload = {"endpoint": endpoint, "params": params, "fetched_at": fetched_at, + "data": json.loads(raw.to_json(orient="records", force_ascii=False))} + data = json.dumps(payload, sort_keys=True, ensure_ascii=False).encode("utf-8") + digest = hashlib.sha256(data).hexdigest() + name = f"raw_{slug}_{digest[:12]}.json" + (out / "objects").mkdir(parents=True, exist_ok=True) + (out / "objects" / name).write_bytes(data) + return name, digest + + +def run_fetch(request_path: str, output_dir: str) -> int: + signal.signal(signal.SIGALRM, _fetch_timeout) + signal.alarm(FETCH_WALL_SECS) + req = _load_json(request_path) + out = Path(output_dir) + out.mkdir(parents=True, exist_ok=True) + t0 = time.monotonic() + errors = _validate_request(req) + result = {"status": "failed", "errors": errors, "warnings": []} if errors else None + if result: + _dump_json(out / "result.json", result) + return 1 + + instruments = req["instruments"] + objects, warnings = [], [] + ok = True + for inst in instruments: + w = list(inst.get("warnings", [])) + try: + res, src_warn = data_mod.fetch_source_with_warnings( + inst, req["start_date"], req["end_date"], req["frequency"], + req["adjustment"], req.get("fields") or BASE_FIELDS, + source=req.get("source", "auto")) + df, endpoint, params = res + raw = res.raw + w.extend(src_warn) + except data_mod.DataError as exc: + objects.append({"instrument": inst, "error": {"code": exc.code, "message": str(exc)}}) + warnings.append(f"{inst['symbol']}: {exc}") + ok = False + continue + except Exception as exc: + objects.append({"instrument": inst, + "error": {"code": "provider_error", "message": f"{type(exc).__name__}: {exc}"}}) + warnings.append(f"{inst['symbol']}: provider error {exc}") + ok = False + continue + + fetched_at = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()) + requested_flds = req.get("fields") or BASE_FIELDS + kept = [c for c in df.columns if c in ("date", "symbol", "open", "high", "low", "close", "volume") + or c in {alt for f in requested_flds if f in EXTRA_MAP for alt in EXTRA_MAP[f]}] + df = df[kept] + norm = normalize_frame(df, inst["symbol"]) + raw_name, raw_hash = _write_raw(out, _slug(inst), res.endpoint, res.params, raw, fetched_at) + csv_path = f"objects/{_slug(inst)}.csv" + (out / "objects").mkdir(exist_ok=True) + norm.to_csv(out / csv_path, index=False) + cov = coverage_summary(norm) + if cov["actual_start"] != req["start_date"] or cov["actual_end"] != req["end_date"]: + w.append(f"actual coverage {cov['actual_start']}..{cov['actual_end']} differs from request; gaps kept") + entry = build_manifest_entry( + norm, + instrument={**inst, "symbol": inst["symbol"].upper()}, + provider=res.provider, endpoint=res.endpoint, params=res.params, + adjustment=req["adjustment"], + requested_start=req["start_date"], requested_end=req["end_date"], + raw_object_hash=raw_hash, warnings=w, path=csv_path, + fetched_at=fetched_at, akshare_version=ak.__version__, + ) + objects.append(entry) + + if not ok: + result = {"status": "failed", "errors": [f"{o['instrument']['symbol']}: {o.get('error', {}).get('message')}" + for o in objects if "error" in o], + "warnings": warnings} + _dump_json(out / "result.json", result) + return 1 + + manifest = { + "schema_version": SCHEMA_VERSION, + "normalization_version": NORMALIZATION_VERSION, + "fetched_at": objects[0]["fetched_at"], + "frequency": req["frequency"], + "adjustment": req["adjustment"], + "objects": objects, + "immutable": True, + } + content = json.dumps(objects, sort_keys=True, ensure_ascii=False) + manifest["hash"] = hashlib.sha256(content.encode()).hexdigest() + + result = { + "status": "ready", + "manifest": manifest, + "preview": {"columns": [c for c in objects[0]["columns"] if c != "symbol"], + "row_count": objects[0]["row_count"], + "rows": json.loads(pd.read_csv(out / objects[0]["path"]).tail(20).to_json(orient="records")), + "coverage": coverage_summary(pd.read_csv(out / objects[0]["path"]))}, + "manifest_hash": manifest["hash"], + "cache_key": manifest["hash"], + "warnings": warnings, + "elapsed_ms": int((time.monotonic() - t0) * 1000), + "akshare_version": ak.__version__, + } + _dump_json(out / "result.json", result) + return 0 + + +def run_backtest(request_path: str, output_dir: str) -> int: + req = _load_json(request_path) + out = Path(output_dir) + res = bt_runner.run_backtest(req) + _dump_json(out / "result.json", res) + return 0 if res.get("status") == "succeeded" else 1 + + +def run_search(query: str, limit: int) -> int: + """Bounded provider suggestion search. Prints a JSON envelope with an + explicit status; failed upstream is a real error, never empty success.""" + envelope = data_mod.search_instruments(query, limit) + print(json.dumps(envelope, ensure_ascii=False)) + return 0 if envelope.get("status") == "ready" else 1 + + +def run_probe(output_dir: str) -> int: + """Real AKShare connectivity probe for qa evidence. No fabrication.""" + out = Path(output_dir) + out.mkdir(parents=True, exist_ok=True) + targets = [ + {"endpoint": "stock_zh_a_hist", "symbol": "600000", "asset_type": "stock", + "call": lambda: ak.stock_zh_a_hist(symbol="600000", period="daily", + start_date="20260101", end_date="20260930", adjust="qfq")}, + {"endpoint": "fund_etf_hist_em", "symbol": "510300", "asset_type": "etf", + "call": lambda: ak.fund_etf_hist_em(symbol="510300", period="daily", + start_date="20260101", end_date="20260930", adjust="")}, + {"endpoint": "stock_zh_index_daily_em", "symbol": "000300", "asset_type": "index", + "call": lambda: ak.stock_zh_index_daily_em(symbol="sh000300")}, + {"endpoint": "stock_info_a_code_name", "symbol": "<catalog>", "asset_type": "catalog", + "call": lambda: ak.stock_info_a_code_name()}, + ] + report = {"generated_at": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), + "akshare_version": ak.__version__, "probes": []} + for t in targets: + rec = {"endpoint": t["endpoint"], "symbol": t["symbol"], "asset_type": t["asset_type"]} + try: + df = t["call"]() + rec.update(status="ok", rows=int(len(df)), columns=list(df.columns), + head=[json.loads(df.head(2).to_json(orient="records", force_ascii=False)), + ][0]) + except Exception as exc: + rec.update(status="failed", error=f"{type(exc).__name__}: {exc}") + report["probes"].append(rec) + _dump_json(out / "akshare-probe.json", report) + ok = sum(1 for p in report["probes"] if p["status"] == "ok") + print(json.dumps({"ok": ok, "total": len(targets)})) + return 0 + + +def main(argv=None) -> int: + ap = argparse.ArgumentParser(prog="python -m worker.main") + sub = ap.add_subparsers(dest="cmd", required=True) + for name in ("fetch", "backtest"): + p = sub.add_parser(name) + p.add_argument("--request", required=True) + p.add_argument("--output", required=True) + ps = sub.add_parser("search") + ps.add_argument("--query", required=True) + ps.add_argument("--limit", type=int, default=50) + pp = sub.add_parser("probe") + pp.add_argument("--output", required=True) + args = ap.parse_args(argv) + if args.cmd == "fetch": + return run_fetch(args.request, args.output) + if args.cmd == "backtest": + return run_backtest(args.request, args.output) + if args.cmd == "search": + return run_search(args.query, args.limit) + if args.cmd == "probe": + return run_probe(args.output) + return 2 + + +if __name__ == "__main__": + sys.exit(main()) |
