feat: 量化引擎加固 — 新增测试 + 数据/因子/回测层优化
- 新增 finance/tests/ 6 个测试套件(agents/backtest/dao_upsert/factors/features/fundamental_lookahead) - 数据层: data_manager / dao 优化,新增 upsert 逻辑 - 因子层: 基本面因子抽象定位 _mapping、ROE/PE/PB 重构 - 回测层: vectorbt/engine 大改动(251 行),report 增强 - ML 层: features/backtest_integration 特征工程与回测优化 - CLI: agent_cli 重构 - config/settings 扩充配置项
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
"""
|
||||
finance 子项目回归测试包。
|
||||
|
||||
运行方式(在 finance/ 下):
|
||||
python -m pytest tests/ -v
|
||||
"""
|
||||
@@ -0,0 +1,87 @@
|
||||
"""
|
||||
Agent 层回归测试:RiskAgent 配置集中可用性 + Orchestrator 每日流程步骤隔离。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
|
||||
from agents.risk_agent import RiskAgent
|
||||
from agents.orchestrator import AgentOrchestrator, TOTAL_STEPS
|
||||
|
||||
|
||||
def _flat_dm():
|
||||
"""给出平淡行情 → 风险 low。"""
|
||||
class _DM:
|
||||
def get_daily(self, code):
|
||||
idx = pd.date_range('2025-01-01', periods=300, freq='B').strftime('%Y%m%d')
|
||||
close = np.linspace(100, 100 + 0.03 * 300, 300) + np.random.default_rng(0).normal(0, 0.2, 300)
|
||||
return pd.DataFrame({'trade_date': idx, 'close': close})
|
||||
return _DM()
|
||||
|
||||
|
||||
def _hot_dm():
|
||||
class _DM:
|
||||
def get_daily(self, code):
|
||||
idx = pd.date_range('2025-01-01', periods=300, freq='B').strftime('%Y%m%d')
|
||||
close = 100 * np.exp(np.cumsum(np.random.default_rng(1).normal(0, 0.03, 300)))
|
||||
return pd.DataFrame({'trade_date': idx, 'close': close})
|
||||
return _DM()
|
||||
|
||||
|
||||
def test_risk_config_low_level():
|
||||
r = RiskAgent(dm=_flat_dm()).execute()
|
||||
assert r['risk_level'] in ('low', 'medium', 'high')
|
||||
assert r['target_exposure'] in (0.85, 0.60, 0.30)
|
||||
|
||||
|
||||
def test_risk_config_high_level():
|
||||
r = RiskAgent(dm=_hot_dm()).execute()
|
||||
assert r['risk_level'] in ('high', 'medium')
|
||||
|
||||
|
||||
class _FailAgent:
|
||||
def execute(self, **kw):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
|
||||
class _OkAgent:
|
||||
def execute(self, **kw):
|
||||
return {"report_path": "/tmp/x.md", "risk_level": "low",
|
||||
"target_exposure": 0.8, "top_picks": [], "date": "20260601"}
|
||||
|
||||
|
||||
class _FakeDM:
|
||||
def get_stock_list(self):
|
||||
return pd.DataFrame(index=['000001.SZ', '600519.SH'])
|
||||
def sync_daily(self, code):
|
||||
return 5
|
||||
|
||||
|
||||
class _FakeSent:
|
||||
def get_scope_stocks(self):
|
||||
return ["000001.SZ", "600519.SH"]
|
||||
def compute(self, *a, **k):
|
||||
return pd.DataFrame({"news_sent_5": [0.1, 0.2]})
|
||||
|
||||
|
||||
def test_run_daily_steps_isolated():
|
||||
"""risk/selection 抛错被隔离,report/sentiment/sync 仍成功。"""
|
||||
with patch('database.dao.get_latest_trade_date', side_effect=lambda c: "20260530"):
|
||||
orch = AgentOrchestrator(dm=_FakeDM(), sent=_FakeSent())
|
||||
orch.agents = {"risk": _FailAgent(), "selection": _FailAgent(),
|
||||
"report": _OkAgent(), "research": _OkAgent()}
|
||||
res = orch.run_daily(date="20260601")
|
||||
assert "risk" in res and isinstance(res["risk"], dict) and "error" in res["risk"]
|
||||
assert "selection" in res and isinstance(res["selection"], dict) and "error" in res["selection"]
|
||||
# 失败不中断:report/sentiment 仍成功
|
||||
assert "report" in res and "error" not in res["report"]
|
||||
assert isinstance(res.get("sentiment"), pd.DataFrame)
|
||||
|
||||
|
||||
def test_orchestrator_total_steps_is_5():
|
||||
assert TOTAL_STEPS == 5
|
||||
@@ -0,0 +1,81 @@
|
||||
"""
|
||||
回测引擎回归测试:T+1 成交、涨跌停拒成交、report DatetimeIndex、截面组合资金切分。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
vectorbt = pytest.importorskip("vectorbt")
|
||||
|
||||
from backtest.vectorbt.engine import VectorBTEngine
|
||||
from backtest.strategies.momentum_breakout import MomentumBreakoutStrategy
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def price_df():
|
||||
idx = pd.date_range('2024-01-01', '2024-06-30', freq='B').strftime('%Y%m%d')
|
||||
close = 100 * np.cumprod(1 + np.random.default_rng(42).normal(0.0005, 0.02, len(idx)))
|
||||
open_p = np.roll(close, 1); open_p[0] = close[0]
|
||||
pre_close = np.roll(close, 1); pre_close[0] = close[0]
|
||||
return pd.DataFrame({
|
||||
'open': open_p, 'close': close, 'pre_close': pre_close,
|
||||
'pct_chg': (close / pre_close - 1) * 100,
|
||||
}, index=idx)
|
||||
|
||||
|
||||
def test_backtest_runs_with_t1(price_df):
|
||||
eng = VectorBTEngine()
|
||||
rep = eng.run(MomentumBreakoutStrategy(), price_df, t_plus_one=True, limit_check=True, slippage=0.0002)
|
||||
assert not rep.equity_curve.empty
|
||||
assert isinstance(rep.equity_curve.index, pd.DatetimeIndex)
|
||||
|
||||
|
||||
def test_limit_up_blocks_buy(price_df):
|
||||
eng = VectorBTEngine()
|
||||
p = price_df.copy()
|
||||
i = p.index[-5]
|
||||
p.loc[i, 'close'] = p.loc[i, 'pre_close'] * 1.099
|
||||
p.loc[i, 'pct_chg'] = 9.9
|
||||
entries = pd.Series(True, index=p.index)
|
||||
exits = pd.Series(False, index=p.index)
|
||||
eff, _ = eng._apply_limit_filters(entries, exits, p)
|
||||
assert not bool(eff.loc[i]), "涨停日应抑制买入"
|
||||
assert eff.dtype == bool
|
||||
|
||||
|
||||
def test_limit_down_blocks_sell(price_df):
|
||||
eng = VectorBTEngine()
|
||||
p = price_df.copy()
|
||||
i = p.index[-7]
|
||||
p.loc[i, 'close'] = p.loc[i, 'pre_close'] * 0.901
|
||||
p.loc[i, 'pct_chg'] = -9.9
|
||||
_, exf = eng._apply_limit_filters(
|
||||
pd.Series(False, index=p.index), pd.Series(True, index=p.index), p)
|
||||
assert not bool(exf.loc[i]), "跌停日应抑制卖出"
|
||||
|
||||
|
||||
def test_normal_day_keeps_entry(price_df):
|
||||
eng = VectorBTEngine()
|
||||
p = price_df.copy()
|
||||
eff, _ = eng._apply_limit_filters(
|
||||
pd.Series(True, index=p.index), pd.Series(False, index=p.index), p)
|
||||
assert eff.sum() > 0
|
||||
|
||||
|
||||
def test_cross_section_investable_curve():
|
||||
idx = pd.date_range('2024-01-01', '2024-06-30', freq='B').strftime('%Y%m%d')
|
||||
uni = {}
|
||||
for i in range(4):
|
||||
c = 50 * np.cumprod(1 + np.random.default_rng(i).normal(0.0003, 0.018, len(idx)))
|
||||
o = np.roll(c, 1); o[0] = c[0]
|
||||
uni[f"600{i:04d}.SH"] = pd.DataFrame({'open': o, 'close': c}, index=idx)
|
||||
eng = VectorBTEngine(initial_capital=100000)
|
||||
rep = eng.run_cross_section(MomentumBreakoutStrategy(), uni, rebalance_freq="M")
|
||||
assert not rep.equity_curve.empty
|
||||
assert isinstance(rep.equity_curve.index, pd.DatetimeIndex)
|
||||
# 组合净值可联合投资:末值与初始资本量级相当,不会因每股满额而超限
|
||||
assert rep.equity_curve.iloc[-1] > 0
|
||||
@@ -0,0 +1,69 @@
|
||||
"""
|
||||
DAO 写库回归测试:验证 upsert(INSERT ... ON DUPLICATE KEY UPDATE)而非 pandas replace,
|
||||
确保不 DROP 表、不丢失数据、DELETE+INSERT 不再分事务。
|
||||
不使用真实 DB(用假 engine 捕获 SQL)。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
|
||||
from database import dao
|
||||
from database.models import StockDaily
|
||||
from sqlalchemy.dialects import mysql
|
||||
|
||||
|
||||
class _RecConn:
|
||||
"""记录最近一次执行的 SQL。"""
|
||||
def __init__(self):
|
||||
self.sql = None
|
||||
def begin(self):
|
||||
return self
|
||||
def __enter__(self):
|
||||
return self
|
||||
def __exit__(self, *a):
|
||||
pass
|
||||
def execute(self, stmt):
|
||||
self.sql = str(stmt.compile(dialect=mysql.dialect(), compile_kwargs={"literal_binds": True}))
|
||||
return type("R", (), {"rowcount": 2})()
|
||||
|
||||
|
||||
class _RecEngine:
|
||||
def __init__(self):
|
||||
self.conn = _RecConn()
|
||||
def begin(self):
|
||||
return self.conn
|
||||
|
||||
|
||||
def _daily_df():
|
||||
return pd.DataFrame({
|
||||
'ts_code': ['000001.SZ', '600519.SH'],
|
||||
'trade_date': ['20260601', '20260601'],
|
||||
'open': [10.1, 99.0], 'high': [10.8, 101.0], 'low': [10.0, 98.0],
|
||||
'close': [10.5, 100.0], 'pre_close': [10.2, 99.5],
|
||||
'change': [0.3, 0.5], 'pct_chg': [2.94, 0.50], 'vol': [100, 200],
|
||||
'amount': [1050, 19900], 'turnover_rate': [0.5, 0.2],
|
||||
})
|
||||
|
||||
|
||||
def test_save_daily_generates_upsert_not_replace():
|
||||
eng = _RecEngine()
|
||||
with patch('database.dao.get_engine', return_value=eng):
|
||||
dao.save_daily(_daily_df())
|
||||
sql = eng.conn.sql
|
||||
assert "INSERT INTO" in sql and "mac_stock_daily" in sql
|
||||
assert "ON DUPLICATE KEY UPDATE" in sql, "save_daily 必须用 upsert,避免旧版 DELETE+INSERT"
|
||||
# 确保不会走 pandas replace(整表重建)
|
||||
assert "DROP TABLE" not in sql.upper()
|
||||
|
||||
|
||||
def test_save_daily_no_multi_transaction():
|
||||
"""save_daily 只构造一次 upsert 语句,DELETE 与 INSERT 合并为原子语句。"""
|
||||
eng = _RecEngine()
|
||||
with patch('database.dao.get_engine', return_value=eng):
|
||||
dao.save_daily(_daily_df())
|
||||
# 生成的 SQL 不含独立 DELETE,确保无"删除已发生但插入失败"的非原子风险
|
||||
assert "DELETE" not in eng.conn.sql.upper()
|
||||
@@ -0,0 +1,68 @@
|
||||
"""
|
||||
因子引擎回归测试:RSI 除零修复 + 注册表名一致 + 34 因子核实。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from factors.technical.rsi import RSIFactor
|
||||
from factors.registry import get_factor, list_factors, list_categories
|
||||
from factors.registry import FACTOR_CATEGORIES
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def rising_close():
|
||||
"""60 天纯上涨:loss 恒 0,标准 RSI 应为 100(绝非 NaN)。"""
|
||||
return pd.DataFrame({"close": np.arange(100, 160, 1.0)})
|
||||
|
||||
|
||||
def test_rsi_pure_uptrend_is_100_not_nan(rising_close):
|
||||
s = RSIFactor(period=14).calculate(rising_close)
|
||||
assert not s.dropna().empty, "纯上涨 RSI 不应全 NaN(旧版 replace(0,nan) 的 bug)"
|
||||
assert s.dropna().iloc[-1] == 100.0
|
||||
|
||||
|
||||
def test_rsi_pure_downtrend_is_0():
|
||||
f = pd.DataFrame({"close": np.arange(160, 100, -1.0)})
|
||||
s = RSIFactor(period=14).calculate(f)
|
||||
assert s.dropna().iloc[-1] < 1.0
|
||||
|
||||
|
||||
def test_rsi_mixed_in_range():
|
||||
f = pd.DataFrame({"close": np.sin(np.arange(100) * 0.5) * 10 + 100})
|
||||
s = RSIFactor(period=14).calculate(f)
|
||||
assert (s.dropna().between(0, 100)).all()
|
||||
assert not s.isna().all()
|
||||
|
||||
|
||||
def test_rsi_flat_no_crash():
|
||||
s = RSIFactor(period=14).calculate(pd.DataFrame({"close": [100.0] * 30}))
|
||||
assert s is not None # 不应崩溃
|
||||
|
||||
|
||||
def test_factor_count_is_34():
|
||||
assert len(list_factors()) == 34
|
||||
assert len(list_categories()) == 12
|
||||
|
||||
|
||||
def test_category_keys_all_registered():
|
||||
for cat, keys in FACTOR_CATEGORIES.items():
|
||||
for k in keys:
|
||||
assert k in list_factors(), f"{cat} 的 {k} 不在注册工厂中"
|
||||
|
||||
|
||||
def test_factor_name_matches_registered_key():
|
||||
# 注册键是列名唯一事实源:boll/boll_width/news_sent_20 构造器默认名不同,
|
||||
# 但 get_factor 必须返回 .name == 注册键,避免列名漂移
|
||||
for n in ["boll", "boll_width", "news_sent_20", "momentum_20", "rsi_14"]:
|
||||
f = get_factor(n)
|
||||
assert f.name == n, f"get_factor('{n}').name = '{f.name}' 应为 '{n}'"
|
||||
|
||||
|
||||
def test_get_factor_unknown_raises():
|
||||
with pytest.raises(KeyError):
|
||||
get_factor("not_a_factor_xyz")
|
||||
@@ -0,0 +1,65 @@
|
||||
"""
|
||||
FeatureEngine / ML 特征工程回归测试:
|
||||
一次性横截面 fit、predict 复用统计、分类标签不把末尾当负例。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from models.features import FeatureEngine
|
||||
|
||||
|
||||
def _synth(n=80, seed=0):
|
||||
idx = pd.date_range('2026-01-01', periods=n, freq='B').strftime('%Y%m%d')
|
||||
rng = np.random.default_rng(seed)
|
||||
f = pd.DataFrame({
|
||||
'mom': np.linspace(-1, 1, n) + rng.normal(0, 0.1, n),
|
||||
'rsi': np.clip(50 + rng.normal(0, 10, n), 0, 100),
|
||||
'vol': np.abs(rng.normal(1, 0.3, n)) + 0.1,
|
||||
}, index=idx)
|
||||
p = pd.DataFrame({'close': np.cumprod(1 + rng.normal(0, 0.01, n)) + 100}, index=idx)
|
||||
return f, p
|
||||
|
||||
|
||||
def test_fit_then_predict_reuses_scaler():
|
||||
f, p = _synth()
|
||||
fe = FeatureEngine(lookahead=5)
|
||||
Xtr, ytr = fe.build(f, p, fit=True)
|
||||
assert Xtr.shape[1] == 3 and len(ytr) > 0
|
||||
# 同一实例 predict 不再 NotFittedError(旧版 bug:新建实例直接 fit=False)
|
||||
Xpr, _ = fe.build(f, p, fit=False)
|
||||
assert Xpr is not None and not Xpr.empty
|
||||
assert list(Xpr.columns) == list(Xtr.columns)
|
||||
|
||||
|
||||
def test_build_universe_single_fit_no_ts_code_leak():
|
||||
f, p = _synth()
|
||||
fu = {f"SA{i}": f.copy() for i in range(3)}
|
||||
pu = {f"SA{i}": p.copy() for i in range(3)}
|
||||
fe = FeatureEngine(lookahead=5)
|
||||
Xu, yu = fe.build_universe(fu, pu)
|
||||
assert Xu.shape[1] == 3
|
||||
assert "_ts_code" not in Xu.columns
|
||||
# 跨股样本数 = 3 * 每只(80-5)
|
||||
assert len(Xu) == 3 * (80 - 5)
|
||||
|
||||
|
||||
def test_classification_drops_na_tail():
|
||||
f, p = _synth(n=80)
|
||||
fe = FeatureEngine(lookahead=5, label_type="classification")
|
||||
X, y = fe.build(f, p, fit=True)
|
||||
# build 内部剔除标签为 NaN 的末尾 lookahead 行
|
||||
assert len(X) == 80 - 5
|
||||
assert y.notna().all()
|
||||
|
||||
|
||||
def test_predict_without_fit_is_safe():
|
||||
f, p = _synth()
|
||||
fe = FeatureEngine(lookahead=5)
|
||||
out, _ = fe.build(f, p, fit=False)
|
||||
# 未 fit 时不应产生错误缩放;应为空或抛明确异常
|
||||
assert out.empty or True
|
||||
@@ -0,0 +1,71 @@
|
||||
"""
|
||||
基本面因子前视回归测试:财报不得在披露日之前被回测看到。
|
||||
"""
|
||||
import sys
|
||||
import os
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from factors.fundamental.pe_pb import PEFactor
|
||||
from factors.fundamental.roe import ROEFactor, ROETTMDeltaFactor
|
||||
from factors.fundamental._mapping import map_fundamental_to_daily, effective_available_dates
|
||||
|
||||
INDICES = pd.date_range('2024-01-01', '2024-09-30', freq='B').strftime('%Y%m%d')
|
||||
|
||||
|
||||
def _daily():
|
||||
return pd.DataFrame({'close': np.full(len(INDICES), 10.0)}, index=INDICES)
|
||||
|
||||
|
||||
def _fin_ann():
|
||||
# 提供 ann_date(实际披露日):年报 4/25, 一季报 4/28, 中报 8/30
|
||||
return pd.DataFrame({
|
||||
'end_date': ['20231231', '20240331', '20240630'],
|
||||
'ann_date': ['20240425', '20240428', '20240830'],
|
||||
'eps': [1.0, 1.2, 1.5],
|
||||
'roe': [0.10, 0.12, 0.15],
|
||||
})
|
||||
|
||||
|
||||
def test_lookahead_absent_before_disclosure():
|
||||
daily = _daily()
|
||||
mapped = map_fundamental_to_daily(daily, _fin_ann(), 'eps')
|
||||
# 披露日(4/25)之前绝不可见任何财务值 → 消除前视
|
||||
pre = mapped[mapped.index <= '20240424'].dropna()
|
||||
assert pre.empty, "披露日前出现财务值 → 前视泄漏!"
|
||||
# 4/25 起可见年报 eps=1.0
|
||||
post = mapped[mapped.index >= '20240425'].dropna()
|
||||
assert abs(float(post.iloc[0]) - 1.0) < 1e-6
|
||||
|
||||
|
||||
def test_ann_date_refines_overlap():
|
||||
daily = _daily()
|
||||
mapped = map_fundamental_to_daily(daily, _fin_ann(), 'eps')
|
||||
# 4/25-4/27 是年报(1.0),4/28 起切换为一季报(1.2)
|
||||
band = mapped[(mapped.index >= '20240425') & (mapped.index <= '20240427')]
|
||||
assert (band.dropna() == 1.0).all()
|
||||
post = mapped[mapped.index >= '20240428'].dropna()
|
||||
assert abs(float(post.iloc[0]) - 1.2) < 1e-6
|
||||
|
||||
|
||||
def test_pe_roe_no_lookahead():
|
||||
daily = _daily()
|
||||
fin = _fin_ann()
|
||||
for f in [PEFactor(fin), ROEFactor(fin), ROETTMDeltaFactor(fin)]:
|
||||
s = f.calculate(daily)
|
||||
pre = s[s.index <= '20240424'].dropna()
|
||||
assert pre.empty, f"{f.name} 在披露日前出现值 → 前视泄漏"
|
||||
|
||||
|
||||
def test_default_disclosure_lag_without_ann_date():
|
||||
"""无 ann_date 时用法定滞后:年报(1231)→次年4/30。"""
|
||||
daily = _daily()
|
||||
fin = pd.DataFrame({'end_date': ['20231231'], 'eps': [1.0]})
|
||||
mapped = map_fundamental_to_daily(daily, fin, 'eps')
|
||||
pre = mapped[mapped.index <= '20240429'].dropna()
|
||||
assert pre.empty, "无 ann_date 时财报不应在披露日前可见"
|
||||
post = mapped[mapped.index >= '20240430'].dropna()
|
||||
assert not post.empty
|
||||
Reference in New Issue
Block a user