Files
sanguo_vnpy_v2/tests/factor/test_eval_store.py
T

58 lines
2.5 KiB
Python

# tests/factor/test_eval_store.py
"""eval_store:建库/建run/UPSERT落盘/查询过滤/finish."""
import sys, os
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
from sanguo_factor import eval_store
def _rows():
return [
{"factor": "alpha2", "category": "alpha101", "expression": "(-1)*ts_corr(...)",
"metrics": {"1": {"ic_mean": 0.05}, "turnover": 0.3}},
{"factor": "kmid", "category": "alpha158", "expression": "(close-open)/open",
"metrics": {"1": {"ic_mean": -0.02}, "turnover": 0.6}},
]
def test_roundtrip(tmp_path):
db = str(tmp_path / "factor_eval.db")
eval_store.init_db(db)
run_id = eval_store.create_run(db, label="冒烟", universe="custom", symbols_count=50,
factors_total=2, start="2024-01-01", end="2024-12-31",
params={"periods": [1, 5, 10]})
eval_store.save_results(db, run_id, _rows())
eval_store.finish_run(db, run_id, "done", factors_done=2)
runs = eval_store.list_runs(db)
assert len(runs) == 1 and runs[0]["run_id"] == run_id and runs[0]["status"] == "done"
rows = eval_store.get_rows(db, run_id)
assert len(rows) == 2
assert rows[0]["metrics"]["1"]["ic_mean"] == 0.05
assert eval_store.get_detail(db, run_id, "kmid")["category"] == "alpha158"
assert eval_store.get_detail(db, run_id, "nope") is None
def test_filter_by_category_and_search(tmp_path):
db = str(tmp_path / "factor_eval.db")
eval_store.init_db(db)
run_id = eval_store.create_run(db, label="x", universe="all_a", symbols_count=1,
factors_total=2, start="", end="", params={})
eval_store.save_results(db, run_id, _rows())
assert len(eval_store.get_rows(db, run_id, category="alpha101")) == 1
assert len(eval_store.get_rows(db, run_id, search="corr")) == 1 # 搜表达式
assert len(eval_store.get_rows(db, run_id, search="kmid")) == 1 # 搜因子名
def test_upsert_same_run(tmp_path):
db = str(tmp_path / "factor_eval.db")
eval_store.init_db(db)
run_id = eval_store.create_run(db, label="x", universe="all_a", symbols_count=1,
factors_total=1, start="", end="", params={})
eval_store.save_results(db, run_id, [_rows()[0]])
eval_store.save_results(db, run_id, [{**_rows()[0], "metrics": {"1": {"ic_mean": 0.09}}}] * 1)
rows = eval_store.get_rows(db, run_id)
assert len(rows) == 1 and rows[0]["metrics"]["1"]["ic_mean"] == 0.09