summaryrefslogtreecommitdiff
path: root/worker/main.py
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 /worker/main.py
downloadstrategy-lab-5c0ba37eda80d39e6ceca59bb1d5f4942f858995.tar.gz
chore: establish Strategy Lab source baseline (development, not release)
Diffstat (limited to 'worker/main.py')
-rw-r--r--worker/main.py267
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())