Files
Chan/tests/generate_golden.py
jackyu66gitandCursor 276481e02c feat(ECR-007): Wyckoff Live Structure with Confirmed/Live isolation
Add live.py lifecycle and event candidates; assemble confirmed vs live
in engine; Summary partition; execution_signal source=confirmed only.
Keep strategies untouched; do not lower Confirmed thresholds for Live.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-08-07 03:14:19 +08:00

243 lines
7.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""生成 / 校验缠论流水线 golden(行为冻结基线)。"""
from __future__ import annotations
import argparse
import json
import sys
from pathlib import Path
import numpy as np
import pandas as pd
ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
def make_ohlcv(n: int = 400, seed: int = 42) -> pd.DataFrame:
rng = np.random.default_rng(seed)
dates = pd.date_range("2024-01-01", periods=n, freq="5min", tz="UTC")
# 随机游走 + 波动,保证高低开收合法
rets = rng.normal(0, 0.002, size=n)
close = 100 * np.exp(np.cumsum(rets))
open_ = np.roll(close, 1)
open_[0] = close[0]
spread = np.abs(rng.normal(0, 0.0015, size=n)) * close
high = np.maximum(open_, close) + spread
low = np.minimum(open_, close) - spread
volume = rng.uniform(100, 1000, size=n)
return pd.DataFrame(
{
"date": dates,
"open": open_,
"high": high,
"low": low,
"close": close,
"volume": volume,
}
)
def _enum_name(v):
if v is None:
return None
name = getattr(v, "name", None)
if name:
return name
s = str(v)
return s.split(".")[-1] if "." in s else s
def _t(obj, attr="time"):
t = getattr(obj, attr, None)
if t is None:
return None
if hasattr(t, "strftime"):
return t.strftime("%Y-%m-%d %H:%M:%S")
return str(t)
def _f(v, nd=10):
if v is None:
return None
try:
return round(float(v), nd)
except (TypeError, ValueError):
return None
def serialize_pipeline(tf) -> dict:
bi_list = [
{
"idx": getattr(b, "idx", i),
"dir": _enum_name(getattr(b, "dir", None)),
"high": _f(b.get_high() if hasattr(b, "get_high") else getattr(b, "high", 0)),
"low": _f(b.get_low() if hasattr(b, "get_low") else getattr(b, "low", 0)),
"is_sure": bool(getattr(b, "is_sure", True)),
"begin": _t(getattr(b, "begin_klc", None), "end_time")
or _t(getattr(b, "begin_klc", None), "time"),
"end": _t(getattr(b, "end_klc", None), "end_time")
or _t(getattr(b, "end_klc", None), "time"),
}
for i, b in enumerate(getattr(tf, "bi_list", []) or [])
]
seg_list = [
{
"idx": getattr(s, "idx", i),
"dir": _enum_name(getattr(s, "dir", None)),
"is_sure": bool(getattr(s, "is_sure", True)),
"high": _f(s.get_high()) if hasattr(s, "get_high") else None,
"low": _f(s.get_low()) if hasattr(s, "get_low") else None,
}
for i, s in enumerate(getattr(tf, "seg_list", []) or [])
]
zs_list = [
{
"idx": getattr(z, "idx", i),
"dir": _enum_name(getattr(z, "dir", None)),
"zg": _f(getattr(z, "zg", 0) or 0),
"zd": _f(getattr(z, "zd", 0) or 0),
"is_sure": bool(getattr(z, "is_sure", True)),
}
for i, z in enumerate(getattr(tf, "zs_list", []) or [])
]
bsp_list = []
for i, p in enumerate(getattr(tf, "bsp_list", []) or []):
bsp_list.append(
{
"idx": getattr(p, "idx", i),
"type": _enum_name(getattr(p, "type", None) or getattr(p, "bsp_type", None)),
"dir": _enum_name(getattr(p, "dir", None)),
"price": _f(getattr(p, "price", 0) or getattr(p, "val", 0) or 0),
}
)
return {
"counts": {
"klu": len(getattr(tf, "klu_list", []) or []),
"klc": len(getattr(tf, "klc_list", []) or []),
"bi": len(bi_list),
"seg": len(seg_list),
"zs": len(zs_list),
"bsp": len(bsp_list),
},
"bi_list": bi_list,
"seg_list": seg_list,
"zs_list": zs_list,
"bsp_list": bsp_list,
}
def analyze_contract_keys() -> dict:
"""文档化 /api/analyze 主周期关键字段(契约冒烟用)。"""
return {
"required": sorted(
[
"timezone",
"kline_data",
"klc_list",
"bi_list",
"uncompleted_bi_list",
"seg_list",
"uncompleted_seg_list",
"zs_list",
"uncompleted_zs_list",
"bi_zs_list",
"bsp_list",
"klc_fx_info",
"macd",
"chan_macd",
"klc_trend",
"wyckoff",
]
),
"optional_when": {
"include_structure_zones": ["structure_zones"],
},
"wyckoff_keys": [
"trading_range",
"bias",
"phases",
"events",
"volume_profile",
"volume_confirm",
],
"notes": "wyckoff 随主周期 analyze 默认返回;有次/次次周期时另附 element_wyckoff / sub_sub_wyckoffinclude_wyckoff=0 可跳过;elements_only 时不返回",
}
def run_pipeline(df: pd.DataFrame):
"""与 web.analyze_chan 同序,避免依赖 TF_DF.__init__ 历史缺口。"""
from chanlun import TF_DF
chan = TF_DF()
df = chan.add_indicators(df.copy())
klu_list = chan.get_kl_data(df)
klc_list = chan.get_klc_list(klu_list)
bi_list = chan.cal_bi_list(klc_list)
seg_list = chan.get_seg_list(bi_list)
zs_list = chan.calculate_seg_zs(seg_list)
bi_zs_list = chan.cal_bi_zs(seg_list)
bsp_list = chan.find_all_bsp(bi_list, bi_zs_list) if bi_zs_list else []
# 挂到伪对象供 serialize_pipeline 使用
chan.klu_list = klu_list
chan.klc_list = klc_list
chan.bi_list = bi_list
chan.seg_list = seg_list
chan.zs_list = zs_list
chan.bsp_list = bsp_list
chan.bi_zs_list = bi_zs_list
return chan
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--generate", action="store_true")
parser.add_argument("--check", action="store_true")
parser.add_argument(
"--fixture-dir",
type=Path,
default=ROOT / "tests" / "fixtures",
)
args = parser.parse_args()
args.fixture_dir.mkdir(parents=True, exist_ok=True)
ohlcv_path = args.fixture_dir / "ohlcv_5m.csv"
golden_path = args.fixture_dir / "golden_pipeline.json"
keys_path = args.fixture_dir / "analyze_contract_keys.json"
if args.generate:
df = make_ohlcv()
df.to_csv(ohlcv_path, index=False)
# 与 --check 同一路径:经 CSV 往返,避免浮点路径差异
df = pd.read_csv(ohlcv_path, parse_dates=["date"])
tf = run_pipeline(df)
payload = serialize_pipeline(tf)
golden_path.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8")
keys_path.write_text(json.dumps(analyze_contract_keys(), indent=2), encoding="utf-8")
print(f"wrote {ohlcv_path}")
print(f"wrote {golden_path} counts={payload['counts']}")
return 0
if args.check:
df = pd.read_csv(ohlcv_path, parse_dates=["date"])
tf = run_pipeline(df)
actual = serialize_pipeline(tf)
expected = json.loads(golden_path.read_text(encoding="utf-8"))
if actual != expected:
print("GOLDEN MISMATCH")
print("expected counts", expected.get("counts"))
print("actual counts", actual.get("counts"))
diff_path = args.fixture_dir / "golden_pipeline.actual.json"
diff_path.write_text(json.dumps(actual, ensure_ascii=False, indent=2), encoding="utf-8")
print(f"wrote actual to {diff_path}")
return 1
print("GOLDEN OK", actual["counts"])
return 0
parser.print_help()
return 2
if __name__ == "__main__":
raise SystemExit(main())