"""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 [--limit 50] (network OK) probe --output (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": "", "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())