feat: 股息率案例口径 + 策略库与图表统一 + 回测存档完整化
汇总三轮未提交的开发(每轮均在本机 MariaDB + 真实浏览器上验证):
1) 股息率案例(全市场股息率最高 n 只,默认 20,每 m 月择股)
- 新增日频估值表 daily_basic + 迁移;股息率因子(dv_ratio / dividend_yield / TTM)
- 名称历史表 stock_name_history:剔除 ST 按**择股日当时名称**判定,消除
「曾高股息后 ST」的股息陷阱(实测 3.70pp 偏差)
- 区间择股/调仓双周期(m 择股 / y 调仓)、指数成分与白名单、停牌近似剔除
- 复权因子口径核对(4,164,742 行、缺失 0.0%)、收盘价成交与涨跌停拦单
- 案例实测:2020-01-01~2026-09-04 总收益 +24.86%(年化 3.52%、回撤 -28.58%)
2) 策略库与前端统一
- strategy 表 + CRUD/PUT 原地更新 + `describe_strategy` 按 spec 真实推导
「一句话说明 + 计算公式 + 执行步骤 + 注意事项」(与引擎实执行规则同源)
- 任何出现股票代码处都成对显示名称且可点击进个股页
- 全站图表基座统一 TradingView Lightweight Charts(ECharts 依赖、
锁文件、组件与文档标注一并清除),买卖点标记只落在真实交易日上
3) 回测存档完整化(可往复查看)
- 同步端点(POST /api/backtests、/api/factor-tests)此前完全不落库 → 现在同样归档,
归档 id 经响应头 X-Experiment-Id 返回(不破坏 response_model)
- data_version 首次真实写入(数据快照指纹:最新交易日 + 各表规模)
- 个股收益曲线默认**全量保存**(此前硬截断 60 只);超出体积预算才裁剪,
并写 archive_meta(机器可读)+ unimplemented(人可读)如实标注
- 列表 kind/q 过滤 + X-Total-Count(此前 limit=50 静默截断)、DELETE 归档
- 只读归档页 /experiments/{id}(Server Component,SSR 直出**选股条件**与
**交易执行依据**);结果视图按 kind 分发(backtest/factor_test/selection),
非回测归档不套用回测口径
- 新增 CLI:prune_experiments(保留策略,默认 dry-run)、
restore_experiment_from_job(从 Job 副本按原 id 重建被删的历史归档,默认 dry-run)
门禁:pytest 388 passed、ruff All checks passed、tsc 0 错误、图表单测 7 passed、
next build 成功、契约脚本 verify_strategy_workspace 59/59(含按 kind 逐类验证归档页)。
This commit is contained in:
@@ -0,0 +1,174 @@
|
||||
"""复权折算(v3 §20.5)测试:SQL 侧按 adjust_factor 折算价格列。
|
||||
|
||||
覆盖:
|
||||
- hfq:price × factor(后复权)
|
||||
- qfq:price × factor / 该股最新 factor(前复权,归一)
|
||||
- volume/amount 不折算
|
||||
- 因子缺失按 1.0 兜底,并由 count_price_adjust_gaps 如实统计
|
||||
- hfq 与 qfq 的**收益率序列完全一致**(仅差一个常数倍),这是复权口径自洽性的关键断言
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
from app.domain.entities.market import AdjustFactor, DailyBar
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
)
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
_D = [date(2024, 6, 3), date(2024, 6, 4), date(2024, 6, 5)]
|
||||
|
||||
|
||||
def _bar(day: date, close: str, symbol: str = "600519.SH") -> DailyBar:
|
||||
return DailyBar(
|
||||
symbol=symbol,
|
||||
trade_date=day,
|
||||
source="tushare",
|
||||
adjust="none",
|
||||
open=Decimal(close),
|
||||
high=Decimal(close),
|
||||
low=Decimal(close),
|
||||
close=Decimal(close),
|
||||
volume=Decimal("1000"),
|
||||
amount=Decimal("100000"),
|
||||
)
|
||||
|
||||
|
||||
def _factor(day: date, factor: str, symbol: str = "600519.SH") -> AdjustFactor:
|
||||
return AdjustFactor(symbol=symbol, trade_date=day, factor=Decimal(factor), source="tushare")
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def session(tmp_path) -> Session:
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'adjprice.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
with Session(engine) as s:
|
||||
repo = SqlAlchemyDailyBarRepository(s)
|
||||
# 100 → 除权前 1.0;110 → 之后因子 1.1(模拟一次分红/送股)
|
||||
repo.upsert_many([_bar(_D[0], "100"), _bar(_D[1], "110"), _bar(_D[2], "121")])
|
||||
s.add_all(
|
||||
[
|
||||
_factor_model(_D[0], "1.0"),
|
||||
_factor_model(_D[1], "1.1"),
|
||||
_factor_model(_D[2], "1.1"),
|
||||
]
|
||||
)
|
||||
s.commit()
|
||||
yield s
|
||||
|
||||
|
||||
def _factor_model(day: date, factor: str, symbol: str = "600519.SH"):
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import AdjustFactorModel
|
||||
|
||||
return AdjustFactorModel(symbol=symbol, trade_date=day, factor=Decimal(factor))
|
||||
|
||||
|
||||
def _stream(session, price_adjust: str, columns=("close",)) -> dict[date, tuple]:
|
||||
repo = SqlAlchemyDailyBarRepository(session)
|
||||
rows = list(
|
||||
repo.stream_range_many_columns(
|
||||
["600519.SH"], _D[0], _D[2], list(columns),
|
||||
adjust="none", price_adjust=price_adjust,
|
||||
)
|
||||
)
|
||||
return {date.fromisoformat(r[1]): r[2:] for r in rows}
|
||||
|
||||
|
||||
class TestAdjustedPrices:
|
||||
def test_hfq_multiplies_by_factor(self, session) -> None:
|
||||
out = _stream(session, "hfq")
|
||||
assert out[_D[0]] == (100.0,) # 100 × 1.0
|
||||
assert out[_D[1]] == pytest.approx((121.0,)) # 110 × 1.1
|
||||
assert out[_D[2]] == pytest.approx((133.1,)) # 121 × 1.1
|
||||
|
||||
def test_qfq_normalizes_by_latest_factor(self, session) -> None:
|
||||
out = _stream(session, "qfq")
|
||||
assert out[_D[0]] == pytest.approx((100 / 1.1,)) # 100 × 1.0 / 1.1
|
||||
assert out[_D[1]] == pytest.approx((110.0,)) # 110 × 1.1 / 1.1
|
||||
assert out[_D[2]] == pytest.approx((121.0,))
|
||||
|
||||
def test_volume_and_amount_not_adjusted(self, session) -> None:
|
||||
out = _stream(session, "hfq", columns=("close", "volume", "amount"))
|
||||
close, volume, amount = out[_D[2]]
|
||||
assert close == pytest.approx(133.1)
|
||||
assert volume == 1000.0 and amount == 100000.0 # 不随复权缩放
|
||||
|
||||
def test_none_is_raw(self, session) -> None:
|
||||
out = _stream(session, "none")
|
||||
assert out[_D[0]] == (100.0,) and out[_D[2]] == (121.0,)
|
||||
|
||||
def test_hfq_and_qfq_give_identical_returns(self, session) -> None:
|
||||
hfq, qfq = _stream(session, "hfq"), _stream(session, "qfq")
|
||||
for a, b in zip(_D, _D[1:], strict=False):
|
||||
r_hfq = hfq[b][0] / hfq[a][0]
|
||||
r_qfq = qfq[b][0] / qfq[a][0]
|
||||
assert r_hfq == pytest.approx(r_qfq, rel=1e-12)
|
||||
|
||||
def test_invalid_price_adjust_rejected(self, session) -> None:
|
||||
repo = SqlAlchemyDailyBarRepository(session)
|
||||
with pytest.raises(ValueError, match="复权"):
|
||||
list(
|
||||
repo.stream_range_many_columns(
|
||||
["600519.SH"], _D[0], _D[2], ["close"], price_adjust="bad"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class TestAdjustGapReporting:
|
||||
def test_missing_factor_falls_back_and_is_reported(self, session) -> None:
|
||||
repo = SqlAlchemyDailyBarRepository(session)
|
||||
# 追加一行无因子的行情(模拟 25 只蓝筹 2020-2022 缺因子的情形)
|
||||
repo.upsert_many([_bar(date(2024, 6, 6), "200")])
|
||||
session.commit()
|
||||
out = dict(
|
||||
(date.fromisoformat(r[1]), r[2])
|
||||
for r in repo.stream_range_many_columns(
|
||||
["600519.SH"], _D[0], date(2024, 6, 6), ["close"], price_adjust="hfq"
|
||||
)
|
||||
)
|
||||
assert out[date(2024, 6, 6)] == 200.0 # 缺因子 → 系数 1.0(未折算)
|
||||
total, missing = repo.count_price_adjust_gaps(
|
||||
["600519.SH"], _D[0], date(2024, 6, 6)
|
||||
)
|
||||
assert total == 4 and missing == 1
|
||||
|
||||
def test_no_gap_when_all_covered(self, session) -> None:
|
||||
repo = SqlAlchemyDailyBarRepository(session)
|
||||
total, missing = repo.count_price_adjust_gaps(["600519.SH"], _D[0], _D[2])
|
||||
assert total == 3 and missing == 0
|
||||
|
||||
def test_empty_symbols(self, session) -> None:
|
||||
repo = SqlAlchemyDailyBarRepository(session)
|
||||
assert repo.count_price_adjust_gaps([], _D[0], _D[2]) == (0, 0)
|
||||
|
||||
class TestQfqMissingFactor:
|
||||
"""qfq 缺因子行不得被错误缩放。
|
||||
|
||||
回归用例:`coalesce(factor,1)/latest` 会把缺口行缩放到 `1/latest`
|
||||
(latest=5 → 121/5=24.2,凭空 −80% 单日跌幅),正确写法必须把 COALESCE
|
||||
放在最外层,使缺口行保持原始价(口径:缺失因子按 1.0 兜底、不折算)。
|
||||
"""
|
||||
|
||||
def test_qfq_gap_row_is_not_rescaled(self, session) -> None:
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import AdjustFactorModel
|
||||
|
||||
session.query(AdjustFactorModel).delete() # 清掉 fixture 的 1.0/1.1/1.1
|
||||
session.add_all([_factor_model(_D[0], "5.0"), _factor_model(_D[1], "5.0")])
|
||||
session.commit() # 第 3 天无因子 → 缺口
|
||||
|
||||
out = _stream(session, "qfq")
|
||||
# 有因子的行:price × factor / max(factor=5.0) → 等于原始价
|
||||
assert out[_D[0]] == pytest.approx((100.0,))
|
||||
assert out[_D[1]] == pytest.approx((110.0,))
|
||||
# 缺口行:保持原始价(若被错误归一则为 121/5 = 24.2)
|
||||
assert out[_D[2]] == pytest.approx((121.0,))
|
||||
|
||||
repo = SqlAlchemyDailyBarRepository(session)
|
||||
total, missing = repo.count_price_adjust_gaps(["600519.SH"], _D[0], _D[2])
|
||||
assert (total, missing) == (3, 1)
|
||||
@@ -6,7 +6,9 @@ from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
from app.core.config import PROJECT_ROOT, _load_env_file, get_settings
|
||||
import pytest
|
||||
import yaml
|
||||
from app.core.config import CONFIG_PATH, PROJECT_ROOT, _load_env_file, get_settings
|
||||
|
||||
|
||||
def test_project_root_points_to_repo_root() -> None:
|
||||
@@ -22,15 +24,24 @@ def test_test_runner_isolation_uses_tmp_sqlite() -> None:
|
||||
|
||||
|
||||
def test_mysql_default_from_config_yaml(monkeypatch) -> None:
|
||||
"""未设 DATABASE_URL 时,config.yaml database.mysql 段组装 MySQL URL(不连接)。"""
|
||||
"""未设 DATABASE_URL 时,config.yaml database.mysql 段组装 MySQL URL(不连接)。
|
||||
|
||||
期望值直接取自 config.yaml,而不硬编码部署地址:切库(192.168.1.10 → 127.0.0.1)
|
||||
属于配置变更,不应让本用例失效。
|
||||
"""
|
||||
mysql = yaml.safe_load(CONFIG_PATH.read_text(encoding="utf-8"))["database"]["mysql"]
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
get_settings.cache_clear()
|
||||
try:
|
||||
url = get_settings().database_url
|
||||
finally:
|
||||
get_settings.cache_clear()
|
||||
assert url.startswith("mysql+pymysql://qlib:")
|
||||
assert "192.168.1.10:3306/qlib?charset=utf8mb4" in url
|
||||
assert url.startswith(f"mysql+pymysql://{mysql['user']}:")
|
||||
expected = (
|
||||
f"{mysql['host']}:{mysql['port']}/{mysql['db']}?charset={mysql['charset']}"
|
||||
)
|
||||
# 断言只比较 @ 之后的部分:失败时 pytest 不会把含密码的 userinfo 打进日志
|
||||
assert expected in url.split("@")[-1]
|
||||
|
||||
|
||||
def test_build_mysql_url(monkeypatch) -> None:
|
||||
@@ -117,3 +128,95 @@ def test_llm_env_overrides_yaml(monkeypatch) -> None:
|
||||
assert s.llm_model == "env-model"
|
||||
assert s.llm_base_url == "https://example.com/v1"
|
||||
assert s.llm_api_key == "sk-test"
|
||||
|
||||
|
||||
class TestForbiddenDbTarget:
|
||||
"""用户约束:**禁止使用 192.168.1.10 作为 DB 目标**(只允许本机 MariaDB)。
|
||||
|
||||
连错库是「最危险的静默错误」:不报错、界面正常,只是数据在另一台机器上读写。
|
||||
因此这里是硬失败(抛错),不是警告;本类把守卫生效与放行口子都钉住。
|
||||
"""
|
||||
|
||||
def _clear(self):
|
||||
get_settings.cache_clear()
|
||||
|
||||
def test_forbidden_host_rejected_from_env_url(self, monkeypatch) -> None:
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL", "mysql+pymysql://qlib:pw@192.168.1.10:3306/qlib"
|
||||
)
|
||||
monkeypatch.delenv("QLIB_FORBIDDEN_DB_HOSTS", raising=False)
|
||||
self._clear()
|
||||
try:
|
||||
with pytest.raises(RuntimeError, match="192.168.1.10"):
|
||||
get_settings()
|
||||
finally:
|
||||
self._clear()
|
||||
|
||||
def test_forbidden_host_rejected_from_config_yaml(self, monkeypatch, tmp_path) -> None:
|
||||
"""即使走 config.yaml 的 mysql 段(而非 env)也必须拦住。"""
|
||||
import app.core.config as config_mod
|
||||
|
||||
cfg = {
|
||||
"database": {
|
||||
"mysql": {
|
||||
"enabled": True,
|
||||
"host": "192.168.1.10",
|
||||
"port": 3306,
|
||||
"db": "qlib",
|
||||
"user": "qlib",
|
||||
"charset": "utf8mb4",
|
||||
}
|
||||
}
|
||||
}
|
||||
path = tmp_path / "config.yaml"
|
||||
path.write_text(yaml.safe_dump(cfg, allow_unicode=True), encoding="utf-8")
|
||||
monkeypatch.delenv("DATABASE_URL", raising=False)
|
||||
monkeypatch.delenv("QLIB_FORBIDDEN_DB_HOSTS", raising=False)
|
||||
monkeypatch.setattr(config_mod, "CONFIG_PATH", path)
|
||||
self._clear()
|
||||
try:
|
||||
with pytest.raises(RuntimeError, match="已被禁止"):
|
||||
get_settings()
|
||||
finally:
|
||||
self._clear()
|
||||
|
||||
def test_local_host_allowed(self, monkeypatch) -> None:
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL", "mysql+pymysql://qlib:pw@127.0.0.1:3306/qlib"
|
||||
)
|
||||
self._clear()
|
||||
try:
|
||||
assert get_settings().database_url.startswith("mysql+pymysql://qlib")
|
||||
finally:
|
||||
self._clear()
|
||||
|
||||
def test_forbidden_list_is_env_overridable(self, monkeypatch) -> None:
|
||||
"""显式放开的口子必须存在(空串=不限制),否则紧急情况下无法操作。"""
|
||||
monkeypatch.setenv(
|
||||
"DATABASE_URL", "mysql+pymysql://qlib:pw@192.168.1.10:3306/qlib"
|
||||
)
|
||||
monkeypatch.setenv("QLIB_FORBIDDEN_DB_HOSTS", "")
|
||||
self._clear()
|
||||
try:
|
||||
assert "192.168.1.10" in get_settings().database_url
|
||||
finally:
|
||||
self._clear()
|
||||
|
||||
def test_sqlite_and_other_hosts_unaffected(self, monkeypatch) -> None:
|
||||
"""SQLite(无 host)与其它主机不受影响 —— 守卫只拦精确匹配的禁用主机。"""
|
||||
from app.core.config import assert_db_target_allowed
|
||||
|
||||
monkeypatch.delenv("QLIB_FORBIDDEN_DB_HOSTS", raising=False)
|
||||
assert_db_target_allowed("sqlite:////tmp/x.db")
|
||||
assert_db_target_allowed("mysql+pymysql://u:p@127.0.0.1:3306/qlib")
|
||||
assert_db_target_allowed("mysql+pymysql://u:p@db.internal:3306/qlib")
|
||||
assert_db_target_allowed("")
|
||||
with pytest.raises(RuntimeError):
|
||||
assert_db_target_allowed("mysql+pymysql://u:p@192.168.1.10:3306/qlib")
|
||||
|
||||
def test_guard_only_matches_host_not_userinfo(self, monkeypatch) -> None:
|
||||
"""只比 host:把禁用 IP 写进用户名/密码不应被误判(避免误报挡住正常连接)。"""
|
||||
from app.core.config import assert_db_target_allowed
|
||||
|
||||
monkeypatch.delenv("QLIB_FORBIDDEN_DB_HOSTS", raising=False)
|
||||
assert_db_target_allowed("mysql+pymysql://192.168.1.10:pw@127.0.0.1:3306/qlib")
|
||||
|
||||
@@ -0,0 +1,244 @@
|
||||
"""用户案例端到端集成测试(合成数据):高股息 Top n → 持仓前 x,m/y 双周期。
|
||||
|
||||
不连真库,用确定性合成行情 + 内存 SQLite(覆盖 daily_basic 与 adjust_factor
|
||||
两条新链路),验证:
|
||||
- dv_ratio 因子参与评分(dividend_yield)
|
||||
- conditions(dv_ratio <= 30)真正过滤掉特殊分红股
|
||||
- hfq 复权生效(除权日不再被计为亏损)
|
||||
- m/y 双周期、x<=n、顺延买入的端到端组合行为
|
||||
- 最低佣金 min_commission
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import (
|
||||
AdjustFactorModel,
|
||||
DailyBasicModel,
|
||||
StockDailyModel,
|
||||
StockModel,
|
||||
)
|
||||
from app.quant.engine import LocalEngine
|
||||
from app.quant.service import ResearchService
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH"]
|
||||
_START = date(2024, 1, 1)
|
||||
|
||||
|
||||
def _seed(session: Session, *, spike_date: date | None = None) -> None:
|
||||
"""4 只股票:A/B 高股息,C 中股息,D 低股息;除权日因子在 2024-06-03 跳到 1.1。
|
||||
|
||||
`spike_date`:把 A 股在该日的 dv_ratio 抬到 45%(特殊分红尖峰,
|
||||
用于验证 `dv_ratio <= 30` 条件确实把它剔除)。
|
||||
"""
|
||||
session.add_all(
|
||||
[
|
||||
StockModel(
|
||||
symbol=s, name=f"股票{s[:6]}", industry="银行" if i < 2 else "制造",
|
||||
market="主板", area="深圳", list_date=date(2000, 1, 1), status="L",
|
||||
)
|
||||
for i, s in enumerate(_SYMS)
|
||||
]
|
||||
)
|
||||
# 股息率:A=8% B=6% C=4% D=2%;special 时 A 在首个交易日出现 45% 的异常高值
|
||||
dv = {"600000.SH": 8.0, "600001.SH": 6.0, "600002.SH": 4.0, "600003.SH": 2.0}
|
||||
dates = pd.bdate_range(_START, periods=130)
|
||||
bars, basics, factors = [], [], []
|
||||
for i, sym in enumerate(_SYMS):
|
||||
price = 10.0 + i
|
||||
for d in dates:
|
||||
price = price * (1 + 0.0008 + 0.0003 * i)
|
||||
bars.append(
|
||||
StockDailyModel(
|
||||
symbol=sym, trade_date=d.date(), source="tushare", adjust="none",
|
||||
open=Decimal(str(price)), high=Decimal(str(price)),
|
||||
low=Decimal(str(price)), close=Decimal(str(price)),
|
||||
volume=Decimal("1000000"), amount=Decimal(str(price * 1e6)),
|
||||
)
|
||||
)
|
||||
rate = dv[sym]
|
||||
if spike_date is not None and sym == "600000.SH" and d.date() == spike_date:
|
||||
rate = 45.0 # 特殊分红尖峰(应按条件在该择股日被剔除)
|
||||
basics.append(
|
||||
DailyBasicModel(
|
||||
symbol=sym, trade_date=d.date(), source="tushare",
|
||||
close=Decimal(str(price)), dv_ratio=Decimal(str(rate)),
|
||||
dv_ttm=Decimal(str(rate)), pe=Decimal("8"), pb=Decimal("1"),
|
||||
total_mv=Decimal("1e11"),
|
||||
)
|
||||
)
|
||||
# 2024-06-03 起因子 1.1(模拟一次除权):hfq 价格应整体上移 10%
|
||||
f = 1.1 if d.date() >= date(2024, 6, 3) else 1.0
|
||||
factors.append(
|
||||
AdjustFactorModel(symbol=sym, trade_date=d.date(), factor=Decimal(str(f)))
|
||||
)
|
||||
session.add_all(bars + basics + factors)
|
||||
session.commit()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def service(tmp_path) -> ResearchService:
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'case.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
_seed(session, spike_date=date(2024, 3, 1))
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyDailyBasicRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
|
||||
# 复权折算发生在 SqlAlchemyDailyBarRepository 的 SQL 内(price_adjust=...),
|
||||
# 业务层无需 adjust_factor 仓储
|
||||
yield ResearchService(
|
||||
SqlAlchemyStockRepository(session),
|
||||
SqlAlchemyDailyBarRepository(session),
|
||||
LocalEngine(),
|
||||
basic_repo=SqlAlchemyDailyBasicRepository(session),
|
||||
)
|
||||
session.close()
|
||||
|
||||
|
||||
def _spec(**kw):
|
||||
from app.domain.entities.research import (
|
||||
ConditionSpec,
|
||||
CostSpec,
|
||||
FactorSpec,
|
||||
ResearchSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
|
||||
base = dict(
|
||||
type="backtest",
|
||||
universe=UniverseSpec(exclude_st=False, min_listing_days=0),
|
||||
price_adjustment="hfq",
|
||||
factors=[FactorSpec(name="dividend_yield")],
|
||||
conditions=[ConditionSpec(field="dv_ratio", op="lte", value=30)],
|
||||
selection=SelectionSpec(
|
||||
top_n=3, hold_top_x=2, allow_substitute=False, defer_buy=True
|
||||
),
|
||||
rebalance="monthly",
|
||||
selection_interval_months=3,
|
||||
rebalance_interval_months=3,
|
||||
period=(date(2024, 3, 1), date(2024, 6, 28)),
|
||||
costs=CostSpec(
|
||||
commission_rate=0.0003, stamp_tax_rate=0.0005, slippage_rate=0.001,
|
||||
min_commission=5.0,
|
||||
),
|
||||
)
|
||||
base.update(kw)
|
||||
return ResearchSpec(**base)
|
||||
|
||||
|
||||
class TestUserCaseEndToEnd:
|
||||
def test_backtest_runs_with_all_new_knobs(self, service) -> None:
|
||||
result = service.run_backtest(_spec())
|
||||
assert result.summary.total_trades >= 1
|
||||
# 配置可溯源(v3 §20.5)
|
||||
snap = result.config_snapshot
|
||||
assert snap["price_adjustment"] == "hfq"
|
||||
assert snap["price_basis"]["execution_price_basis"] == "close_adj"
|
||||
assert snap["selection_interval_months"] == 3
|
||||
assert snap["selection"]["hold_top_x"] == 2
|
||||
assert result.symbol_curves, "应输出个股收益曲线"
|
||||
|
||||
def test_condition_excludes_special_dividend_spike(self, service) -> None:
|
||||
"""A 股在 2024-03-01 的 dv_ratio=45% > 30% → 该日不得进入候选池;
|
||||
其它择股日 A 股息率仍最高(8%)→ 应正常入选(证明过滤是按日求值而非一刀切)。"""
|
||||
result = service.run_backtest(
|
||||
_spec(period=(date(2024, 3, 1), date(2024, 6, 28)),
|
||||
selection_interval_months=3, rebalance_interval_months=3)
|
||||
)
|
||||
by_day: dict = {}
|
||||
for p in result.selection_history:
|
||||
by_day.setdefault(p.date, []).append(p.symbol)
|
||||
assert date(2024, 3, 1) in by_day
|
||||
assert "600000.SH" not in by_day[date(2024, 3, 1)], "尖峰日应按条件剔除"
|
||||
later = [d for d in by_day if d > date(2024, 3, 1)]
|
||||
assert later, "应存在后续择股日"
|
||||
assert all("600000.SH" in by_day[d] for d in later), "尖峰解除后应恢复入选"
|
||||
|
||||
def test_top_x_cap_and_pool_recorded(self, service) -> None:
|
||||
result = service.run_backtest(_spec())
|
||||
# 候选池 = n = 3(4 只中剔除 1 只)
|
||||
per_day: dict = {}
|
||||
for p in result.selection_history:
|
||||
per_day.setdefault(p.date, []).append(p.symbol)
|
||||
assert per_day and all(len(v) <= 3 for v in per_day.values())
|
||||
# 持仓 ≤ x = 2
|
||||
held: dict = {}
|
||||
for pos in result.positions:
|
||||
held.setdefault(pos.date, []).append(pos.symbol)
|
||||
assert held and all(len(v) <= 2 for v in held.values())
|
||||
|
||||
def test_hfq_reflects_dividend_jump(self, service) -> None:
|
||||
"""hfq 下 2024-06-03 不复权价与复权价之比应为 1.1(无横截面跳变计入收益)。"""
|
||||
from app.quant.service import load_daily_df
|
||||
|
||||
daily = load_daily_df(
|
||||
service._daily_repo, ["600000.SH"], date(2024, 5, 31), date(2024, 6, 4),
|
||||
["close"], price_adjust="hfq",
|
||||
)
|
||||
raw = load_daily_df(
|
||||
service._daily_repo, ["600000.SH"], date(2024, 5, 31), date(2024, 6, 4),
|
||||
["close"], price_adjust="none",
|
||||
)
|
||||
adj = daily.sort_values("trade_date")["close"].tolist()
|
||||
rawp = raw.sort_values("trade_date")["close"].tolist()
|
||||
# 6/3 之前因子 1.0;之后 1.1 → 复权价整体上移
|
||||
assert adj[0] == pytest.approx(rawp[0])
|
||||
assert adj[-1] == pytest.approx(rawp[-1] * 1.1, rel=1e-9)
|
||||
|
||||
def test_min_commission_increases_cost(self, service) -> None:
|
||||
"""小资金 + x=2 → 单笔约 1 万元,佣金 3 元低于最低 5 元 → 成本上升。"""
|
||||
free = service.run_backtest(
|
||||
_spec(
|
||||
initial_capital=20_000.0,
|
||||
costs=_spec().costs.model_copy(update={"min_commission": 0.0}),
|
||||
)
|
||||
)
|
||||
costly = service.run_backtest(_spec(initial_capital=20_000.0)) # min_commission=5.0
|
||||
assert costly.summary.final_equity < free.summary.final_equity
|
||||
|
||||
def test_unimplemented_notes_surfaced(self, service) -> None:
|
||||
result = service.run_backtest(_spec())
|
||||
joined = " | ".join(result.unimplemented)
|
||||
assert "后复权" in joined
|
||||
assert "顺延买入" in joined
|
||||
assert "退市" in joined # 幸存者偏差显式标注
|
||||
|
||||
|
||||
class TestSelectionBacktestConsistencyWithConditions:
|
||||
"""v3 §28:回测择股日候选池 == 同日 /api/selections(同一求值器)。"""
|
||||
|
||||
def test_pool_matches_selection_service(self, service) -> None:
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
|
||||
result = service.run_backtest(
|
||||
_spec(selection_interval_months=3, rebalance_interval_months=3)
|
||||
)
|
||||
sel_service = SelectionService(
|
||||
service._stock_repo, service._daily_repo, basic_repo=service._basic_repo
|
||||
)
|
||||
first_day = min(p.date for p in result.selection_history)
|
||||
query = SelectionQuery(
|
||||
method="condition",
|
||||
price_adjustment="hfq",
|
||||
conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}],
|
||||
as_of=first_day,
|
||||
)
|
||||
sel = sel_service.select(query)
|
||||
eligible = {c.symbol for c in sel.candidates}
|
||||
# 回测当日候选池(记为条件通过者中的前 n)必是条件合格集合的子集
|
||||
pool = {p.symbol for p in result.selection_history if p.date == first_day}
|
||||
assert pool <= eligible, f"候选池 {pool} 不在条件合格集 {eligible} 内"
|
||||
assert eligible, "条件合格集不应为空(B/C/D 股 dv_ratio ≤ 30)"
|
||||
@@ -0,0 +1,653 @@
|
||||
"""Experiment 归档补齐(2026-09)测试。
|
||||
|
||||
覆盖五件事:
|
||||
1. **同步端点归档**:`POST /api/backtests` / `POST /api/factor-tests` 落库(此前只写
|
||||
进程内存),并用响应头 `X-Experiment-Id` 返回归档 id;归档失败时仍返回计算结果
|
||||
并用 `X-Archive-Error` 如实暴露(AGENT §7 不静默)。
|
||||
2. **data_version**:真实数据指纹(交易日 + 各表行数),格式与列宽(varchar(40))合规。
|
||||
3. **列表过滤/分页**:`kind` / `q` / `limit` / `offset` + `X-Total-Count`(不静默截断)。
|
||||
4. **DELETE**:200 → 再 GET 404 → 再 DELETE 404,且不触碰 job 记录。
|
||||
5. **归档完整性**:默认不再按 60 只截断个股曲线;超体积预算时裁剪并留下
|
||||
`archive_meta`(机器可读)+ `unimplemented`(人可读)证据;
|
||||
完整结果只存 experiment 一份(`GET /api/jobs/{id}` 契约不变,老记录可回退)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from datetime import date, datetime
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from app.api import deps
|
||||
from app.application.services import experiment_archive as ea
|
||||
from app.domain.entities.market import Stock
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
CostSpec,
|
||||
ExperimentRecord,
|
||||
ExperimentSummary,
|
||||
FactorSpec,
|
||||
JobRecord,
|
||||
JobStatus,
|
||||
ResearchSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.models.jobs import ExperimentModel, JobModel
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import StockDailyModel
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
from app.main import app
|
||||
from app.quant.engine import LocalEngine
|
||||
from app.quant.service import ResearchService
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import create_engine, select
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
|
||||
|
||||
_SYMS = ["60000" + str(i) + ".SH" for i in range(5)]
|
||||
|
||||
_BACKTEST_BODY = {
|
||||
"type": "backtest",
|
||||
"universe": {"exclude_st": False, "min_listing_days": 0},
|
||||
"factors": [{"name": "momentum_20", "weight": 1.0}],
|
||||
"selection": {"top_n": 2},
|
||||
"rebalance": "monthly",
|
||||
"period": ["2024-03-01", "2024-10-31"],
|
||||
}
|
||||
|
||||
# data_version 格式:d<YYYYMMDD|->,后跟可选的行数段(MySQL 近似值带 ≈ 与 k)
|
||||
_DATA_VERSION_RE = re.compile(
|
||||
r"^d(\d{8}|-)(;n(≈\d+k|\d+))?(;a(≈\d+k|\d+))?(;b(≈\d+k|\d+))?$"
|
||||
)
|
||||
|
||||
|
||||
class _MemStockRepo:
|
||||
def __init__(self, stocks: list[Stock]) -> None:
|
||||
self._stocks = stocks
|
||||
|
||||
def get_by_symbol(self, symbol: str) -> Stock | None:
|
||||
return next((s for s in self._stocks if s.symbol == symbol), None)
|
||||
|
||||
def list(self) -> list[Stock]:
|
||||
return self._stocks
|
||||
|
||||
|
||||
class _MemDailyRepo:
|
||||
def __init__(self, bars) -> None:
|
||||
self._bars = bars
|
||||
|
||||
def get_range_many(self, symbols, start, end, adjust="none"):
|
||||
return [b for b in self._bars if b.symbol in symbols and start <= b.trade_date <= end]
|
||||
|
||||
def get_range(self, symbol, start, end):
|
||||
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def api(tmp_path):
|
||||
"""TestClient + 临时 SQLite(job/experiment/行情表)+ 内存行情仓储。"""
|
||||
drifts = {sym: 0.003 - 0.0015 * i for i, sym in enumerate(_SYMS)}
|
||||
bars = bars_dataframe_to_daily_bars(synthetic_daily(drifts, n=300))
|
||||
stocks = [
|
||||
Stock(symbol=sym, name=f"测试股份{i}", list_date=date(1999, 11, 10))
|
||||
for i, sym in enumerate(_SYMS)
|
||||
]
|
||||
service = ResearchService(_MemStockRepo(stocks), _MemDailyRepo(bars), LocalEngine())
|
||||
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'archive.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
||||
|
||||
def _session_override():
|
||||
with Session() as s:
|
||||
yield s
|
||||
|
||||
app.dependency_overrides[deps.get_session] = _session_override
|
||||
app.dependency_overrides[deps._service_factory] = lambda: service
|
||||
app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(stocks)
|
||||
with TestClient(app) as client:
|
||||
yield client, Session
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
def _exp_by_id(Session, exp_id: str) -> ExperimentRecord | None:
|
||||
with Session() as session:
|
||||
return SqlAlchemyExperimentRepository(session).get(exp_id)
|
||||
|
||||
|
||||
class TestSyncEndpointsArchive:
|
||||
"""任务 1:同步端点必须落库,并通过 X-Experiment-Id 暴露归档 id。"""
|
||||
|
||||
def test_backtest_archives_to_db_and_returns_header(self, api) -> None:
|
||||
client, Session = api
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["equity_curve"], "body 形状不变:仍是完整 BacktestResult"
|
||||
|
||||
exp_id = resp.headers.get("X-Experiment-Id")
|
||||
assert exp_id and exp_id.startswith("EXP-"), "必须用响应头返回归档 id"
|
||||
exp = _exp_by_id(Session, exp_id)
|
||||
assert exp is not None, "同步回测必须在 experiment 表落库(不再只写进程内存)"
|
||||
assert exp.kind == "backtest"
|
||||
assert json.loads(exp.spec_json)["factors"][0]["name"] == "momentum_20"
|
||||
assert exp.summary_text and "总收益" in exp.summary_text
|
||||
assert exp.code_version, "归档必须带代码版本"
|
||||
assert exp.job_id is None, "同步端点没有 Job,job_id 应为 None"
|
||||
# 归档结果可解码为完整 BacktestResult
|
||||
archived = BacktestResult.model_validate_json(exp.result_json)
|
||||
assert archived.summary.total_return_pct == body["summary"]["total_return_pct"]
|
||||
|
||||
# 详情接口能读回
|
||||
detail = client.get(f"/api/experiments/{exp_id}")
|
||||
assert detail.status_code == 200
|
||||
assert detail.json()["result"]["summary"]["total_return_pct"] == (
|
||||
body["summary"]["total_return_pct"]
|
||||
)
|
||||
|
||||
def test_factor_test_archives_with_kind(self, api) -> None:
|
||||
client, Session = api
|
||||
payload = dict(_BACKTEST_BODY, type="factor_test")
|
||||
resp = client.post("/api/factor-tests", json=payload)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["factor_name"] == "momentum_20"
|
||||
|
||||
exp_id = resp.headers.get("X-Experiment-Id")
|
||||
assert exp_id
|
||||
exp = _exp_by_id(Session, exp_id)
|
||||
assert exp is not None
|
||||
assert exp.kind == "factor_test"
|
||||
assert exp.result_json and json.loads(exp.result_json)["factor_name"] == "momentum_20"
|
||||
|
||||
def test_archive_failure_keeps_result_and_exposes_reason(self, api, monkeypatch) -> None:
|
||||
"""归档失败不得吞掉计算结果,且失败必须如实暴露(AGENT §7 / §24)。"""
|
||||
client, _Session = api
|
||||
|
||||
def _boom(self, experiment): # noqa: ANN001
|
||||
raise RuntimeError("db down")
|
||||
|
||||
monkeypatch.setattr(SqlAlchemyExperimentRepository, "save", _boom)
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
assert resp.status_code == 200, "归档失败也必须返回已算出的结果"
|
||||
assert resp.json()["equity_curve"], "结果本身必须完整返回"
|
||||
assert "X-Experiment-Id" not in resp.headers
|
||||
err = resp.headers.get("X-Archive-Error")
|
||||
assert err and "RuntimeError" in err and "db down" in err, (
|
||||
f"归档失败原因必须如实暴露,实际:{err!r}"
|
||||
)
|
||||
|
||||
|
||||
class TestDataVersion:
|
||||
"""任务 2:data_version 必须是真实、可比较、且塞得进 varchar(40) 的指纹。"""
|
||||
|
||||
def test_format_and_column_width(self, api) -> None:
|
||||
client, Session = api
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
exp = _exp_by_id(Session, resp.headers["X-Experiment-Id"])
|
||||
assert exp is not None
|
||||
assert exp.data_version, "data_version 不允许为空(此前永远是 NULL)"
|
||||
assert _DATA_VERSION_RE.match(exp.data_version), exp.data_version
|
||||
assert len(exp.data_version) <= 40, "experiment.data_version 是 varchar(40)"
|
||||
|
||||
def test_real_values_from_seeded_bars(self, api) -> None:
|
||||
"""种子 3 行日线(SQLite 走精确 COUNT(*))→ 指纹必须是精确的真实值。"""
|
||||
client, Session = api
|
||||
days = [date(2024, 9, 26), date(2024, 9, 27), date(2024, 9, 30)]
|
||||
with Session() as session:
|
||||
session.execute(
|
||||
StockDailyModel.__table__.insert(),
|
||||
[
|
||||
{
|
||||
"symbol": f"60010{i}.SH",
|
||||
"trade_date": d,
|
||||
"source": "tushare",
|
||||
"adjust": "none",
|
||||
"open": 1.0,
|
||||
"high": 1.0,
|
||||
"low": 1.0,
|
||||
"close": 1.0,
|
||||
"volume": 1.0,
|
||||
"amount": 1.0,
|
||||
}
|
||||
for i, d in enumerate(days)
|
||||
],
|
||||
)
|
||||
session.commit()
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
exp = _exp_by_id(Session, resp.headers["X-Experiment-Id"])
|
||||
assert exp is not None
|
||||
assert exp.data_version == "d20240930;n3;a0;b0", exp.data_version
|
||||
|
||||
def test_unavailable_when_tables_missing(self, tmp_path) -> None:
|
||||
"""取不到数据口径时如实降级为 unavailable,绝不编造数字。"""
|
||||
from app.infrastructure.persistence.sqlalchemy.data_version import compute_data_version
|
||||
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'empty.db'}", future=True)
|
||||
with sessionmaker(bind=engine)() as session:
|
||||
assert compute_data_version(session) == "unavailable"
|
||||
|
||||
|
||||
class TestListFilters:
|
||||
"""任务 3:列表过滤/分页在 SQL 层完成,并用 X-Total-Count 暴露总数。"""
|
||||
|
||||
@staticmethod
|
||||
def _seed(Session) -> list[str]:
|
||||
"""直接造 4 条归档:2 个 backtest(其一含 DIVIDEND_YIELD 因子)、1 个
|
||||
factor_test、1 个 selection。"""
|
||||
rows = [
|
||||
("EXP-A1", "backtest", "DIVIDEND_YIELD", "总收益 12.00% · 年化 8.00% · 回撤 5.00%"),
|
||||
("EXP-A2", "backtest", "momentum_20", "总收益 3.00% · 年化 2.00% · 回撤 1.00%"),
|
||||
("EXP-A3", "factor_test", "momentum_20", "IC 0.0100 · RankIC 0.0200 · 样本 100 日"),
|
||||
("EXP-A4", "selection", "momentum_60", "as_of 2024-10-31 · 选出 3 / 评估 100"),
|
||||
]
|
||||
with Session() as session:
|
||||
repo = SqlAlchemyExperimentRepository(session)
|
||||
for i, (exp_id, kind, factor, summary) in enumerate(rows):
|
||||
repo.save(
|
||||
ExperimentRecord(
|
||||
id=exp_id,
|
||||
kind=kind,
|
||||
spec_json=json.dumps(
|
||||
{
|
||||
"factors": [{"name": factor, "weight": 1.0}],
|
||||
"period": ["2024-03-01", "2024-10-31"],
|
||||
"rebalance": "monthly",
|
||||
"selection": {"top_n": i + 1},
|
||||
}
|
||||
),
|
||||
result_json=json.dumps({"kind": kind, "n": i}),
|
||||
summary_text=summary,
|
||||
code_version="abc1234",
|
||||
data_version="d20260904;n1;a2;b3",
|
||||
created_at=datetime(2026, 9, 4, 10, i, 0),
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
return [r[0] for r in rows]
|
||||
|
||||
def test_kind_filter_and_total_count(self, api) -> None:
|
||||
client, Session = api
|
||||
self._seed(Session)
|
||||
resp = client.get("/api/experiments", params={"kind": "backtest"})
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert {e["id"] for e in body} == {"EXP-A1", "EXP-A2"}
|
||||
assert resp.headers["X-Total-Count"] == "2"
|
||||
|
||||
def test_q_is_case_insensitive_over_factor_name_and_summary(self, api) -> None:
|
||||
client, Session = api
|
||||
self._seed(Session)
|
||||
# 因子名(在 spec_json 里,大小写不敏感)
|
||||
r1 = client.get("/api/experiments", params={"q": "dividend_yield"})
|
||||
assert [e["id"] for e in r1.json()] == ["EXP-A1"]
|
||||
assert r1.headers["X-Total-Count"] == "1"
|
||||
# id 片段
|
||||
r2 = client.get("/api/experiments", params={"q": "exp-a"})
|
||||
assert r2.headers["X-Total-Count"] == "4"
|
||||
# summary_text(中文)
|
||||
r3 = client.get("/api/experiments", params={"q": "选出 3"})
|
||||
assert [e["id"] for e in r3.json()] == ["EXP-A4"]
|
||||
# kind + q 组合
|
||||
r4 = client.get("/api/experiments", params={"kind": "backtest", "q": "momentum"})
|
||||
assert [e["id"] for e in r4.json()] == ["EXP-A2"]
|
||||
|
||||
def test_like_wildcards_and_injection_are_bound(self, api) -> None:
|
||||
client, Session = api
|
||||
self._seed(Session)
|
||||
# `%` 是 LIKE 通配符,必须被转义成字面量:种子里只有 EXP-A1/A2 的摘要带
|
||||
# 百分号,若未转义会返回全部 4 条
|
||||
r = client.get("/api/experiments", params={"q": "%"})
|
||||
assert {e["id"] for e in r.json()} == {"EXP-A1", "EXP-A2"}
|
||||
assert r.headers["X-Total-Count"] == "2"
|
||||
# 注入尝试必须是普通字符串(参数绑定),不得改变语义
|
||||
inj = client.get("/api/experiments", params={"q": "' OR 1=1 --"})
|
||||
assert inj.json() == [] and inj.headers["X-Total-Count"] == "0"
|
||||
|
||||
def test_limit_offset_and_total_count_not_silently_truncated(self, api) -> None:
|
||||
client, Session = api
|
||||
self._seed(Session)
|
||||
page1 = client.get("/api/experiments", params={"limit": 2, "offset": 0})
|
||||
assert len(page1.json()) == 2
|
||||
assert page1.headers["X-Total-Count"] == "4", "总数必须暴露,客户端才知道被截断"
|
||||
page2 = client.get("/api/experiments", params={"limit": 2, "offset": 2})
|
||||
assert len(page2.json()) == 2
|
||||
ids = {e["id"] for e in page1.json()} | {e["id"] for e in page2.json()}
|
||||
assert len(ids) == 4, "分页不得重复/漏项(created_at 同秒时靠 id 次级排序稳定)"
|
||||
# 默认 limit(200)远大于旧硬编码的 50
|
||||
assert client.get("/api/experiments").status_code == 200
|
||||
|
||||
def test_meta_includes_data_version_job_id_and_result_bytes(self, api) -> None:
|
||||
client, Session = api
|
||||
seeded = self._seed(Session)
|
||||
row = next(e for e in client.get("/api/experiments").json() if e["id"] == "EXP-A1")
|
||||
assert row["data_version"] == "d20260904;n1;a2;b3", "列表必须暴露数据指纹"
|
||||
assert row["job_id"] is None
|
||||
assert row["result_bytes"] == len(json.dumps({"kind": "backtest", "n": 0}))
|
||||
|
||||
# 详情路径(result_json 在内存,len() 零成本)与列表口径必须一致
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
exp_id = resp.headers["X-Experiment-Id"]
|
||||
exp = _exp_by_id(Session, exp_id)
|
||||
assert exp is not None
|
||||
detail = client.get(f"/api/experiments/{exp_id}").json()
|
||||
assert detail["result_bytes"] == len(exp.result_json)
|
||||
assert detail["data_version"] == exp.data_version
|
||||
assert detail["job_id"] is None
|
||||
listed = next(e for e in client.get("/api/experiments").json() if e["id"] == exp_id)
|
||||
assert listed["result_bytes"] == detail["result_bytes"], "列表与详情体积口径一致"
|
||||
assert set(seeded) == {"EXP-A1", "EXP-A2", "EXP-A3", "EXP-A4"}
|
||||
|
||||
def test_repository_list_filtered_returns_summaries_without_result_json(self, api) -> None:
|
||||
_client, Session = api
|
||||
self._seed(Session)
|
||||
with Session() as session:
|
||||
repo = SqlAlchemyExperimentRepository(session)
|
||||
rows = repo.list_filtered(kind="backtest", limit=10)
|
||||
assert all(isinstance(r, ExperimentSummary) for r in rows)
|
||||
assert not hasattr(rows[0], "result_json"), "列表不得拉取 MEDIUMTEXT 大字段"
|
||||
assert repo.count_filtered() == 4
|
||||
assert repo.count_filtered(kind="selection") == 1
|
||||
|
||||
|
||||
class TestDelete:
|
||||
"""任务 4:DELETE 只删归档,不删 Job 记录。"""
|
||||
|
||||
def test_delete_then_404(self, api) -> None:
|
||||
client, Session = api
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
exp_id = resp.headers["X-Experiment-Id"]
|
||||
|
||||
deleted = client.delete(f"/api/experiments/{exp_id}")
|
||||
assert deleted.status_code == 200
|
||||
assert deleted.json() == {"deleted": exp_id}
|
||||
assert _exp_by_id(Session, exp_id) is None
|
||||
assert client.get(f"/api/experiments/{exp_id}").status_code == 404
|
||||
again = client.delete(f"/api/experiments/{exp_id}")
|
||||
assert again.status_code == 404
|
||||
assert "不存在" in again.json()["detail"]
|
||||
|
||||
def test_delete_keeps_job_record_and_reports_missing_result(self, api) -> None:
|
||||
"""删除归档不触碰 job 表;job 的结果只剩归档一份,故如实说明原因而非静默 200 空值。"""
|
||||
client, Session = api
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
exp_id = resp.headers["X-Experiment-Id"]
|
||||
with Session() as session:
|
||||
job = JobRecord(
|
||||
id="JOB-DEL-1",
|
||||
kind="backtest",
|
||||
spec_json=json.dumps(_BACKTEST_BODY),
|
||||
status=JobStatus.SUCCESS,
|
||||
experiment_id=exp_id,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
SqlAlchemyJobRepository(session).create(job)
|
||||
session.commit()
|
||||
|
||||
assert client.delete(f"/api/experiments/{exp_id}").status_code == 200
|
||||
with Session() as session:
|
||||
kept = SqlAlchemyJobRepository(session).get("JOB-DEL-1")
|
||||
assert kept is not None, "job 记录是执行历史,DELETE 不得连带删除"
|
||||
assert kept.experiment_id == exp_id
|
||||
view = client.get("/api/jobs/JOB-DEL-1")
|
||||
assert view.status_code == 200
|
||||
assert view.json()["result"] is None
|
||||
assert "已不存在" in view.json()["result_unavailable_reason"]
|
||||
|
||||
|
||||
class TestCurveCompleteness:
|
||||
"""P0:默认完整存档(不再按 60 只截断);超预算时裁剪并留证据。"""
|
||||
|
||||
@staticmethod
|
||||
def _spec(top_n: int, symbols: list[str]) -> ResearchSpec:
|
||||
return ResearchSpec(
|
||||
type="backtest",
|
||||
universe=UniverseSpec(
|
||||
exclude_st=False, min_listing_days=0, symbols=list(symbols)
|
||||
),
|
||||
factors=[FactorSpec(name="momentum_20")],
|
||||
selection=SelectionSpec(top_n=top_n, hold_top_x=top_n, allow_substitute=False, defer_buy=True),
|
||||
rebalance="monthly",
|
||||
selection_interval_months=6,
|
||||
rebalance_interval_months=6,
|
||||
period=(date(2024, 3, 1), date(2024, 6, 30)),
|
||||
costs=CostSpec(
|
||||
commission_rate=0.0, stamp_tax_rate=0.0, slippage_rate=0.0
|
||||
),
|
||||
)
|
||||
|
||||
def test_default_stores_every_held_symbol_beyond_60(self) -> None:
|
||||
syms = [f"60{i:04d}.SH" for i in range(70)]
|
||||
drifts = {s: 0.001 * (i % 4) + 0.0005 for i, s in enumerate(syms)}
|
||||
daily = synthetic_daily(drifts, n=120)
|
||||
res = LocalEngine().run_backtest(daily, self._spec(len(syms), syms))
|
||||
|
||||
held = {a.symbol for a in res.fills if a.symbol}
|
||||
assert len(held) > 60, f"构造的样本应持有 60 只以上,实际 {len(held)}"
|
||||
curve_syms = {c.symbol for c in res.symbol_curves}
|
||||
assert held <= curve_syms, "期内持有的每只都必须有曲线(完整存档,默认不截断)"
|
||||
assert not [n for n in res.unimplemented if "个股收益曲线" in n], "默认不得出现截断说明"
|
||||
|
||||
def test_configured_limit_truncates_and_notes_actual_limit(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""配置了上限就截断 —— 且 note 里写清「实际保留几只 / 共几只」。"""
|
||||
from app.core import config as config_mod
|
||||
from app.quant import local_engine as le
|
||||
|
||||
monkeypatch.setattr(
|
||||
config_mod,
|
||||
"get_settings",
|
||||
lambda: SimpleNamespace(research_archive_curve_limit=3),
|
||||
)
|
||||
assert le._MAX_SYMBOL_CURVES is None, "默认代码路径不设上限(配置优先)"
|
||||
syms = [f"60{i:04d}.SH" for i in range(10)]
|
||||
drifts = {s: 0.001 * (i % 3) + 0.0005 for i, s in enumerate(syms)}
|
||||
daily = synthetic_daily(drifts, n=90)
|
||||
res = LocalEngine().run_backtest(daily, self._spec(len(syms), syms))
|
||||
|
||||
assert len(res.symbol_curves) == 3
|
||||
note = [n for n in res.unimplemented if "个股收益曲线" in n]
|
||||
assert note and "3 只" in note[0] and "共持有" in note[0]
|
||||
|
||||
def test_archive_budget_truncates_with_machine_readable_evidence(
|
||||
self, api, monkeypatch
|
||||
) -> None:
|
||||
"""归档字节预算:超预算时裁剪曲线,并把证据同时写进 archive_meta 与 unimplemented。
|
||||
|
||||
预算不靠猜:先按默认预算归档一次拿到「完整归档的字节数」L,再把预算设为
|
||||
L - 2,000 重新归档 —— 必然需要裁掉若干条曲线,且裁剪后必须真的落到预算内。
|
||||
"""
|
||||
client, Session = api
|
||||
body = {
|
||||
**_BACKTEST_BODY,
|
||||
"selection": {"top_n": 5},
|
||||
"universe": {**_BACKTEST_BODY["universe"], "symbols": _SYMS},
|
||||
}
|
||||
first = client.post("/api/backtests", json=body)
|
||||
assert first.status_code == 200
|
||||
assert first.json()["archive_meta"]["truncated"] is False, "默认预算下应完整存档"
|
||||
full = _exp_by_id(Session, first.headers["X-Experiment-Id"])
|
||||
assert full is not None
|
||||
full_bytes = len(full.result_json.encode("utf-8"))
|
||||
assert first.json()["archive_meta"]["curves_total"] >= 3, "样本需有足够曲线才能验证裁剪"
|
||||
|
||||
budget = full_bytes - 2_000
|
||||
monkeypatch.setattr(ea, "_archive_budget_bytes", lambda: budget)
|
||||
resp = client.post("/api/backtests", json=body)
|
||||
assert resp.status_code == 200
|
||||
meta = resp.json()["archive_meta"]
|
||||
assert meta["truncated"] is True, f"预算 {budget} < 完整体积 {full_bytes},必须触发裁剪"
|
||||
assert 0 < meta["curves_stored"] < meta["curves_total"], meta
|
||||
assert meta["budget_bytes"] == budget
|
||||
assert meta["over_budget"] is False
|
||||
assert any("个股收益曲线" in n for n in resp.json()["unimplemented"])
|
||||
|
||||
exp = _exp_by_id(Session, resp.headers["X-Experiment-Id"])
|
||||
assert exp is not None
|
||||
assert len(exp.result_json.encode("utf-8")) <= budget, "裁剪后必须真的放进预算"
|
||||
archived = json.loads(exp.result_json)
|
||||
assert archived["archive_meta"]["truncated"] is True
|
||||
assert archived["archive_meta"]["curves_stored"] == len(archived["symbol_curves"])
|
||||
assert any("个股收益曲线" in n for n in archived["unimplemented"])
|
||||
assert archived["archive_meta"]["result_bytes"] == len(exp.result_json.encode("utf-8"))
|
||||
|
||||
def test_archive_over_budget_is_labelled_not_hidden(self, api, monkeypatch) -> None:
|
||||
"""预算小到连归档主体都放不下:必须标 over_budget + 说明,绝不假装达标。"""
|
||||
client, Session = api
|
||||
budget = 5_000
|
||||
monkeypatch.setattr(ea, "_archive_budget_bytes", lambda: budget)
|
||||
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
meta = body["archive_meta"]
|
||||
assert meta["curves_stored"] == 0 and meta["truncated"] is True
|
||||
assert meta["over_budget"] is True, "主体超预算必须显式标注,而不是静默超标"
|
||||
assert any("超过配置的体积预算" in n for n in body["unimplemented"])
|
||||
exp = _exp_by_id(Session, resp.headers["X-Experiment-Id"])
|
||||
assert exp is not None
|
||||
archived = json.loads(exp.result_json)
|
||||
assert archived["archive_meta"]["over_budget"] is True
|
||||
assert any("超过配置的体积预算" in n for n in archived["unimplemented"])
|
||||
# 如实记录「确实超标」这一事实(不掩盖,也不假装裁剪到位)
|
||||
assert archived["archive_meta"]["result_bytes"] > budget
|
||||
|
||||
def test_archive_meta_complete_when_under_budget(self, api) -> None:
|
||||
client, Session = api
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
exp = _exp_by_id(Session, resp.headers["X-Experiment-Id"])
|
||||
assert exp is not None
|
||||
archived = json.loads(exp.result_json)
|
||||
meta = archived["archive_meta"]
|
||||
assert meta["truncated"] is False
|
||||
assert meta["curves_stored"] == meta["curves_total"] == len(archived["symbol_curves"])
|
||||
assert meta["result_bytes"] == len(exp.result_json.encode("utf-8"))
|
||||
assert meta["result_chars"] == len(exp.result_json)
|
||||
|
||||
|
||||
class TestSingleResultCopy:
|
||||
"""P1:完整结果只存 experiment 一份;GET /api/jobs/{id} 契约与老记录兼容。"""
|
||||
|
||||
def test_submit_and_run_reads_through_in_memory_only(self, tmp_path) -> None:
|
||||
"""`submit_and_run` 的返回对象仍带完整结果(既有调用方如
|
||||
scripts/run_dividend_case.py 依赖它),但**数据库里不存第二份**。"""
|
||||
from app.application.services.job_executor import submit_and_run
|
||||
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'submit.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
||||
drifts = {sym: 0.003 - 0.0015 * i for i, sym in enumerate(_SYMS)}
|
||||
bars = bars_dataframe_to_daily_bars(synthetic_daily(drifts, n=300))
|
||||
stocks = [
|
||||
Stock(symbol=sym, name=f"测试股份{i}", list_date=date(1999, 11, 10))
|
||||
for i, sym in enumerate(_SYMS)
|
||||
]
|
||||
spec = ResearchSpec.model_validate(_BACKTEST_BODY)
|
||||
done = submit_and_run(
|
||||
spec,
|
||||
factories={
|
||||
"session_factory": Session,
|
||||
"job_repo_factory": lambda s: SqlAlchemyJobRepository(s),
|
||||
"experiment_repo_factory": lambda s: SqlAlchemyExperimentRepository(s),
|
||||
"stock_repo_factory": lambda s: _MemStockRepo(stocks),
|
||||
"daily_repo_factory": lambda s: _MemDailyRepo(bars),
|
||||
"engine": LocalEngine(),
|
||||
},
|
||||
)
|
||||
assert done.status == JobStatus.SUCCESS
|
||||
assert done.experiment_id
|
||||
assert done.result_json, "返回对象必须带完整结果(内存读透,供既有调用方使用)"
|
||||
assert BacktestResult.model_validate_json(done.result_json).symbol_curves
|
||||
|
||||
with Session() as session:
|
||||
persisted = SqlAlchemyJobRepository(session).get(done.id)
|
||||
exp = SqlAlchemyExperimentRepository(session).get(done.experiment_id)
|
||||
assert persisted is not None and persisted.result_json is None, "数据库不得存第二份"
|
||||
assert exp is not None and exp.result_json, "完整结果在 experiment 侧"
|
||||
|
||||
def test_job_view_reads_experiment_and_keeps_full_result(self, api) -> None:
|
||||
from app.api.jobs import _job_view
|
||||
|
||||
client, Session = api
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
exp_id = resp.headers["X-Experiment-Id"]
|
||||
with Session() as session:
|
||||
job = JobRecord(
|
||||
id="JOB-VIEW-1",
|
||||
kind="backtest",
|
||||
spec_json=json.dumps(_BACKTEST_BODY),
|
||||
status=JobStatus.SUCCESS,
|
||||
result_json=None, # 新形态:job 侧不再保存结果副本
|
||||
experiment_id=exp_id,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
SqlAlchemyJobRepository(session).create(job)
|
||||
session.commit()
|
||||
view = _job_view(job, SqlAlchemyExperimentRepository(session))
|
||||
assert view["result_source"] == "experiment"
|
||||
assert view["result"].symbol_curves, "契约不变:GET /api/jobs/{id} 仍返回完整结果"
|
||||
assert (
|
||||
view["result"].summary.total_return_pct
|
||||
== resp.json()["summary"]["total_return_pct"]
|
||||
)
|
||||
assert client.get("/api/jobs/JOB-VIEW-1").json()["result"]["symbol_curves"]
|
||||
|
||||
def test_legacy_job_record_still_decodes_from_result_json(self, api) -> None:
|
||||
"""老形态(result_json 有值、experiment_id 为空)必须继续可用。"""
|
||||
from app.api.jobs import _job_view
|
||||
|
||||
client, Session = api
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
full = json.dumps(resp.json(), ensure_ascii=False)
|
||||
with Session() as session:
|
||||
job = JobRecord(
|
||||
id="JOB-LEGACY-1",
|
||||
kind="backtest",
|
||||
spec_json=json.dumps(_BACKTEST_BODY),
|
||||
status=JobStatus.SUCCESS,
|
||||
result_json=full,
|
||||
experiment_id=None,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
SqlAlchemyJobRepository(session).create(job)
|
||||
session.commit()
|
||||
view = _job_view(job, SqlAlchemyExperimentRepository(session))
|
||||
assert view["result_source"] == "job"
|
||||
assert view["result"].symbol_curves
|
||||
assert view["result"].summary.total_return_pct == resp.json()["summary"]["total_return_pct"]
|
||||
|
||||
def test_job_table_no_longer_duplicates_result(self, api) -> None:
|
||||
"""json 列确实没有第二份结果(体积审计的机器可读证据)。"""
|
||||
client, Session = api
|
||||
resp = client.post("/api/backtests", json=_BACKTEST_BODY)
|
||||
exp_id = resp.headers["X-Experiment-Id"]
|
||||
with Session() as session:
|
||||
job = JobRecord(
|
||||
id="JOB-SIZE-1",
|
||||
kind="backtest",
|
||||
spec_json=json.dumps(_BACKTEST_BODY),
|
||||
status=JobStatus.SUCCESS,
|
||||
experiment_id=exp_id,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
SqlAlchemyJobRepository(session).create(job)
|
||||
session.commit()
|
||||
with Session() as session:
|
||||
row = session.execute(
|
||||
select(JobModel.result_json).where(JobModel.id == "JOB-SIZE-1")
|
||||
).scalar()
|
||||
exp_row = session.execute(
|
||||
select(ExperimentModel.result_json).where(ExperimentModel.id == exp_id)
|
||||
).scalar()
|
||||
assert row is None, "job 侧不得再存结果副本"
|
||||
assert exp_row, "完整结果必须在 experiment 侧"
|
||||
archived = _exp_by_id(Session, exp_id)
|
||||
assert archived is not None
|
||||
assert len(exp_row) == len(archived.result_json)
|
||||
@@ -95,3 +95,50 @@ class TestFactorsApi:
|
||||
resp = client.get("/api/factors")
|
||||
names = {r["name"] for r in resp.json()}
|
||||
assert names == {d.name for d in list_factors()}
|
||||
|
||||
|
||||
class TestRegistrySyncRegression:
|
||||
"""回归:表非空时也必须补齐「注册表有、库里没有」的因子。
|
||||
|
||||
历史 bug:seed 只在表为空时触发,导致 `dividend_yield` 等后加的因子永远不进目录
|
||||
(实测真实库表里 9 条、注册表 11 条),前端因子下拉与归档说明都取不到它们。
|
||||
"""
|
||||
|
||||
def test_missing_registry_factor_is_seeded_when_table_not_empty(self, client, tmp_path) -> None:
|
||||
from sqlalchemy import text
|
||||
|
||||
# 1) 先正常读一次 → 目录完整(含股息率因子)
|
||||
full = {r["name"] for r in client.get("/api/factors").json()}
|
||||
assert "dividend_yield" in full, "注册表里的股息率因子必须出现在目录中"
|
||||
|
||||
# 2) 人为删除一行,复现「表非空但缺因子」的历史状态
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
|
||||
with engine.begin() as conn:
|
||||
conn.execute(text("DELETE FROM factor_definition WHERE name = 'dividend_yield'"))
|
||||
rows = client.get("/api/factors").json()
|
||||
# 3) 再读 → 缺失因子被当场补齐
|
||||
names = {r["name"] for r in rows}
|
||||
assert "dividend_yield" in names, "表非空时也必须补齐缺失的注册表因子"
|
||||
assert names == {d.name for d in list_factors()}
|
||||
# 4) 用户登记的自定义因子元数据不被覆盖/删除(只补不删)
|
||||
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
||||
with Session() as s2:
|
||||
SqlAlchemyFactorRepository(s2).upsert_many(
|
||||
[
|
||||
FactorDefinition(
|
||||
name="my_custom_factor",
|
||||
description="自定义因子",
|
||||
formula="x",
|
||||
brief="b",
|
||||
frequency="daily",
|
||||
lookback=5,
|
||||
direction="higher_is_better",
|
||||
requires=[],
|
||||
version="1",
|
||||
)
|
||||
]
|
||||
)
|
||||
s2.commit()
|
||||
names2 = {r["name"] for r in client.get("/api/factors").json()}
|
||||
assert "my_custom_factor" in names2, "只补不删:自定义因子必须保留"
|
||||
assert "dividend_yield" in names2
|
||||
|
||||
@@ -117,3 +117,46 @@ class TestUniverseIndexCodeFilter:
|
||||
|
||||
def test_index_code_filter(self, tmp_path) -> None:
|
||||
self._build(tmp_path)
|
||||
|
||||
|
||||
class TestDelistedUniverse:
|
||||
"""退市股的时点股票池语义(幸存者偏差修正的核心断言)。
|
||||
|
||||
退市股必须在**退市日之前**纳入池子、退市之后排除;否则回测只剩「活下来的
|
||||
赢家」,收益被系统性高估(高股息策略尤其容易被股息陷阱的退市股反噬)。
|
||||
"""
|
||||
|
||||
def _stock(self, symbol: str, list_date: date, delist_date: date | None):
|
||||
from app.domain.entities.market import Stock
|
||||
|
||||
return Stock(
|
||||
symbol=symbol, name=f"测试{symbol}", list_date=list_date, delist_date=delist_date
|
||||
)
|
||||
|
||||
def test_delisted_included_before_excluded_after(self) -> None:
|
||||
from app.domain.entities.research import UniverseSpec
|
||||
from app.quant.universe import filter_stocks
|
||||
|
||||
stocks = [
|
||||
self._stock("600000.SH", date(1999, 11, 10), None),
|
||||
self._stock("000005.SZ", date(1990, 12, 10), date(2024, 4, 26)),
|
||||
]
|
||||
u = UniverseSpec(exclude_st=False, min_listing_days=0)
|
||||
before = {s.symbol for s in filter_stocks(stocks, u, as_of=date(2024, 1, 2))}
|
||||
after = {s.symbol for s in filter_stocks(stocks, u, as_of=date(2024, 6, 3))}
|
||||
on_delist_day = {s.symbol for s in filter_stocks(stocks, u, as_of=date(2024, 4, 26))}
|
||||
assert "000005.SZ" in before and "000005.SZ" not in after
|
||||
assert "000005.SZ" in on_delist_day # 退市日当天仍在(数据截至当日)
|
||||
assert "600000.SH" in before and "600000.SH" in after
|
||||
|
||||
def test_st_name_excludes_regardless_of_period(self) -> None:
|
||||
"""已知局限:exclude_st 用**最新名称**判定,会把曾用名非 ST 的标的整段排除。"""
|
||||
from app.domain.entities.research import UniverseSpec
|
||||
from app.quant.universe import filter_stocks
|
||||
|
||||
stocks = [self._stock("000005.SZ", date(1990, 12, 10), date(2024, 4, 26))]
|
||||
stocks[0].name = "ST星源(退)"
|
||||
u = UniverseSpec(exclude_st=True, min_listing_days=0)
|
||||
assert filter_stocks(stocks, u, as_of=date(2024, 1, 2)) == []
|
||||
u2 = UniverseSpec(exclude_st=False, min_listing_days=0)
|
||||
assert len(filter_stocks(stocks, u2, as_of=date(2024, 1, 2))) == 1
|
||||
|
||||
@@ -0,0 +1,328 @@
|
||||
"""择股/调仓双周期(m/y)、两级截断(n/x)与顺延买入(defer_buy)测试。
|
||||
|
||||
对应用户案例:「全市场股息率最高的 n 只 → 持仓前 x 只;每 m 个月择股一次;
|
||||
每 y 个月调仓,默认 y=m;买卖点为收盘价;买不进时顺延到之后不涨停的交易日买入」。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from app.domain.entities.research import (
|
||||
CostSpec,
|
||||
FactorSpec,
|
||||
ResearchSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.quant.engine import LocalEngine
|
||||
from app.quant.local_engine import rebalance_dates
|
||||
from pydantic import ValidationError
|
||||
|
||||
from conftest_quant import synthetic_daily
|
||||
|
||||
|
||||
def _spec(
|
||||
top_n: int = 2,
|
||||
hold_top_x: int | None = None,
|
||||
start: date = date(2024, 3, 1),
|
||||
end: date = date(2024, 12, 20),
|
||||
m: int | None = None,
|
||||
y: int | None = None,
|
||||
allow_substitute: bool = False,
|
||||
defer_buy: bool = True,
|
||||
) -> ResearchSpec:
|
||||
"""合成行情自 2024-01-01 起(因子预热),回测自 2024-03-01 起(首个调仓日即有信号)。"""
|
||||
return ResearchSpec(
|
||||
type="backtest",
|
||||
universe=UniverseSpec(exclude_st=False, min_listing_days=0),
|
||||
factors=[FactorSpec(name="momentum_20")],
|
||||
selection=SelectionSpec(
|
||||
top_n=top_n,
|
||||
hold_top_x=hold_top_x,
|
||||
allow_substitute=allow_substitute,
|
||||
defer_buy=defer_buy,
|
||||
),
|
||||
rebalance="monthly",
|
||||
selection_interval_months=m,
|
||||
rebalance_interval_months=y,
|
||||
period=(start, end),
|
||||
costs=CostSpec(commission_rate=0.0, stamp_tax_rate=0.0, slippage_rate=0.0),
|
||||
)
|
||||
|
||||
|
||||
class TestIntervalSchedule:
|
||||
def test_every_n_months_anchored_at_start_month(self) -> None:
|
||||
idx = pd.bdate_range("2020-01-01", "2021-12-31")
|
||||
out = rebalance_dates(idx, "monthly", date(2020, 1, 1), every_months=6)
|
||||
assert [d.strftime("%Y-%m-%d") for d in out] == [
|
||||
"2020-01-01", "2020-07-01", "2021-01-01", "2021-07-01",
|
||||
]
|
||||
|
||||
def test_anchor_moves_with_start_month(self) -> None:
|
||||
idx = pd.bdate_range("2020-01-01", "2021-12-31")
|
||||
out = rebalance_dates(idx, "monthly", date(2020, 3, 1), every_months=6)
|
||||
assert [d.strftime("%Y-%m-%d") for d in out] == [
|
||||
"2020-03-02", "2020-09-01", "2021-03-01", "2021-09-01",
|
||||
]
|
||||
|
||||
def test_start_not_first_trading_day_keeps_anchor_month(self) -> None:
|
||||
"""起始日非月初时不得跳过锚点月(否则白等 m 个月才首次建仓)。"""
|
||||
idx = pd.bdate_range("2024-01-01", "2025-12-31")
|
||||
out = rebalance_dates(idx, "monthly", date(2024, 3, 15), every_months=6)
|
||||
# 2024-03-15 本身是交易日 → 锚点即当日;后续按 +6 个月推进
|
||||
assert [d.strftime("%Y-%m-%d") for d in out] == [
|
||||
"2024-03-15", "2024-09-02", "2025-03-03", "2025-09-01",
|
||||
]
|
||||
# 起始日落在非交易日(2024-03-16/17 为周末)→ 取之后首个交易日,仍属 3 月
|
||||
out2 = rebalance_dates(idx, "monthly", date(2024, 3, 16), every_months=6)
|
||||
assert out2[0].strftime("%Y-%m-%d") == "2024-03-18"
|
||||
|
||||
def test_start_after_last_trading_day_of_month_moves_anchor(self) -> None:
|
||||
"""起始日晚于该月最后一个交易日时,锚点自然落到下一个月(不产生空区间)。"""
|
||||
idx = pd.bdate_range("2024-01-01", "2025-12-31")
|
||||
out = rebalance_dates(idx, "monthly", date(2024, 3, 31), every_months=6)
|
||||
assert [d.strftime("%Y-%m-%d") for d in out][:2] == ["2024-04-01", "2024-10-01"]
|
||||
|
||||
def test_non_month_start_does_not_leave_early_cash(self) -> None:
|
||||
"""非月初起始 + m=y=6:首个择股/调仓日不应晚于起始月,净值不得长期恒为初始值。"""
|
||||
daily = synthetic_daily({"600000.SH": 0.004, "600001.SH": 0.002}, n=260)
|
||||
res = LocalEngine().run_backtest(
|
||||
daily, _spec(top_n=1, m=6, y=6, start=date(2024, 3, 15), end=date(2024, 12, 20))
|
||||
)
|
||||
first_pos = min(p.date for p in res.positions)
|
||||
assert first_pos == date(2024, 3, 15), first_pos
|
||||
# 起始日之后的前 20 个交易日里不应出现「净值恒为初始资金」
|
||||
head = res.equity_curve[:20]
|
||||
assert len({p.value for p in head}) > 1, "起始月即应建仓,净值不应恒为初始资金"
|
||||
|
||||
def test_selection_and_rebalance_schedules_are_independent(self) -> None:
|
||||
"""m=6、y=3:择股 2 次,调仓 4 次。"""
|
||||
daily = synthetic_daily({f"60000{i}.SH": 0.004 - 0.001 * i for i in range(4)}, n=260)
|
||||
res = LocalEngine().run_backtest(daily, _spec(top_n=2, m=6, y=3))
|
||||
sel_dates = sorted({p.date for p in res.selection_history})
|
||||
rebal_dates_ = sorted({p.date for p in res.positions})
|
||||
assert len(sel_dates) == 2, sel_dates # 2024-03-01 / 2024-09-02
|
||||
assert len(rebal_dates_) == 4, rebal_dates_ # 03/06/09/12 各一次
|
||||
assert set(sel_dates) < set(rebal_dates_)
|
||||
|
||||
def test_default_y_equals_m(self) -> None:
|
||||
spec = _spec(m=6)
|
||||
assert spec.effective_selection_months == 6
|
||||
assert spec.effective_rebalance_months == 6 # y 缺省 → 跟随 m
|
||||
|
||||
def test_y_without_m_rejected(self) -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
_spec(m=None, y=3)
|
||||
|
||||
|
||||
class TestTwoLevelTruncation:
|
||||
def test_hold_only_top_x_of_pool(self) -> None:
|
||||
"""n=4、x=2:候选池记录 4 只,实际只持仓前 2 只。"""
|
||||
daily = synthetic_daily({f"60000{i}.SH": 0.006 - 0.0015 * i for i in range(6)}, n=260)
|
||||
res = LocalEngine().run_backtest(daily, _spec(top_n=4, hold_top_x=2, m=2, y=2))
|
||||
pool_dates = {p.date for p in res.selection_history}
|
||||
for d in pool_dates:
|
||||
assert len([p for p in res.selection_history if p.date == d]) == 4
|
||||
# 持仓数量不超过 x=2
|
||||
per_date: dict = {}
|
||||
for p in res.positions:
|
||||
per_date.setdefault(p.date, []).append(p.symbol)
|
||||
assert per_date and all(len(v) <= 2 for v in per_date.values())
|
||||
|
||||
def test_x_cannot_exceed_n(self) -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
SelectionSpec(top_n=5, hold_top_x=10)
|
||||
|
||||
def test_substitute_and_defer_are_mutually_exclusive(self) -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
SelectionSpec(top_n=5, allow_substitute=True, defer_buy=True)
|
||||
|
||||
def test_no_substitute_keeps_pool_membership(self) -> None:
|
||||
"""defer 模式(不替补):持仓必属候选池,绝不出现池外标的。"""
|
||||
daily = synthetic_daily({f"60000{i}.SH": 0.006 - 0.0015 * i for i in range(6)}, n=260)
|
||||
res = LocalEngine().run_backtest(daily, _spec(top_n=2, hold_top_x=2, m=6, y=6))
|
||||
for pos in res.positions:
|
||||
same_day_pool = {p.symbol for p in res.selection_history if p.date == pos.date}
|
||||
assert pos.symbol in same_day_pool, f"{pos.symbol} 不在当日候选池 {same_day_pool}"
|
||||
|
||||
|
||||
class TestDeferBuy:
|
||||
"""顺延买入:涨停当日不成交,之后首个不涨停交易日按收盘价买入。"""
|
||||
|
||||
def _limit_up_frame(self) -> tuple[pd.DataFrame, str, date, date]:
|
||||
"""构造:600000.SH 在首个调仓日涨停,之后恢复;动量高于 600001.SH。"""
|
||||
daily = synthetic_daily({"600000.SH": 0.006, "600001.SH": 0.001}, n=260)
|
||||
dates = sorted(pd.to_datetime(daily["trade_date"].unique()))
|
||||
d0 = next(d for d in dates if d.date() >= date(2024, 3, 1)) # 首个调仓日
|
||||
d1 = dates[dates.index(d0) + 1]
|
||||
prev = dates[dates.index(d0) - 1]
|
||||
prev_close = float(
|
||||
daily[(daily["symbol"] == "600000.SH") & (daily["trade_date"] == prev.date())][
|
||||
"close"
|
||||
].iloc[0]
|
||||
)
|
||||
mask = (daily["symbol"] == "600000.SH") & (daily["trade_date"] == d0.date())
|
||||
daily.loc[mask, "close"] = prev_close * 1.10 # 主板涨停
|
||||
daily.loc[mask, "high"] = prev_close * 1.10
|
||||
return daily, "600000.SH", d0.date(), d1.date()
|
||||
|
||||
def test_buy_deferred_to_next_tradable_day(self) -> None:
|
||||
daily, sym, d0, d1 = self._limit_up_frame()
|
||||
res = LocalEngine().run_backtest(daily, _spec(top_n=2, hold_top_x=2, m=6, y=6))
|
||||
# 调仓日意图登记为未成交,原因含「涨停」与「顺延」
|
||||
rejects = [
|
||||
a for a in res.signal_history
|
||||
if a.date == d0 and a.symbol == sym and a.signal == "BUY" and not a.filled
|
||||
]
|
||||
assert rejects and "涨停" in (rejects[0].reject_reason or "")
|
||||
assert "顺延" in (rejects[0].reject_reason or "")
|
||||
# 当日无成交
|
||||
assert not [a for a in res.fills if a.date == d0 and a.symbol == sym]
|
||||
# 次一交易日按收盘价成交(价格 = 该日 close,滑点为 0)
|
||||
fills = [a for a in res.fills if a.symbol == sym and a.signal == "BUY"]
|
||||
assert fills, "顺延后应成交"
|
||||
assert fills[0].date == d1
|
||||
px = float(
|
||||
daily[(daily["symbol"] == sym) & (daily["trade_date"] == d1)]["close"].iloc[0]
|
||||
)
|
||||
assert fills[0].price == pytest.approx(px, rel=1e-6)
|
||||
|
||||
def test_defer_disabled_gives_up_immediately(self) -> None:
|
||||
"""defer_buy=False:涨停当日被拒后直接放弃,不顺延(区间内无下一次调仓)。"""
|
||||
daily, sym, d0, d1 = self._limit_up_frame()
|
||||
res = LocalEngine().run_backtest(
|
||||
daily, _spec(top_n=2, hold_top_x=2, m=6, y=6, end=date(2024, 6, 28), defer_buy=False)
|
||||
)
|
||||
assert not [a for a in res.fills if a.symbol == sym and a.signal == "BUY"]
|
||||
rejects = [
|
||||
a for a in res.signal_history
|
||||
if a.date == d0 and a.symbol == sym and a.signal == "BUY" and not a.filled
|
||||
]
|
||||
assert rejects
|
||||
assert "顺延" not in (rejects[0].reject_reason or "")
|
||||
|
||||
def test_pending_order_cleared_at_next_rebalance(self) -> None:
|
||||
"""顺延单不跨调仓:3/1 挂的单在 4/1 调仓时作废,4/1 是新的调仓尝试。"""
|
||||
daily = synthetic_daily({"600000.SH": 0.006, "600001.SH": 0.001}, n=260)
|
||||
dates = sorted(pd.to_datetime(daily["trade_date"].unique()))
|
||||
d0 = min(d for d in dates if d.date() >= date(2024, 3, 1))
|
||||
d1 = min(d for d in dates if d.date() >= date(2024, 4, 1))
|
||||
# 2024-03-01 ~ 04-01 连续每日涨停(链式 +10%)
|
||||
seed = dates[dates.index(d0) - 1]
|
||||
prev_close = float(
|
||||
daily[(daily["symbol"] == "600000.SH") & (daily["trade_date"] == seed.date())][
|
||||
"close"
|
||||
].iloc[0]
|
||||
)
|
||||
for d in dates:
|
||||
if not (d0.date() <= d.date() <= d1.date()):
|
||||
continue
|
||||
prev_close = prev_close * 1.10
|
||||
mask = (daily["symbol"] == "600000.SH") & (daily["trade_date"] == d.date())
|
||||
daily.loc[mask, "close"] = prev_close
|
||||
daily.loc[mask, "high"] = prev_close
|
||||
res = LocalEngine().run_backtest(
|
||||
daily, _spec(top_n=1, hold_top_x=1, end=date(2024, 4, 30), m=12, y=1)
|
||||
)
|
||||
# 3/1 与 4/1 两次调仓尝试均被涨停拒绝
|
||||
reject_dates = {
|
||||
a.date for a in res.signal_history
|
||||
if a.symbol == "600000.SH" and a.signal == "BUY" and not a.filled
|
||||
}
|
||||
assert d0.date() in reject_dates and d1.date() in reject_dates
|
||||
# 3/1 挂出的顺延单在 4/1 之前一次都没成交(期间每日涨停)
|
||||
assert not [a for a in res.fills if a.date < d1.date()]
|
||||
# 4/1 之后涨停解除 → 新一次调仓的顺延单成交
|
||||
fills = [a for a in res.fills if a.symbol == "600000.SH"]
|
||||
assert len(fills) == 1 and fills[0].date > d1.date()
|
||||
|
||||
|
||||
class TestSymbolCurves:
|
||||
def test_symbol_curves_and_marks(self) -> None:
|
||||
daily = synthetic_daily({"600000.SH": 0.004, "600001.SH": -0.002}, n=260)
|
||||
res = LocalEngine().run_backtest(daily, _spec(top_n=1, m=6, y=6))
|
||||
assert res.symbol_curves, "应输出个股收益曲线"
|
||||
for curve in res.symbol_curves:
|
||||
assert curve.points, f"{curve.symbol} 曲线无数据点"
|
||||
if curve.marks:
|
||||
assert all(m.symbol == curve.symbol for m in curve.marks)
|
||||
assert all(m.filled for m in curve.marks)
|
||||
# 成交记录中的股票都应有曲线
|
||||
traded = {a.symbol for a in res.fills if a.symbol}
|
||||
assert traded <= {c.symbol for c in res.symbol_curves}
|
||||
|
||||
def test_symbol_curve_pct_matches_holding_gain(self) -> None:
|
||||
"""单股全程持有:曲线期末收益 ≈ 期末价/建仓日收盘 - 1(零成本下)。"""
|
||||
daily = synthetic_daily({"600000.SH": 0.004}, n=130)
|
||||
res = LocalEngine().run_backtest(
|
||||
daily, _spec(top_n=1, m=12, y=12, end=date(2024, 6, 28))
|
||||
)
|
||||
curve = res.symbol_curves[0]
|
||||
closes = daily[daily["symbol"] == "600000.SH"].sort_values("trade_date")
|
||||
entry_close = float(closes[closes["trade_date"] == date(2024, 3, 1)]["close"].iloc[0])
|
||||
exit_close = float(closes["close"].iloc[-1])
|
||||
expected = (exit_close / entry_close - 1) * 100
|
||||
assert curve.final_return_pct == pytest.approx(expected, rel=1e-3)
|
||||
# 建仓当日曲线为 0%(当日收盘成交,不计当日涨跌)
|
||||
assert curve.points[0].date == date(2024, 3, 1)
|
||||
assert curve.points[0].value == pytest.approx(0.0, abs=1e-9)
|
||||
|
||||
class TestFirstDayWithoutPrevClose:
|
||||
"""数据窗口起点无上一有效收盘价时的买入处理(案例 start=2020-01-01 实测命中)。
|
||||
|
||||
`_buyable` 在无前收时无法判定涨停:若按「不可买」处理,回测首个调仓日会被
|
||||
整体放弃(顺延到次日,白付一天空仓);现按「可买」处理并在 unimplemented
|
||||
中如实标注次数(AGENT.md §24)。
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _runner(daily: pd.DataFrame, spec: ResearchSpec):
|
||||
"""用常量 score 直接驱动 Runner:绕开因子预热,隔离「首个交易日」这一场景。"""
|
||||
from app.quant.local_engine import TopKBacktestRunner
|
||||
|
||||
close = daily.pivot(index="trade_date", columns="symbol", values="close")
|
||||
close.index = pd.to_datetime(close.index)
|
||||
score = pd.DataFrame(1.0, index=close.index, columns=close.columns)
|
||||
return TopKBacktestRunner(spec, score, close)
|
||||
|
||||
def test_buy_on_first_bar_without_prev_close(self) -> None:
|
||||
daily = synthetic_daily({"600000.SH": 0.004, "600001.SH": 0.002}, n=20)
|
||||
first_day = min(daily["trade_date"])
|
||||
res = self._runner(
|
||||
daily, _spec(top_n=1, m=1, y=1, start=first_day, end=max(daily["trade_date"]))
|
||||
).run()
|
||||
buy_days = [a.date for a in res.fills if a.signal == "BUY"]
|
||||
assert buy_days and min(buy_days) == first_day, "首个交易日即应成交,不应被整体顺延"
|
||||
assert any("无法判定涨停" in n for n in res.unimplemented), "应如实标注无前收的判定降级"
|
||||
|
||||
def test_no_note_when_prev_close_available(self) -> None:
|
||||
daily = synthetic_daily({"600000.SH": 0.004, "600001.SH": 0.002}, n=60)
|
||||
res = LocalEngine().run_backtest(daily, _spec(top_n=1, m=1, y=1))
|
||||
assert not any("无法判定涨停" in n for n in res.unimplemented)
|
||||
|
||||
|
||||
class TestSymbolCurvePayloadBound:
|
||||
def test_curves_capped_with_note(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""曲线数量上限生效并如实标注(结果体积约束:超限会让落库/传输不可用)。"""
|
||||
from app.quant import local_engine as le
|
||||
|
||||
monkeypatch.setattr(le, "_MAX_SYMBOL_CURVES", 2)
|
||||
drifts = {f"60000{i}.SH": 0.001 * (i % 3) for i in range(10)}
|
||||
daily = synthetic_daily(drifts, n=90)
|
||||
res = LocalEngine().run_backtest(daily, _spec(top_n=10, hold_top_x=10, m=6, y=6))
|
||||
assert len(res.symbol_curves) == 2
|
||||
note = [n for n in res.unimplemented if "个股收益曲线" in n]
|
||||
assert note and "共持有" in note[0]
|
||||
|
||||
def test_no_flat_points_when_not_held(self) -> None:
|
||||
"""未持有期间不落点(只落持仓日 + 建仓基准点),显著压缩结果体积。"""
|
||||
daily = synthetic_daily({"600000.SH": 0.004, "600001.SH": 0.002}, n=260)
|
||||
res = LocalEngine().run_backtest(daily, _spec(top_n=1, m=6, y=6))
|
||||
all_days = sorted(daily["trade_date"].unique())
|
||||
for curve in res.symbol_curves:
|
||||
point_days = {p.date for p in curve.points}
|
||||
assert point_days <= set(all_days)
|
||||
assert len(point_days) < len(all_days), f"{curve.symbol} 不应逐日落点"
|
||||
@@ -112,11 +112,13 @@ class TestJobExecutor:
|
||||
done = SqlAlchemyJobRepository(session).get("JOB-TEST-1")
|
||||
assert done is not None
|
||||
assert done.status == JobStatus.SUCCESS
|
||||
assert done.result_json is not None
|
||||
# 完整结果只在 experiment 存一份(job 侧不再重复落库)
|
||||
assert done.result_json is None
|
||||
assert done.experiment_id is not None
|
||||
exp = SqlAlchemyExperimentRepository(session).get(done.experiment_id)
|
||||
assert exp is not None
|
||||
assert exp.kind == "backtest"
|
||||
assert exp.result_json # 归档里必须有完整结果
|
||||
assert exp.summary_text is not None
|
||||
assert "总收益" in (exp.summary_text or "")
|
||||
|
||||
|
||||
@@ -0,0 +1,427 @@
|
||||
"""名称变更历史(时点 ST)测试:Repository、Provider 归一、universe 时点口径。
|
||||
|
||||
背景(实测):`stock.name` 只是最新名称快照。用它做 `exclude_st` 会把
|
||||
「曾是高股息、后来才变 ST/退市」的标的在**整段历史**里排除 —— 而那正是
|
||||
「股息陷阱」样本。实测对照(同一 spec 仅改 exclude_st):+35.71% → +32.01%,
|
||||
即约 3.70pp 收益被名称快照口径隐藏。本模块锁定修复后的时点语义。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from app.domain.entities.market import Stock, StockNameHistory
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import (
|
||||
DailyBasicModel,
|
||||
StockDailyModel,
|
||||
StockModel,
|
||||
StockNameHistoryModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyStockNameHistoryRepository,
|
||||
)
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def session(tmp_path):
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'name.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
s = Session(engine)
|
||||
yield s
|
||||
s.close()
|
||||
|
||||
|
||||
def _seed_dima(session: Session) -> None:
|
||||
"""600565.SH 迪马股份:2002-07-23 上市 → 2024-05-06 变 ST迪马(真实数据)。"""
|
||||
session.add_all(
|
||||
[
|
||||
StockModel(
|
||||
symbol="600565.SH", name="ST迪马", list_date=date(2002, 7, 23), status="L"
|
||||
),
|
||||
StockModel(
|
||||
symbol="600519.SH", name="贵州茅台", list_date=date(2001, 8, 27), status="L"
|
||||
),
|
||||
StockNameHistoryModel(
|
||||
symbol="600565.SH", name="迪马股份", start_date=date(2002, 7, 23),
|
||||
end_date=date(2024, 5, 5), ann_date=date(2002, 7, 23), change_reason="其他",
|
||||
),
|
||||
StockNameHistoryModel(
|
||||
symbol="600565.SH", name="ST迪马", start_date=date(2024, 5, 6),
|
||||
end_date=None, ann_date=date(2024, 4, 30), change_reason="ST",
|
||||
),
|
||||
StockNameHistoryModel(
|
||||
symbol="600519.SH", name="贵州茅台", start_date=date(2001, 8, 27),
|
||||
end_date=None, ann_date=date(2001, 8, 27), change_reason="其他",
|
||||
),
|
||||
]
|
||||
)
|
||||
session.commit()
|
||||
|
||||
|
||||
class TestRepository:
|
||||
def test_names_as_of_picks_effective_span(self, session) -> None:
|
||||
_seed_dima(session)
|
||||
repo = SqlAlchemyStockNameHistoryRepository(session)
|
||||
symbols = ["600565.SH", "600519.SH"]
|
||||
# 2020 年:迪马股份(不是 ST)
|
||||
assert repo.names_as_of(symbols, date(2020, 1, 2))["600565.SH"] == "迪马股份"
|
||||
# 区间边界:末日仍为旧名,次日起为新名
|
||||
assert repo.names_as_of(["600565.SH"], date(2024, 5, 5))["600565.SH"] == "迪马股份"
|
||||
assert repo.names_as_of(["600565.SH"], date(2024, 5, 6))["600565.SH"] == "ST迪马"
|
||||
# end_date 为空的区间至今有效
|
||||
assert repo.names_as_of(["600565.SH"], date(2026, 9, 4))["600565.SH"] == "ST迪马"
|
||||
|
||||
def test_unknown_symbol_absent_from_map(self, session) -> None:
|
||||
"""无记录 → 不返回该键,由调用方回退最新名称(不抛错)。"""
|
||||
_seed_dima(session)
|
||||
repo = SqlAlchemyStockNameHistoryRepository(session)
|
||||
assert repo.names_as_of(["000001.SZ"], date(2020, 1, 2)) == {}
|
||||
|
||||
def test_upsert_is_idempotent(self, session) -> None:
|
||||
repo = SqlAlchemyStockNameHistoryRepository(session)
|
||||
row = StockNameHistory(
|
||||
symbol="600565.SH", name="迪马股份", start_date=date(2002, 7, 23),
|
||||
end_date=date(2024, 5, 5), change_reason="其他",
|
||||
)
|
||||
assert repo.upsert_many([row]) == 1
|
||||
assert repo.upsert_many([row]) == 1 # 重跑无副作用
|
||||
assert repo.count_rows() == 1
|
||||
assert repo.namechange_dates() == (date(2002, 7, 23), date(2002, 7, 23))
|
||||
|
||||
def test_name_spans_grouped_by_symbol(self, session) -> None:
|
||||
_seed_dima(session)
|
||||
repo = SqlAlchemyStockNameHistoryRepository(session)
|
||||
spans = repo.name_spans(["600565.SH"])
|
||||
assert [n for _s, _e, n in spans["600565.SH"]] == ["迪马股份", "ST迪马"]
|
||||
|
||||
|
||||
class TestUniversePointInTime:
|
||||
"""filter_stocks(name_at=...) 的时点语义(案例收益口径的关键)。"""
|
||||
|
||||
def _stocks(self) -> list[Stock]:
|
||||
return [
|
||||
Stock(symbol="600565.SH", name="ST迪马", list_date=date(2002, 7, 23)),
|
||||
Stock(symbol="600519.SH", name="贵州茅台", list_date=date(2001, 8, 27)),
|
||||
]
|
||||
|
||||
def test_name_at_overrides_snapshot(self) -> None:
|
||||
from app.domain.entities.research import UniverseSpec
|
||||
from app.quant.universe import filter_stocks
|
||||
|
||||
u = UniverseSpec(exclude_st=True, min_listing_days=0)
|
||||
# 旧口径(无 name_at):最新名称含 ST → 2020 年就被排除(股息陷阱被隐藏)
|
||||
assert "600565.SH" not in {s.symbol for s in filter_stocks(self._stocks(), u, date(2020, 1, 2))}
|
||||
# 时点口径:2020 年它叫「迪马股份」→ 必须纳入
|
||||
got = {
|
||||
s.symbol
|
||||
for s in filter_stocks(
|
||||
self._stocks(), u, date(2020, 1, 2), name_at={"600565.SH": "迪马股份"}
|
||||
)
|
||||
}
|
||||
assert "600565.SH" in got
|
||||
# 2024-05-06 起为 ST迪马 → 排除
|
||||
got2 = {
|
||||
s.symbol
|
||||
for s in filter_stocks(
|
||||
self._stocks(), u, date(2024, 6, 3), name_at={"600565.SH": "ST迪马"}
|
||||
)
|
||||
}
|
||||
assert "600565.SH" not in got2
|
||||
|
||||
def test_name_at_missing_falls_back_to_snapshot(self) -> None:
|
||||
from app.domain.entities.research import UniverseSpec
|
||||
from app.quant.universe import filter_stocks
|
||||
|
||||
u = UniverseSpec(exclude_st=True, min_listing_days=0)
|
||||
got = filter_stocks(self._stocks(), u, date(2020, 1, 2), name_at={})
|
||||
assert "600565.SH" not in {s.symbol for s in got} # 回退快照,行为不变
|
||||
|
||||
def test_exclude_st_false_ignores_names(self) -> None:
|
||||
from app.domain.entities.research import UniverseSpec
|
||||
from app.quant.universe import filter_stocks
|
||||
|
||||
u = UniverseSpec(exclude_st=False, min_listing_days=0)
|
||||
got = filter_stocks(self._stocks(), u, date(2020, 1, 2), name_at={"600565.SH": "ST迪马"})
|
||||
assert "600565.SH" in {s.symbol for s in got}
|
||||
|
||||
|
||||
class TestNamesAsOfHelper:
|
||||
def test_none_repo_reports_snapshot_basis(self) -> None:
|
||||
from app.quant.universe import names_as_of
|
||||
|
||||
name_at, applied = names_as_of([], date(2020, 1, 2), None)
|
||||
assert name_at is None
|
||||
assert applied == (False, 0)
|
||||
|
||||
def test_repo_failure_degrades_without_raising(self) -> None:
|
||||
"""名称历史查询异常不得让选股/回测整体失败(降级为快照口径)。"""
|
||||
from app.quant.universe import names_as_of
|
||||
|
||||
class Boom:
|
||||
def names_as_of(self, symbols, as_of):
|
||||
raise RuntimeError("表不存在")
|
||||
|
||||
name_at, applied = names_as_of([], date(2020, 1, 2), Boom())
|
||||
assert name_at is None and applied == (False, 0)
|
||||
|
||||
|
||||
class TestProviderNormalize:
|
||||
def test_normalize_name_history(self) -> None:
|
||||
from app.infrastructure.data_sources.tushare import TushareProvider
|
||||
|
||||
rows = TushareProvider.normalize_name_history(
|
||||
[
|
||||
{
|
||||
"ts_code": "600565.SH", "name": "ST迪马", "start_date": "20240506",
|
||||
"end_date": None, "ann_date": "20240430", "change_reason": "ST",
|
||||
},
|
||||
{
|
||||
"ts_code": "600565.SH", "name": "迪马股份", "start_date": "20020723",
|
||||
"end_date": "20240505", "ann_date": "20020723", "change_reason": "其他",
|
||||
},
|
||||
]
|
||||
)
|
||||
assert rows[1].start_date == date(2002, 7, 23)
|
||||
assert rows[0].end_date is None
|
||||
assert rows[0].is_risk_warned is True
|
||||
assert rows[1].is_risk_warned is False
|
||||
|
||||
def test_nan_end_date_and_bad_code_skipped(self) -> None:
|
||||
"""实测:namechange 的 end_date 是 float NaN,曾让 32/37 个分片整体失败。"""
|
||||
from app.infrastructure.data_sources.tushare import TushareProvider
|
||||
|
||||
rows = TushareProvider.normalize_name_history(
|
||||
[
|
||||
{
|
||||
"ts_code": "000001.SZ", "name": "平安银行", "start_date": "19910403",
|
||||
"end_date": float("nan"), "ann_date": float("nan"), "change_reason": None,
|
||||
},
|
||||
{"ts_code": "T600018.SH", "name": "上港集箱(退)", "start_date": "19960101"},
|
||||
]
|
||||
)
|
||||
assert len(rows) == 1
|
||||
assert rows[0].end_date is None and rows[0].ann_date is None
|
||||
|
||||
def test_to_date_rejects_nan_variants(self) -> None:
|
||||
from app.infrastructure.data_sources.tushare import _to_date
|
||||
|
||||
assert _to_date(float("nan")) is None
|
||||
assert _to_date("nan") is None
|
||||
assert _to_date("None") is None
|
||||
assert _to_date("") is None
|
||||
assert _to_date(None) is None
|
||||
assert _to_date("20240506") == date(2024, 5, 6)
|
||||
|
||||
|
||||
class TestPerDateStFilter:
|
||||
"""回测的 exclude_st 必须**逐择股日**重判(与 /api/selections 单时点同口径)。
|
||||
|
||||
否则「入池时非 ST、之后才变 ST」的标的会在之后所有择股日继续被选中 ——
|
||||
正是高股息策略最危险的股息陷阱路径。
|
||||
"""
|
||||
|
||||
def _build(self, tmp_path):
|
||||
"""A 高股息但 2024-06-01 变 ST;B 低股息且始终非 ST。"""
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyDailyBasicRepository,
|
||||
SqlAlchemyStockNameHistoryRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'pit.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
syms = ["600000.SH", "600001.SH"]
|
||||
session.add_all(
|
||||
[
|
||||
StockModel(symbol=syms[0], name="ST阿甲", list_date=date(2000, 1, 1), status="L"),
|
||||
StockModel(symbol=syms[1], name="阿乙", list_date=date(2000, 1, 1), status="L"),
|
||||
# 名称历史:A 2024-06-01 起为 ST阿甲
|
||||
StockNameHistoryModel(
|
||||
symbol=syms[0], name="阿甲", start_date=date(2000, 1, 1),
|
||||
end_date=date(2024, 5, 31), change_reason="其他",
|
||||
),
|
||||
StockNameHistoryModel(
|
||||
symbol=syms[0], name="ST阿甲", start_date=date(2024, 6, 1),
|
||||
end_date=None, change_reason="ST",
|
||||
),
|
||||
StockNameHistoryModel(
|
||||
symbol=syms[1], name="阿乙", start_date=date(2000, 1, 1),
|
||||
end_date=None, change_reason="其他",
|
||||
),
|
||||
]
|
||||
)
|
||||
days = pd.bdate_range("2023-12-01", "2024-12-31")
|
||||
for sym, dv in ((syms[0], 12.0), (syms[1], 4.0)):
|
||||
session.add_all(
|
||||
[
|
||||
StockDailyModel(
|
||||
symbol=sym, trade_date=d.date(), open=Decimal("10"),
|
||||
high=Decimal("10"), low=Decimal("10"), close=Decimal("10"),
|
||||
volume=Decimal("1000"), amount=Decimal("10000"), source="tushare",
|
||||
)
|
||||
for d in days
|
||||
]
|
||||
)
|
||||
session.add_all(
|
||||
[
|
||||
DailyBasicModel(symbol=sym, trade_date=d.date(), dv_ratio=Decimal(str(dv)))
|
||||
for d in days
|
||||
]
|
||||
)
|
||||
session.commit()
|
||||
from app.quant.engine import LocalEngine
|
||||
from app.quant.service import ResearchService
|
||||
|
||||
svc = ResearchService(
|
||||
SqlAlchemyStockRepository(session),
|
||||
SqlAlchemyDailyBarRepository(session),
|
||||
LocalEngine(),
|
||||
basic_repo=SqlAlchemyDailyBasicRepository(session),
|
||||
name_repo=SqlAlchemyStockNameHistoryRepository(session),
|
||||
)
|
||||
return svc, session, syms
|
||||
|
||||
def _spec(self):
|
||||
from app.domain.entities.research import (
|
||||
FactorSpec,
|
||||
ResearchSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
|
||||
return ResearchSpec(
|
||||
type="backtest",
|
||||
universe=UniverseSpec(exclude_st=True, min_listing_days=0),
|
||||
factors=[FactorSpec(name="dividend_yield", weight=1.0)],
|
||||
selection=SelectionSpec(top_n=2, hold_top_x=1, allow_substitute=True),
|
||||
period=("2024-01-01", "2024-12-31"),
|
||||
selection_interval_months=6,
|
||||
rebalance_interval_months=6,
|
||||
initial_capital=100000.0,
|
||||
)
|
||||
|
||||
def test_st_stock_dropped_from_later_selection(self, tmp_path) -> None:
|
||||
svc, session, syms = self._build(tmp_path)
|
||||
try:
|
||||
result = svc.run_backtest(self._spec())
|
||||
finally:
|
||||
session.close()
|
||||
picks = {}
|
||||
for p in result.selection_history:
|
||||
picks.setdefault(p.date, []).append(p.symbol)
|
||||
dates = sorted(picks)
|
||||
assert len(dates) >= 2, picks
|
||||
first, second = dates[0], dates[-1]
|
||||
# 首个择股日:A 非 ST(高股息)→ 入选
|
||||
assert syms[0] in picks[first]
|
||||
# 变 ST 之后的择股日:A 必须消失,B 顶上
|
||||
assert syms[0] not in picks[second]
|
||||
assert syms[1] in picks[second]
|
||||
# 口径标注:时点名称(不是快照回退)
|
||||
assert result.config_snapshot["price_basis"]["name_basis"]["point_in_time"] is True
|
||||
|
||||
def test_snapshot_fallback_without_repo(self, tmp_path) -> None:
|
||||
"""未注入名称历史 → 回退最新名称(旧行为),并如实标注 point_in_time=False。"""
|
||||
svc, session, syms = self._build(tmp_path)
|
||||
try:
|
||||
svc._name_repo = None
|
||||
result = svc.run_backtest(self._spec())
|
||||
finally:
|
||||
session.close()
|
||||
# 最新名称是 ST阿甲 → 首个择股日就被排除(股息陷阱被隐藏,已标注)
|
||||
first = min(p.date for p in result.selection_history)
|
||||
assert syms[0] not in [p.symbol for p in result.selection_history if p.date == first]
|
||||
assert result.config_snapshot["price_basis"]["name_basis"]["point_in_time"] is False
|
||||
|
||||
|
||||
class TestReviewRegressions:
|
||||
"""代码审查发现的缺陷回归(P1/P2/P3):不得复活。"""
|
||||
|
||||
def test_empty_table_does_not_claim_point_in_time(self) -> None:
|
||||
"""[P1] 表存在但为空时必须降级为快照口径,不得声称 point_in_time=true。
|
||||
|
||||
否则结果页会把「未被修正的 10.85pp 股息陷阱偏差」当成已修正上报(AGENT.md §24)。
|
||||
"""
|
||||
from app.quant.universe import names_as_of
|
||||
|
||||
class EmptyRepo:
|
||||
def names_as_of(self, symbols, as_of):
|
||||
return {}
|
||||
|
||||
stocks = [Stock(symbol="600565.SH", name="ST迪马", list_date=date(2000, 1, 1))]
|
||||
name_at, applied = names_as_of(stocks, date(2020, 1, 2), EmptyRepo())
|
||||
assert name_at is None
|
||||
assert applied == (False, 0)
|
||||
|
||||
def test_empty_table_end_to_end_reports_snapshot_basis(self, tmp_path) -> None:
|
||||
"""[P1] 端到端:名称表为空 → name_basis.point_in_time 必须为 False。"""
|
||||
svc, session, _syms = TestPerDateStFilter()._build(tmp_path)
|
||||
try:
|
||||
session.query(StockNameHistoryModel).delete()
|
||||
session.commit()
|
||||
result = svc.run_backtest(TestPerDateStFilter()._spec())
|
||||
finally:
|
||||
session.close()
|
||||
assert result.config_snapshot["price_basis"]["name_basis"]["point_in_time"] is False
|
||||
assert any("最新名称快照" in n for n in result.unimplemented)
|
||||
|
||||
def test_missing_span_falls_back_to_snapshot_in_backtest(self, tmp_path) -> None:
|
||||
"""[P2] 某股无生效区间时,逐择股日 ST 必须与 filter_stocks 同口径回退最新名称。
|
||||
|
||||
否则同一择股日会出现「回测入选、/api/selections 排除」的打架(v2 §25)。
|
||||
"""
|
||||
svc, session, syms = TestPerDateStFilter()._build(tmp_path)
|
||||
try:
|
||||
# 删掉 A 的**全部**名称区间:只剩最新名称 ST阿甲
|
||||
session.query(StockNameHistoryModel).filter(
|
||||
StockNameHistoryModel.symbol == syms[0]
|
||||
).delete()
|
||||
session.commit()
|
||||
result = svc.run_backtest(TestPerDateStFilter()._spec())
|
||||
dates = sorted({p.date for p in result.selection_history})
|
||||
first = dates[0]
|
||||
picks_first = [p.symbol for p in result.selection_history if p.date == first]
|
||||
finally:
|
||||
session.close()
|
||||
# 快照口径:A 名称含 ST → 首个择股日即被排除(与 filter_stocks 一致)
|
||||
assert syms[0] not in picks_first
|
||||
assert syms[1] in picks_first
|
||||
|
||||
def test_protocol_declares_name_changes(self) -> None:
|
||||
"""[P2] MarketDataProvider 必须声明 get_name_changes(§6 业务层只依赖抽象)。"""
|
||||
|
||||
from app.domain.providers import MarketDataProvider
|
||||
|
||||
assert hasattr(MarketDataProvider, "get_name_changes")
|
||||
assert "get_name_changes" in dir(MarketDataProvider)
|
||||
|
||||
def test_sina_declares_not_supported(self) -> None:
|
||||
"""[P2] 备用源必须显式 NotSupported,不得静默返回空列表(否则时点 ST 静默降级)。"""
|
||||
from app.infrastructure.data_sources.errors import DataSourceNotSupported
|
||||
from app.infrastructure.data_sources.sina import SinaProvider
|
||||
|
||||
provider = SinaProvider.__new__(SinaProvider)
|
||||
with pytest.raises(DataSourceNotSupported):
|
||||
provider.get_name_changes(date(2024, 1, 1), date(2024, 12, 31))
|
||||
|
||||
def test_qlib_engine_accepts_eligibility_fn(self) -> None:
|
||||
"""[P2] 引擎协议一致性:所有引擎都必须接受 eligibility_fn。
|
||||
|
||||
否则注入 QlibEngine 后任何回测都会 TypeError(且条件/时点 ST 会被静默忽略)。
|
||||
"""
|
||||
import inspect
|
||||
|
||||
from app.quant.engine import QuantEngine
|
||||
from app.quant.qlib_adapter.engine import QlibEngine
|
||||
|
||||
for cls in (QuantEngine, QlibEngine):
|
||||
params = inspect.signature(cls.run_backtest).parameters
|
||||
assert "eligibility_fn" in params, cls.__name__
|
||||
@@ -0,0 +1,69 @@
|
||||
"""归档保留策略(`app.cli.prune_experiments`)的候选选择单测。
|
||||
|
||||
删除归档是**不可恢复**操作,因此候选集合的判定必须被测试钉住:多删一条就是丢了
|
||||
一份研究结论,少删一条则保留策略失效。这里只测纯函数与参数护栏(不连数据库)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import pytest
|
||||
from app.cli.prune_experiments import _parse_args, select_prune_candidates
|
||||
|
||||
|
||||
class _Row:
|
||||
"""最小归档摘要替身(只有 id / created_at 参与判定)。"""
|
||||
|
||||
def __init__(self, id_: str, age_days: float | None) -> None:
|
||||
self.id = id_
|
||||
self.created_at = None if age_days is None else NOW - timedelta(days=age_days)
|
||||
|
||||
|
||||
NOW = datetime(2026, 9, 20, 12, 0, 0)
|
||||
|
||||
|
||||
def _rows() -> list[_Row]:
|
||||
return [_Row("E1", 1), _Row("E2", 2), _Row("E3", 10), _Row("E4", 30), _Row("E5", None)]
|
||||
|
||||
|
||||
class TestSelectPruneCandidates:
|
||||
def test_keep_newest_n(self) -> None:
|
||||
got = [e.id for e in select_prune_candidates(_rows(), keep=2, older_than_days=None, now=NOW)]
|
||||
assert got == ["E3", "E4", "E5"], "只保留最新 2 条(E1/E2),其余都是候选"
|
||||
|
||||
def test_keep_zero_deletes_all(self) -> None:
|
||||
got = [e.id for e in select_prune_candidates(_rows(), keep=0, older_than_days=None, now=NOW)]
|
||||
assert set(got) == {"E1", "E2", "E3", "E4", "E5"}
|
||||
|
||||
def test_older_than_only(self) -> None:
|
||||
got = [e.id for e in select_prune_candidates(_rows(), keep=None, older_than_days=5, now=NOW)]
|
||||
assert got == ["E3", "E4", "E5"], "5 天内创建的 E1/E2 必须保留"
|
||||
|
||||
def test_intersection_of_both_conditions(self) -> None:
|
||||
got = [e.id for e in select_prune_candidates(_rows(), keep=4, older_than_days=5, now=NOW)]
|
||||
# older_than 选出 E3/E4/E5,keep=4 保留 E1..E4 → 交集只剩 E5
|
||||
assert got == ["E5"]
|
||||
|
||||
def test_missing_created_at_treated_as_oldest(self) -> None:
|
||||
"""created_at 为空按最旧处理:不能被当作「最新」从而永久留库。"""
|
||||
got = [e.id for e in select_prune_candidates(_rows(), keep=4, older_than_days=None, now=NOW)]
|
||||
assert "E5" in got
|
||||
|
||||
def test_result_is_newest_first(self) -> None:
|
||||
got = [e.id for e in select_prune_candidates(_rows(), keep=0, older_than_days=None, now=NOW)]
|
||||
assert got == ["E1", "E2", "E3", "E4", "E5"], "输出按新→旧,便于人工核对"
|
||||
|
||||
|
||||
class TestArgGuard:
|
||||
def test_requires_keep_or_older_than(self) -> None:
|
||||
with pytest.raises(SystemExit):
|
||||
_parse_args([]) # 两个条件都不给 → 拒绝执行(避免误删全部)
|
||||
|
||||
def test_negative_keep_rejected(self) -> None:
|
||||
with pytest.raises(SystemExit):
|
||||
_parse_args(["--keep", "-1"])
|
||||
|
||||
def test_default_is_dry_run(self) -> None:
|
||||
args = _parse_args(["--keep", "10"])
|
||||
assert args.apply is False, "默认必须 dry-run,只有显式 --apply 才写库"
|
||||
@@ -0,0 +1,222 @@
|
||||
"""`app.cli.restore_experiment_from_job` 测试:从 Job 结果副本重建归档。
|
||||
|
||||
为什么要测:删除归档是正常功能,但**历史归档在 `job.result_json` 里另存了一份完整结果**
|
||||
(完整存档上线前的双写遗留),所以「删了能不能救回来」是有确定答案的:
|
||||
历史归档能按原 id 重建,新归档(`job.result_json IS NULL`)不能。
|
||||
这个工具承载该结论,必须验证真写库路径与全部拒绝路径,而不是只做 dry-run。
|
||||
|
||||
测试全部在 conftest 强制的 /tmp sqlite 上跑,不触碰真实 MySQL。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import date, datetime
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.cli import restore_experiment_from_job as rj
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
BacktestSummary,
|
||||
CurvePoint,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.models.jobs import ExperimentModel, JobModel
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
|
||||
def _payload(final_equity: float = 1_100_000.0) -> str:
|
||||
"""一份最小可用的回测结果 JSON(与归档里存的结构同源)。"""
|
||||
result = BacktestResult(
|
||||
summary=BacktestSummary(
|
||||
start=date(2024, 1, 2),
|
||||
end=date(2024, 6, 28),
|
||||
initial_capital=1_000_000.0,
|
||||
final_equity=final_equity,
|
||||
total_return_pct=10.0,
|
||||
annual_return_pct=21.5,
|
||||
sharpe=1.2,
|
||||
max_drawdown_pct=-8.5,
|
||||
volatility_pct=18.0,
|
||||
win_rate_pct=55.0,
|
||||
total_trades=12,
|
||||
avg_turnover_pct=30.0,
|
||||
),
|
||||
equity_curve=[CurvePoint(date=date(2024, 1, 2), value=1_000_000.0)],
|
||||
drawdown=[],
|
||||
monthly_returns=[],
|
||||
yearly_returns=[],
|
||||
positions=[],
|
||||
trades=[],
|
||||
turnover_pct=30.0,
|
||||
)
|
||||
return json.dumps(result.model_dump(mode="json"), ensure_ascii=False)
|
||||
|
||||
|
||||
def _make_session(tmp_path):
|
||||
engine = create_engine(f"sqlite:///{tmp_path}/restore.db", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
return sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
|
||||
|
||||
|
||||
def _seed_job(session, *, job_id="JOB-TEST", exp_id="EXP-TEST", result_json=None):
|
||||
session.add(
|
||||
JobModel(
|
||||
id=job_id,
|
||||
kind="backtest",
|
||||
status="success",
|
||||
spec_json=json.dumps({"type": "backtest", "factors": []}),
|
||||
result_json=result_json,
|
||||
experiment_id=exp_id,
|
||||
created_at=datetime(2026, 1, 5, 9, 0, 0),
|
||||
finished_at=datetime(2026, 1, 5, 9, 0, 1),
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
|
||||
|
||||
def _run(session_factory, argv):
|
||||
with patch.object(rj, "SessionLocal", session_factory):
|
||||
return rj.main(argv)
|
||||
|
||||
|
||||
def test_dry_run_writes_nothing(tmp_path, capsys) -> None:
|
||||
factory = _make_session(tmp_path)
|
||||
with factory() as s:
|
||||
_seed_job(s, result_json=_payload())
|
||||
code = _run(factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234"])
|
||||
out = capsys.readouterr().out
|
||||
assert code == 0
|
||||
assert "(dry-run)" in out
|
||||
with factory() as s:
|
||||
assert s.get(ExperimentModel, "EXP-TEST") is None
|
||||
|
||||
|
||||
def test_restores_archive_from_job_copy(tmp_path, capsys) -> None:
|
||||
"""真写库路径:原 id、逐字复制、摘要按归档口径重算、data_version 留空。"""
|
||||
factory = _make_session(tmp_path)
|
||||
payload = _payload()
|
||||
with factory() as s:
|
||||
_seed_job(s, result_json=payload)
|
||||
code = _run(
|
||||
factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234", "--apply"]
|
||||
)
|
||||
assert code == 0
|
||||
assert "已重建" in capsys.readouterr().out
|
||||
with factory() as s:
|
||||
row = s.get(ExperimentModel, "EXP-TEST")
|
||||
assert row is not None
|
||||
assert row.kind == "backtest"
|
||||
assert row.result_json == payload # 逐字复制,不重新计算
|
||||
assert row.job_id == "JOB-TEST"
|
||||
assert row.code_version == "abc1234"
|
||||
# 历史归档当年没有数据指纹:留空而不是补今天的(伪造复现依据)
|
||||
assert row.data_version is None
|
||||
assert row.created_at == datetime(2026, 1, 5, 9, 0, 1)
|
||||
assert row.summary_text == "总收益 10.00% · 年化 21.50% · 回撤 -8.50%"
|
||||
|
||||
|
||||
def test_refuses_when_archive_already_exists(tmp_path, capsys) -> None:
|
||||
factory = _make_session(tmp_path)
|
||||
with factory() as s:
|
||||
_seed_job(s, result_json=_payload())
|
||||
s.add(
|
||||
ExperimentModel(
|
||||
id="EXP-TEST",
|
||||
kind="backtest",
|
||||
spec_json="{}",
|
||||
result_json=_payload(1.0),
|
||||
summary_text="原有",
|
||||
code_version="old",
|
||||
data_version=None,
|
||||
job_id="JOB-TEST",
|
||||
created_at=datetime(2026, 1, 5, 9, 0, 1),
|
||||
)
|
||||
)
|
||||
s.commit()
|
||||
code = _run(
|
||||
factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234", "--apply"]
|
||||
)
|
||||
assert code == 1
|
||||
assert "已存在" in capsys.readouterr().err
|
||||
with factory() as s:
|
||||
assert s.get(ExperimentModel, "EXP-TEST").summary_text == "原有" # 未被改动
|
||||
|
||||
|
||||
def test_force_overwrites_existing_archive(tmp_path) -> None:
|
||||
factory = _make_session(tmp_path)
|
||||
payload = _payload()
|
||||
with factory() as s:
|
||||
_seed_job(s, result_json=payload)
|
||||
s.add(
|
||||
ExperimentModel(
|
||||
id="EXP-TEST",
|
||||
kind="backtest",
|
||||
spec_json="{}",
|
||||
result_json="{}",
|
||||
summary_text="旧的",
|
||||
code_version="old",
|
||||
data_version="d20260101",
|
||||
job_id="JOB-TEST",
|
||||
created_at=datetime(2026, 1, 5, 9, 0, 1),
|
||||
)
|
||||
)
|
||||
s.commit()
|
||||
code = _run(
|
||||
factory,
|
||||
["--job-id", "JOB-TEST", "--code-version", "abc1234", "--apply", "--force"],
|
||||
)
|
||||
assert code == 0
|
||||
with factory() as s:
|
||||
row = s.get(ExperimentModel, "EXP-TEST")
|
||||
assert row.result_json == payload
|
||||
assert row.code_version == "abc1234"
|
||||
assert row.data_version is None # 覆盖后也留空,不继承旧指纹
|
||||
|
||||
|
||||
def test_refuses_new_archive_without_job_copy(tmp_path, capsys) -> None:
|
||||
"""完整存档上线后结果只存归档一份 → 删了不可恢复,工具必须如实拒绝。"""
|
||||
factory = _make_session(tmp_path)
|
||||
with factory() as s:
|
||||
_seed_job(s, result_json=None)
|
||||
code = _run(
|
||||
factory, ["--job-id", "JOB-TEST", "--code-version", "abc1234", "--apply"]
|
||||
)
|
||||
assert code == 1
|
||||
err = capsys.readouterr().err
|
||||
assert "result_json 为空" in err
|
||||
assert "不可恢复" in err
|
||||
|
||||
|
||||
def test_requires_explicit_code_version(tmp_path, capsys) -> None:
|
||||
"""不猜版本:缺失时直接拒绝(退出码 2),避免写进一个编造的复现依据。"""
|
||||
factory = _make_session(tmp_path)
|
||||
with factory() as s:
|
||||
_seed_job(s, result_json=_payload())
|
||||
code = _run(factory, ["--job-id", "JOB-TEST", "--apply"])
|
||||
assert code == 2
|
||||
assert "code-version" in capsys.readouterr().err
|
||||
|
||||
|
||||
def test_empty_code_version_is_allowed_as_unknown(tmp_path) -> None:
|
||||
"""显式传空串=如实表示"版本未知",此时写 NULL 而不是空字符串。"""
|
||||
factory = _make_session(tmp_path)
|
||||
with factory() as s:
|
||||
_seed_job(s, result_json=_payload())
|
||||
code = _run(
|
||||
factory, ["--job-id", "JOB-TEST", "--code-version", "", "--apply"]
|
||||
)
|
||||
assert code == 0
|
||||
with factory() as s:
|
||||
assert s.get(ExperimentModel, "EXP-TEST").code_version is None
|
||||
|
||||
|
||||
def test_rejects_unknown_job_and_job_without_experiment_id(tmp_path, capsys) -> None:
|
||||
factory = _make_session(tmp_path)
|
||||
with factory() as s:
|
||||
_seed_job(s, job_id="JOB-NOEXP", exp_id=None, result_json=_payload())
|
||||
assert _run(factory, ["--job-id", "JOB-MISSING", "--code-version", "x"]) == 1
|
||||
assert "不存在" in capsys.readouterr().err
|
||||
assert _run(factory, ["--job-id", "JOB-NOEXP", "--code-version", "x"]) == 1
|
||||
assert "experiment_id" in capsys.readouterr().err
|
||||
@@ -0,0 +1,86 @@
|
||||
"""run_job 研究子进程入口的平台/配置分支单测。
|
||||
|
||||
覆盖点(Review 补测):
|
||||
- `_apply_memory_limit` 成功路径确实施加了请求的上限;
|
||||
- 环境变量缺失时回退 `Settings.job_memory_limit_gb`;
|
||||
- setrlimit 失败(如 macOS 上限额低于进程虚拟地址空间基线 → EINVAL/ValueError)
|
||||
**不得抛出**,否则研究子进程启动即失败;
|
||||
- `_subprocess_log_target` 落盘失败时退回 DEVNULL —— 记日志失败不能拖垮 Job。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import resource
|
||||
|
||||
import pytest
|
||||
from app.application.services import job_executor as je
|
||||
from app.cli import run_job
|
||||
|
||||
|
||||
def test_apply_memory_limit_sets_requested_value(monkeypatch) -> None:
|
||||
captured: dict = {}
|
||||
monkeypatch.setenv("QLIB_JOB_MEM_LIMIT_GB", "3")
|
||||
monkeypatch.setattr(resource, "setrlimit", lambda which, lim: captured.update(w=which, lim=lim))
|
||||
|
||||
run_job._apply_memory_limit()
|
||||
|
||||
assert captured["w"] == resource.RLIMIT_AS
|
||||
assert captured["lim"] == (3 * 1024**3, 3 * 1024**3)
|
||||
|
||||
|
||||
def test_apply_memory_limit_falls_back_to_settings(monkeypatch) -> None:
|
||||
"""未设 QLIB_JOB_MEM_LIMIT_GB 时用 Settings 的上限(config.yaml job.max_memory_gb)。"""
|
||||
from app.core.config import get_settings
|
||||
|
||||
captured: dict = {}
|
||||
monkeypatch.delenv("QLIB_JOB_MEM_LIMIT_GB", raising=False)
|
||||
monkeypatch.setattr(resource, "setrlimit", lambda which, lim: captured.update(lim=lim))
|
||||
|
||||
run_job._apply_memory_limit()
|
||||
|
||||
gb = get_settings().job_memory_limit_gb
|
||||
assert captured["lim"] == (gb * 1024**3, gb * 1024**3)
|
||||
|
||||
|
||||
def test_apply_memory_limit_failure_does_not_raise(monkeypatch, capsys) -> None:
|
||||
"""setrlimit 失败(macOS 限额低于地址空间基线的典型情形)只提示、不抛出。"""
|
||||
|
||||
def _boom(which, lim): # noqa: ANN001
|
||||
raise ValueError("current limit exceeds maximum limit")
|
||||
|
||||
monkeypatch.setenv("QLIB_JOB_MEM_LIMIT_GB", "6")
|
||||
monkeypatch.setattr(resource, "setrlimit", _boom)
|
||||
|
||||
run_job._apply_memory_limit() # 不应抛异常
|
||||
|
||||
err = capsys.readouterr().err
|
||||
assert "内存上限 6GB 设置失败" in err
|
||||
assert "无内存隔离保护" in err
|
||||
|
||||
|
||||
def test_subprocess_log_target_falls_back_to_devnull(tmp_path, monkeypatch) -> None:
|
||||
"""日志目录不可创建时退回 DEVNULL,不影响 Job 执行。"""
|
||||
blocker = tmp_path / "not-a-dir"
|
||||
blocker.write_text("x", encoding="utf-8")
|
||||
monkeypatch.setattr(je, "BACKEND_ROOT", blocker / "sub")
|
||||
|
||||
assert je._subprocess_log_target() is je.subprocess.DEVNULL
|
||||
|
||||
|
||||
def test_subprocess_log_target_writes_under_project_logs(tmp_path, monkeypatch) -> None:
|
||||
"""正常路径:在 <项目根>/logs/job-subprocess.log 追加写入。"""
|
||||
monkeypatch.setattr(je, "BACKEND_ROOT", tmp_path / "backend")
|
||||
|
||||
target = je._subprocess_log_target()
|
||||
try:
|
||||
target.write("hello-subprocess\n")
|
||||
finally:
|
||||
target.close()
|
||||
|
||||
assert (tmp_path / "logs" / "job-subprocess.log").read_text(encoding="utf-8") == (
|
||||
"hello-subprocess\n"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__": # pragma: no cover
|
||||
raise SystemExit(pytest.main([__file__, "-q"]))
|
||||
@@ -76,9 +76,11 @@ class TestSelectionJobExecutor:
|
||||
job = SqlAlchemyJobRepository(session).get("JOB-SEL-1")
|
||||
exp = SqlAlchemyExperimentRepository(session).get(job.experiment_id or "")
|
||||
assert job.status == JobStatus.SUCCESS
|
||||
result = SelectionResult.model_validate_json(job.result_json or "{}")
|
||||
assert len(result.candidates) == 3
|
||||
# 完整结果只在 experiment 存一份:从归档解码验证(job.result_json 为新形态空值)
|
||||
assert job.result_json is None
|
||||
assert exp is not None and exp.kind == "selection"
|
||||
result = SelectionResult.model_validate_json(exp.result_json)
|
||||
assert len(result.candidates) == 3
|
||||
assert "as_of" in (exp.summary_text or "") and "选出 3" in (exp.summary_text or "")
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,544 @@
|
||||
"""策略说明书生成器(app.quant.strategy_doc)测试 + 策略说明/公式/编辑接口测试。
|
||||
|
||||
目标(AGENT.md §24/§31):说明书必须由 spec **真实推导**,并且
|
||||
- 条件表达式的渲染必须与 `app.quant.selection._eval_condition` 的实际求值一致
|
||||
(含 ref 语义、缺失值语义、in/not_in 的列表语义);
|
||||
- 步骤必须与 `app.quant.local_engine` 的真实行为一致(择股日/调仓日/顺延/期末);
|
||||
- 未知因子、未建模约束一律进 warnings(不编造、不假装支持)。
|
||||
|
||||
同一文件还覆盖策略 API 的说明接口与「原地更新(PUT)」,因为三者共用同一契约。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from app.api import deps
|
||||
from app.domain.entities.research import (
|
||||
ConditionSpec,
|
||||
CostSpec,
|
||||
FactorSpec,
|
||||
ResearchSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.main import app
|
||||
from app.quant.factors import FactorDef
|
||||
from app.quant.selection import build_condition_fields, eligible_symbols
|
||||
from app.quant.strategy_doc import describe_strategy
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from conftest_quant import synthetic_daily
|
||||
|
||||
|
||||
def _spec(**kw) -> ResearchSpec:
|
||||
"""高股息案例参数(n=30 → x=20、m=y=6、hfq、显式成本与最低佣金)。"""
|
||||
base = dict(
|
||||
type="backtest",
|
||||
universe=UniverseSpec(exclude_st=False, min_listing_days=250),
|
||||
price_adjustment="hfq",
|
||||
factors=[FactorSpec(name="dividend_yield", weight=1.0)],
|
||||
conditions=[ConditionSpec(field="dv_ratio", op="lte", value=30)],
|
||||
selection=SelectionSpec(top_n=30, hold_top_x=20, allow_substitute=False, defer_buy=True),
|
||||
rebalance="monthly",
|
||||
selection_interval_months=6,
|
||||
rebalance_interval_months=6,
|
||||
period=(date(2020, 1, 1), date(2024, 12, 31)),
|
||||
costs=CostSpec(min_commission=5.0),
|
||||
)
|
||||
base.update(kw)
|
||||
return ResearchSpec(**base)
|
||||
|
||||
|
||||
class TestSummary:
|
||||
def test_summary_contains_factors_truncation_cycles_and_costs(self) -> None:
|
||||
doc = describe_strategy(_spec())
|
||||
s = doc.summary
|
||||
assert "股息率" in s, "summary 必须由因子元数据推导出因子含义"
|
||||
assert "dividend_yield" in s
|
||||
assert "前 30 只候选池" in s and "前 20 只等权" in s, "必须说清两级截断 n → x"
|
||||
assert "每 6 个月重新择股" in s and "每 6 个月调仓" in s, "必须说清 m / y"
|
||||
assert "后复权" in s and "收盘价成交" in s
|
||||
assert "佣金 0.03%" in s and "印花税 0.05%" in s and "滑点 0.1%" in s
|
||||
|
||||
def test_summary_lower_is_better_factor_方向(self) -> None:
|
||||
doc = describe_strategy(_spec(factors=[FactorSpec(name="volatility_20", weight=1.0)]))
|
||||
assert "从低到高排序" in doc.summary
|
||||
|
||||
def test_summary_without_hold_top_x_uses_n(self) -> None:
|
||||
doc = describe_strategy(
|
||||
_spec(selection=SelectionSpec(top_n=20, allow_substitute=False), conditions=[])
|
||||
)
|
||||
assert "前 20 只等权持有" in doc.summary
|
||||
assert "先通过" not in doc.summary, "无 conditions 时不得声称有过滤条件"
|
||||
|
||||
def test_summary_states_buy_fallback_mode(self) -> None:
|
||||
"""买不进的处置决定实际持仓,必须出现在一句话说明里(三种模式各一)。"""
|
||||
defer = describe_strategy(_spec(conditions=[]))
|
||||
assert "买不进则顺延买入" in defer.summary
|
||||
subst = describe_strategy(
|
||||
_spec(selection=SelectionSpec(top_n=5, allow_substitute=True), conditions=[])
|
||||
)
|
||||
assert "从候选池之外替补" in subst.summary
|
||||
none = describe_strategy(
|
||||
_spec(selection=SelectionSpec(top_n=5, allow_substitute=False), conditions=[])
|
||||
)
|
||||
assert "买不进则放弃" in none.summary
|
||||
|
||||
|
||||
class TestFormula:
|
||||
def test_formula_has_factor_name_meaning_and_formula(self) -> None:
|
||||
doc = describe_strategy(_spec())
|
||||
f = doc.formula
|
||||
assert "dividend_yield" in f
|
||||
# 因子元数据(AGENT.md §22)逐项落进公式:含义 / 公式 / 方向 / 输入列
|
||||
assert "含义:股息率" in f
|
||||
assert "公式:dv_ratio(Tushare daily_basic,逐日时点值)" in f
|
||||
assert "越高越好" in f and "higher_is_better" in f
|
||||
assert "输入列:dv_ratio" in f
|
||||
assert "score_i = Σ_f w_f × d_f × z_f,i" in f
|
||||
assert "z_f,i = (x_f,i - mean_i(x_f)) / std_i(x_f)" in f
|
||||
|
||||
def test_formula_has_two_level_truncation_and_cycles(self) -> None:
|
||||
f = describe_strategy(_spec()).formula
|
||||
assert "n = 30" in f and "x = 20" in f
|
||||
assert "每 6 个月" in f
|
||||
assert "择股日" in f and "调仓日" in f
|
||||
|
||||
def test_formula_has_costs_and_execution_price_basis(self) -> None:
|
||||
f = describe_strategy(_spec()).formula
|
||||
assert "佣金 = max(投入资金 × 0.03%, 最低佣金 5 元/笔)" in f
|
||||
assert "印花税 0.05%(仅卖出)" in f, "印花税只对卖出计提(local_engine._rebalance)"
|
||||
assert "滑点 0.1%" in f
|
||||
assert "后复权" in f and "乘以 adjust_factor" not in f # 口径文案由 _ADJUST_TEXT 给出
|
||||
assert "成交时点:调仓日收盘" in f
|
||||
assert "1,000,000 元" in f
|
||||
|
||||
def test_formula_renders_conditions_readably(self) -> None:
|
||||
spec = _spec(
|
||||
conditions=[
|
||||
ConditionSpec(field="dv_ratio", op="lte", value=30),
|
||||
ConditionSpec(field="close", op="gte", ref="ma60"),
|
||||
ConditionSpec(field="fundamental.roe", op="gte", value=15),
|
||||
ConditionSpec(field="static.industry", op="in", value=["白酒", "银行"]),
|
||||
ConditionSpec(field="static.industry", op="ne", value="白酒"),
|
||||
ConditionSpec(field="pe", op="lt", value=30.5),
|
||||
]
|
||||
)
|
||||
f = describe_strategy(spec).formula
|
||||
assert "dv_ratio <= 30" in f
|
||||
assert "close >= ma60" in f, "ref 条件必须渲染成「字段 op 字段」(同一股票同一日比较)"
|
||||
assert "fundamental.roe >= 15" in f
|
||||
assert 'static.industry ∈ ["白酒", "银行"]' in f
|
||||
assert 'static.industry != "白酒"' in f
|
||||
assert "pe < 30.5" in f
|
||||
# 说明字段域语义,避免使用者误以为 fundamental.* 是当日值
|
||||
assert "announce_date <= 择股日" in f
|
||||
|
||||
def test_empty_conditions_does_not_crash(self) -> None:
|
||||
doc = describe_strategy(_spec(conditions=[]))
|
||||
assert "(无)" in doc.formula
|
||||
assert any("未配置过滤条件" in w for w in doc.warnings)
|
||||
assert doc.steps
|
||||
|
||||
def test_unknown_factor_goes_to_warnings_not_crash(self) -> None:
|
||||
doc = describe_strategy(_spec(factors=[FactorSpec(name="no_such_factor", weight=2.0)]))
|
||||
assert any("未知因子 no_such_factor" in w for w in doc.warnings)
|
||||
assert "元数据缺失" in doc.formula
|
||||
assert doc.summary # 仍然可用(说明不因未知因子而失败)
|
||||
|
||||
def test_factor_meta_override_is_used(self) -> None:
|
||||
meta = {
|
||||
"my_signal": FactorDef(
|
||||
name="my_signal", description="自定义信号", formula="close / ma20",
|
||||
brief="实验因子", direction="lower_is_better", requires=("close",),
|
||||
)
|
||||
}
|
||||
doc = describe_strategy(
|
||||
_spec(factors=[FactorSpec(name="my_signal", weight=0.5)], conditions=[]),
|
||||
factor_meta=meta,
|
||||
)
|
||||
assert "自定义信号" in doc.formula and "公式:close / ma20" in doc.formula
|
||||
assert "越低越好" in doc.formula
|
||||
assert not any("未知因子" in w for w in doc.warnings)
|
||||
assert "从低到高排序" in doc.summary
|
||||
|
||||
def test_unknown_condition_field_goes_to_warnings(self) -> None:
|
||||
doc = describe_strategy(
|
||||
_spec(conditions=[ConditionSpec(field="no_such_col", op="gt", value=1)])
|
||||
)
|
||||
assert any("条件字段 no_such_col 未识别" in w for w in doc.warnings)
|
||||
assert "no_such_col > 1" in doc.formula
|
||||
|
||||
|
||||
class TestConditionSemanticsMatchSelectionEvaluator:
|
||||
"""渲染文本必须与 selection 的真实求值一致(以代码为准,不以注释/直觉为准)。"""
|
||||
|
||||
def test_gte_ref_and_lte_value_match_eligible_symbols(self) -> None:
|
||||
syms = ["600000.SH", "600001.SH", "600002.SH"]
|
||||
daily = synthetic_daily({"600000.SH": 0.006, "600001.SH": -0.006, "600002.SH": 0.0005}, n=300)
|
||||
obs = pd.Timestamp("2024-12-31")
|
||||
conds = [
|
||||
ConditionSpec(field="close", op="gte", ref="ma60"),
|
||||
ConditionSpec(field="momentum_20", op="lte", value=0.05),
|
||||
]
|
||||
fields = build_condition_fields(daily, conds, obs)
|
||||
passed = eligible_symbols(sorted(syms), conds, {}, fields, {})
|
||||
|
||||
# 期望值直接按「渲染出的表达式」的语义计算,与求值器结果对齐
|
||||
for sym in syms:
|
||||
expected = bool(
|
||||
fields["close"][sym] >= fields["ma60"][sym]
|
||||
and fields["momentum_20"][sym] <= 0.05
|
||||
)
|
||||
assert (sym in passed) == expected, f"{sym} 的渲染语义与求值器不一致"
|
||||
|
||||
# 且文档里出现的就是这两条表达式
|
||||
doc = describe_strategy(_spec(conditions=conds))
|
||||
assert "close >= ma60" in doc.formula
|
||||
assert "momentum_20 <= 0.05" in doc.formula
|
||||
|
||||
def test_in_not_in_and_missing_value_semantics(self) -> None:
|
||||
"""in/not_in 的右操作数是列表,语义为 `left in right`;缺失值语义见下。"""
|
||||
conds = [
|
||||
ConditionSpec(field="static.industry", op="in", value=["白酒"]),
|
||||
ConditionSpec(field="static.pe", op="ne", value=10), # 缺失字段
|
||||
]
|
||||
statics = {"A": {"industry": "白酒"}, "B": {"industry": "银行"}}
|
||||
passed = eligible_symbols(["A", "B"], conds, statics, {}, {})
|
||||
# 代码事实(selection._compare):op == "ne" 在 None 判定**之前**返回 left != right,
|
||||
# 因此字段缺失时 `ne` 判定为「不等于」→ 通过;其余运算符在缺失值上不通过。
|
||||
assert set(passed) == {"A"}
|
||||
doc = describe_strategy(_spec(conditions=conds))
|
||||
assert 'static.industry ∈ ["白酒"]' in doc.formula
|
||||
assert "static.pe != 10" in doc.formula
|
||||
assert "除 != 外一律判为「未通过」" in doc.formula
|
||||
|
||||
def test_string_comparison_is_lexicographic_in_real_evaluator(self) -> None:
|
||||
"""对照 selection._compare 的实际实现,修正其内联注释的说法。
|
||||
|
||||
`_compare` 的注释写「字符串会 ValueError → False」,但代码里 gt/gte/lt/lte
|
||||
最终走 `_num_cmp(left, right, op)`,而 `_num_cmp` **不做 float 转换**,
|
||||
因此两个字符串是按 Python 原生(字典序)比较的,并不会 ValueError。
|
||||
说明书按代码事实渲染为 `static.industry > "M"`(同一运算符、同一语义)。
|
||||
"""
|
||||
conds = [ConditionSpec(field="static.industry", op="gt", value="M")]
|
||||
statics = {"A": {"industry": "白酒"}} # '白酒' > 'M'(Unicode 码位更大)
|
||||
passed = eligible_symbols(["A"], conds, statics, {}, {})
|
||||
assert set(passed) == {"A"}
|
||||
assert 'static.industry > "M"' in describe_strategy(_spec(conditions=conds)).formula
|
||||
|
||||
|
||||
class TestStepsAndWarnings:
|
||||
def test_steps_follow_local_engine_behaviour(self) -> None:
|
||||
doc = describe_strategy(_spec())
|
||||
joined = " | ".join(doc.steps)
|
||||
assert "300 个自然日" in joined, "预热窗口来自 service._load_daily"
|
||||
assert "候选池" in joined and "selection_history" in joined
|
||||
assert "收盘" in joined, "成交发生在调仓日收盘"
|
||||
assert "涨停" in joined and "停牌" in joined
|
||||
assert "顺延" in joined, "defer_buy=True 必须如实说明顺延买入"
|
||||
assert "次日起计收益" in joined, "调仓日收盘生效(无未来函数)"
|
||||
assert "不强制平仓" in joined, "期末不平仓是引擎的真实行为,必须披露"
|
||||
|
||||
def test_steps_describe_substitute_mode(self) -> None:
|
||||
doc = describe_strategy(
|
||||
_spec(selection=SelectionSpec(top_n=5, allow_substitute=True), conditions=[])
|
||||
)
|
||||
assert any("替补" in s for s in doc.steps)
|
||||
|
||||
def test_universe_step_conditional_and_time_accurate(self) -> None:
|
||||
"""股票池步骤只写实际生效的过滤,且时点口径必须正确(起始日过滤一次)。
|
||||
|
||||
以 service._load_daily / universe.filter_stocks 的代码为准:universe 过滤在
|
||||
as_of=回测起始日执行一次,之后只有 exclude_st 会按当日名称逐择股日重判。
|
||||
"""
|
||||
doc = describe_strategy(_spec(universe=UniverseSpec(min_listing_days=250), conditions=[]))
|
||||
pool_step = doc.steps[1]
|
||||
assert "回测起始日按 universe 过滤一次" in pool_step
|
||||
assert "当日名称" in pool_step
|
||||
assert "指数成分" not in pool_step, "未配置 index_code 时不得声称有指数成分过滤"
|
||||
|
||||
idx = describe_strategy(
|
||||
_spec(
|
||||
universe=UniverseSpec(index_code="000300.SH", exclude_st=False, min_listing_days=0),
|
||||
conditions=[],
|
||||
)
|
||||
)
|
||||
assert "000300.SH 的当日历史成分" in idx.steps[1]
|
||||
assert "名称含 ST" not in idx.steps[1]
|
||||
|
||||
def test_warnings_cover_unmodelled_constraints(self) -> None:
|
||||
# exclude_st 默认 True(UniverseSpec 默认值)时才需要标注时点/快照口径
|
||||
doc = describe_strategy(_spec(universe=UniverseSpec(min_listing_days=250)))
|
||||
joined = " | ".join(doc.warnings)
|
||||
assert "exclude_suspended" in joined, "停牌未建模必须标注"
|
||||
assert "名称变更历史" in joined, "exclude_st 时点/快照口径必须标注"
|
||||
assert "幸存者偏差" in joined
|
||||
assert "涨跌停" in joined
|
||||
|
||||
def test_warnings_flag_none_adjustment_and_industry_cap(self) -> None:
|
||||
from app.domain.entities.research import PortfolioSpec
|
||||
|
||||
doc = describe_strategy(
|
||||
_spec(price_adjustment="none", portfolio=PortfolioSpec(max_industry_weight_pct=0.3))
|
||||
)
|
||||
joined = " | ".join(doc.warnings)
|
||||
assert "不复权" in joined
|
||||
assert "未建模" in joined and "行业权重" in joined
|
||||
|
||||
def test_stale_pool_warning_when_y_lt_m(self) -> None:
|
||||
doc = describe_strategy(
|
||||
_spec(selection_interval_months=12, rebalance_interval_months=3, conditions=[])
|
||||
)
|
||||
assert any("池子陈旧" in w for w in doc.warnings)
|
||||
|
||||
|
||||
class TestDefinitionInput:
|
||||
def test_strategy_definition_input_uses_placeholder_period(self) -> None:
|
||||
st = StrategyDefinition(
|
||||
name="演示",
|
||||
factors=[FactorSpec(name="momentum_60", weight=1.0)],
|
||||
selection=SelectionSpec(top_n=10),
|
||||
)
|
||||
doc = describe_strategy(st)
|
||||
assert doc.summary and doc.formula and doc.steps
|
||||
assert "momentum_60" in doc.formula
|
||||
assert "1900" not in doc.formula, "占位区间不得泄漏到展示文本"
|
||||
assert any("无回测区间" in w for w in doc.warnings)
|
||||
|
||||
def test_real_spec_has_no_placeholder_warning(self) -> None:
|
||||
doc = describe_strategy(_spec())
|
||||
assert not any("无回测区间" in w for w in doc.warnings)
|
||||
assert "2020-01-01 ~ 2024-12-31" in doc.formula
|
||||
|
||||
def test_describe_does_not_mutate_input(self) -> None:
|
||||
"""纯函数约定:只生成文本,不改写 spec(更不改写研究结果)。"""
|
||||
spec = _spec()
|
||||
before = spec.model_dump()
|
||||
describe_strategy(spec)
|
||||
assert spec.model_dump() == before
|
||||
|
||||
st = StrategyDefinition(name="演示", factors=[FactorSpec(name="momentum_60")])
|
||||
before_st = st.model_dump()
|
||||
describe_strategy(st)
|
||||
assert st.model_dump() == before_st
|
||||
|
||||
def test_unsupported_input_raises_clear_type_error(self) -> None:
|
||||
with pytest.raises(TypeError, match="ResearchSpec 或 StrategyDefinition"):
|
||||
describe_strategy({"factors": []}) # type: ignore[arg-type]
|
||||
|
||||
|
||||
# ---------- 策略 API:说明接口 + 原地更新(PUT) ----------
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def client(tmp_path) -> TestClient:
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'stratdoc.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
||||
|
||||
def _session_override():
|
||||
with Session() as s:
|
||||
yield s
|
||||
|
||||
app.dependency_overrides[deps.get_session] = _session_override
|
||||
with TestClient(app) as c:
|
||||
yield c
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
_SPEC_BODY = {
|
||||
"type": "backtest",
|
||||
"universe": {"exclude_st": False, "min_listing_days": 0},
|
||||
"factors": [{"name": "dividend_yield", "weight": 1.0}],
|
||||
"selection": {"top_n": 20, "hold_top_x": 20, "allow_substitute": False, "defer_buy": True},
|
||||
"rebalance": "monthly",
|
||||
"selection_interval_months": 6,
|
||||
"rebalance_interval_months": 6,
|
||||
"period": ["2020-01-01", "2024-12-31"],
|
||||
}
|
||||
|
||||
|
||||
class TestStrategyDocApi:
|
||||
def test_post_describe_from_research_spec_without_saving(self, client: TestClient) -> None:
|
||||
resp = client.post("/api/strategies/describe", json=_SPEC_BODY)
|
||||
assert resp.status_code == 200
|
||||
doc = resp.json()
|
||||
assert set(doc) == {"summary", "formula", "steps", "warnings"}
|
||||
assert "dividend_yield" in doc["formula"]
|
||||
assert "每 6 个月" in doc["summary"]
|
||||
# 未保存:列表仍为空(POST /describe 与 POST /strategies 不冲突)
|
||||
assert client.get("/api/strategies").json() == []
|
||||
|
||||
def test_post_describe_invalid_spec_422(self, client: TestClient) -> None:
|
||||
bad = dict(_SPEC_BODY, period=["2024-01-01", "2020-01-01"])
|
||||
assert client.post("/api/strategies/describe", json=bad).status_code == 422
|
||||
|
||||
def test_get_describe_saved_strategy_and_404(self, client: TestClient) -> None:
|
||||
created = client.post(
|
||||
"/api/strategies",
|
||||
json={"name": "高股息", "factors": [{"name": "dividend_yield", "weight": 1}]},
|
||||
)
|
||||
sid = created.json()["id"]
|
||||
resp = client.get(f"/api/strategies/{sid}/describe")
|
||||
assert resp.status_code == 200
|
||||
doc = resp.json()
|
||||
assert doc["summary"]
|
||||
assert any("无回测区间" in w for w in doc["warnings"])
|
||||
assert client.get("/api/strategies/STG-NOT-EXIST/describe").status_code == 404
|
||||
|
||||
def test_post_fills_empty_description(self, client: TestClient) -> None:
|
||||
"""需求:策略必须有说明 —— description 为空/纯空白时由 summary 自动补全。"""
|
||||
for desc in ("", " "):
|
||||
resp = client.post(
|
||||
"/api/strategies",
|
||||
json={
|
||||
"name": f"无说明策略-{len(desc)}",
|
||||
"description": desc,
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["description"].strip()
|
||||
assert "momentum_60" in body["description"]
|
||||
# 响应中的说明与 GET 读回一致(真的落库了)
|
||||
assert client.get(f"/api/strategies/{body['id']}").json()["description"] == body[
|
||||
"description"
|
||||
]
|
||||
|
||||
def test_post_keeps_explicit_description(self, client: TestClient) -> None:
|
||||
resp = client.post(
|
||||
"/api/strategies",
|
||||
json={
|
||||
"name": "有说明策略",
|
||||
"description": "我自己写的说明",
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
},
|
||||
)
|
||||
assert resp.json()["description"] == "我自己写的说明"
|
||||
|
||||
def test_auto_description_fits_column_width(self, client: TestClient) -> None:
|
||||
"""自动说明必须落在 `StrategyModel.description = String(300)` 之内。
|
||||
|
||||
超长在 SQLite(测试库)不会报错、到 MySQL 严格模式会 Data too long,
|
||||
因此这里显式断言列宽;截断必须带省略号(显式标记,不静默改短)。
|
||||
"""
|
||||
from app.api.strategies import _DESCRIPTION_MAX_CHARS
|
||||
from app.quant.factors import list_factors
|
||||
|
||||
all_factors = [{"name": f.name, "weight": 1} for f in list_factors()]
|
||||
resp = client.post(
|
||||
"/api/strategies", json={"name": "全因子策略", "factors": all_factors}
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
desc = resp.json()["description"]
|
||||
assert len(desc) <= _DESCRIPTION_MAX_CHARS
|
||||
assert desc.endswith("…"), "超长被截断时必须显式带省略号"
|
||||
|
||||
# 常规(单因子)说明远短于列宽:不应被截断
|
||||
normal = client.post(
|
||||
"/api/strategies",
|
||||
json={"name": "单因子策略", "factors": [{"name": "dividend_yield", "weight": 1}]},
|
||||
).json()["description"]
|
||||
assert len(normal) <= _DESCRIPTION_MAX_CHARS
|
||||
assert not normal.endswith("…")
|
||||
|
||||
|
||||
class TestStrategyUpdateApi:
|
||||
def _create(self, client: TestClient, name: str = "策略A") -> dict:
|
||||
resp = client.post(
|
||||
"/api/strategies",
|
||||
json={
|
||||
"name": name,
|
||||
"description": "初始说明",
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
"selection": {"top_n": 10},
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
return resp.json()
|
||||
|
||||
def test_put_updates_in_place_and_keeps_created_at(self, client: TestClient) -> None:
|
||||
created = self._create(client)
|
||||
sid = created["id"]
|
||||
created_at = created["created_at"]
|
||||
assert created_at is not None
|
||||
|
||||
resp = client.put(
|
||||
f"/api/strategies/{sid}",
|
||||
json={
|
||||
"name": "策略A",
|
||||
"description": "改后的说明",
|
||||
"factors": [{"name": "momentum_20", "weight": 2}],
|
||||
"selection": {"top_n": 5},
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["id"] == sid, "原地更新必须保持 id 不变(不新建)"
|
||||
assert body["created_at"] == created_at, "PUT 不得刷新创建时间"
|
||||
assert body["selection"]["top_n"] == 5
|
||||
assert body["factors"][0]["name"] == "momentum_20"
|
||||
assert body["description"] == "改后的说明"
|
||||
|
||||
# 再读一次:落库后的创建时间同样未变,且列表仍只有一条
|
||||
assert client.get(f"/api/strategies/{sid}").json()["created_at"] == created_at
|
||||
assert len(client.get("/api/strategies").json()) == 1
|
||||
|
||||
def test_put_path_id_wins_over_body_id(self, client: TestClient) -> None:
|
||||
created = self._create(client, "策略B")
|
||||
sid = created["id"]
|
||||
resp = client.put(
|
||||
f"/api/strategies/{sid}",
|
||||
json={
|
||||
"id": "STG-FAKE",
|
||||
"name": "策略B",
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["id"] == sid
|
||||
assert client.get("/api/strategies/STG-FAKE").status_code == 404
|
||||
|
||||
def test_put_missing_strategy_404(self, client: TestClient) -> None:
|
||||
resp = client.put(
|
||||
"/api/strategies/STG-NOT-EXIST",
|
||||
json={"name": "X", "factors": [{"name": "momentum_60", "weight": 1}]},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
assert "不存在" in resp.json()["detail"]
|
||||
|
||||
def test_put_duplicate_name_400(self, client: TestClient) -> None:
|
||||
first = self._create(client, "策略C")
|
||||
second = self._create(client, "策略D")
|
||||
resp = client.put(
|
||||
f"/api/strategies/{second['id']}",
|
||||
json={"name": "策略C", "factors": [{"name": "momentum_60", "weight": 1}]},
|
||||
)
|
||||
assert resp.status_code == 400
|
||||
assert "策略名已存在" in resp.json()["detail"]
|
||||
# 撞车失败后第二条策略名不变(未被部分写入)
|
||||
assert client.get(f"/api/strategies/{second['id']}").json()["name"] == "策略D"
|
||||
assert client.get(f"/api/strategies/{first['id']}").json()["name"] == "策略C"
|
||||
|
||||
def test_put_fills_empty_description(self, client: TestClient) -> None:
|
||||
created = self._create(client, "策略E")
|
||||
resp = client.put(
|
||||
f"/api/strategies/{created['id']}",
|
||||
json={
|
||||
"name": "策略E",
|
||||
"description": " ",
|
||||
"factors": [{"name": "dividend_yield", "weight": 1}],
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["description"].strip()
|
||||
assert "dividend_yield" in resp.json()["description"]
|
||||
@@ -0,0 +1,239 @@
|
||||
"""股票名称(`name`)填充与 `GET /api/stocks/names` 接口测试。
|
||||
|
||||
覆盖三件事(需求:前端任何出现代码的地方都要能显示名称):
|
||||
1) 回测结果的展示结构(symbol_curves / positions / trades / fills / signal_history /
|
||||
selection_history)由 `ResearchService.run_backtest` 统一回填名称;
|
||||
股票池为空时**静默跳过**(名称只是展示增强,不得让已算完的回测失败);
|
||||
2) 选股结果 `candidates[].name` 由 `SelectionService` 用已装配股票池回填;
|
||||
3) `GET /api/stocks/names` 返回 `{symbol: name}` **dict**(前端契约),
|
||||
且 `GET /api/stocks/{symbol}` 未被 `/names` 破坏(路由顺序)。
|
||||
|
||||
用内存 Repository + 合成行情(不连真库、不跑长区间),秒级完成。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from app.api import deps
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.domain.entities.market import Stock
|
||||
from app.domain.entities.research import (
|
||||
FactorSpec,
|
||||
ResearchSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
from app.main import app
|
||||
from app.quant.engine import LocalEngine
|
||||
from app.quant.service import ResearchService, _fill_names
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
|
||||
|
||||
_SYMS = ["600000.SH", "600001.SH", "600002.SH"]
|
||||
_NAMES = {"600000.SH": "浦发银行", "600001.SH": "邯郸钢铁", "600002.SH": "齐鲁石化"}
|
||||
|
||||
|
||||
def _stocks(symbols: list[str] | None = None) -> list[Stock]:
|
||||
return [
|
||||
Stock(symbol=s, name=_NAMES[s], list_date=date(1999, 11, 10))
|
||||
for s in (symbols or _SYMS)
|
||||
]
|
||||
|
||||
|
||||
class _MemStockRepo:
|
||||
def __init__(self, stocks: list[Stock]) -> None:
|
||||
self._stocks = stocks
|
||||
|
||||
def list(self) -> list[Stock]:
|
||||
return self._stocks
|
||||
|
||||
def get_by_symbol(self, symbol: str) -> Stock | None:
|
||||
return next((s for s in self._stocks if s.symbol == symbol), None)
|
||||
|
||||
|
||||
class _MemDailyRepo:
|
||||
def __init__(self, df: pd.DataFrame) -> None:
|
||||
self._bars = bars_dataframe_to_daily_bars(df)
|
||||
|
||||
def get_range(self, symbol, start, end):
|
||||
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
|
||||
|
||||
def get_range_many(self, symbols, start, end, adjust="none"):
|
||||
syms = set(symbols)
|
||||
return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end]
|
||||
|
||||
|
||||
def _daily_df() -> pd.DataFrame:
|
||||
return synthetic_daily(
|
||||
{"600000.SH": 0.004, "600001.SH": 0.001, "600002.SH": -0.003}, n=300
|
||||
)
|
||||
|
||||
|
||||
def _backtest_spec() -> ResearchSpec:
|
||||
return ResearchSpec(
|
||||
type="backtest",
|
||||
universe=UniverseSpec(exclude_st=False, min_listing_days=0),
|
||||
factors=[FactorSpec(name="momentum_20", weight=1.0)],
|
||||
selection=SelectionSpec(top_n=2),
|
||||
rebalance="monthly",
|
||||
# 半年区间(约 150 个交易日):够触发月度调仓与卖出,又足够快
|
||||
period=(date(2024, 6, 3), date(2024, 12, 31)),
|
||||
)
|
||||
|
||||
|
||||
class TestBacktestNameFill:
|
||||
def test_all_display_structures_carry_name(self) -> None:
|
||||
service = ResearchService(
|
||||
_MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), LocalEngine()
|
||||
)
|
||||
result = service.run_backtest(_backtest_spec())
|
||||
|
||||
assert result.symbol_curves, "应持有过股票并输出个股曲线"
|
||||
assert result.positions, "应有调仓后仓位记录"
|
||||
assert result.selection_history, "应有择股记录"
|
||||
assert result.signal_history, "应有交易意图记录"
|
||||
|
||||
for curve in result.symbol_curves:
|
||||
assert curve.name == _NAMES[curve.symbol]
|
||||
for pos in result.positions:
|
||||
assert pos.name == _NAMES[pos.symbol]
|
||||
for pick in result.selection_history:
|
||||
assert pick.name == _NAMES[pick.symbol]
|
||||
# 成交/意图记录里存在 symbol="" 的池子不足提示(非股票),其 name 保持 None
|
||||
for action in result.signal_history + result.fills:
|
||||
if action.symbol:
|
||||
assert action.name == _NAMES[action.symbol]
|
||||
else:
|
||||
assert action.name is None
|
||||
|
||||
def test_trades_and_marks_carry_name(self) -> None:
|
||||
service = ResearchService(
|
||||
_MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), LocalEngine()
|
||||
)
|
||||
result = service.run_backtest(_backtest_spec())
|
||||
assert result.trades, "月度调仓 + 因子换手应产生已卖出往返"
|
||||
for trade in result.trades:
|
||||
assert trade.name == _NAMES[trade.symbol]
|
||||
# marks 与 signal_history 同源,也应带上名称(前端个股曲线标注用)
|
||||
for curve in result.symbol_curves:
|
||||
for mark in curve.marks:
|
||||
assert mark.name == _NAMES[mark.symbol]
|
||||
|
||||
def test_fill_names_skips_silently_when_stock_pool_empty(self) -> None:
|
||||
"""空股票池 → 静默跳过(名称缺失只影响展示,不得抛错)。"""
|
||||
service = ResearchService(
|
||||
_MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), LocalEngine()
|
||||
)
|
||||
result = service.run_backtest(_backtest_spec())
|
||||
before = result.symbol_curves[0].name
|
||||
|
||||
returned = _fill_names(result, [])
|
||||
assert returned is result
|
||||
assert result.symbol_curves[0].name == before # 未被清空、也未被改写
|
||||
|
||||
def test_fill_names_is_idempotent_and_keeps_name_without_stock_record(self) -> None:
|
||||
"""名称回填幂等;股票池里没有该代码时保持 None(不伪造名称)。"""
|
||||
service = ResearchService(
|
||||
_MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), LocalEngine()
|
||||
)
|
||||
result = service.run_backtest(_backtest_spec())
|
||||
snapshot = [(c.symbol, c.name) for c in result.symbol_curves]
|
||||
|
||||
_fill_names(result, _stocks())
|
||||
assert [(c.symbol, c.name) for c in result.symbol_curves] == snapshot
|
||||
|
||||
only_other = [Stock(symbol="601398.SH", name="工商银行", list_date=date(2006, 1, 1))]
|
||||
_fill_names(result, only_other)
|
||||
assert all(c.name == _NAMES[c.symbol] for c in result.symbol_curves)
|
||||
|
||||
|
||||
class TestSelectionCandidateName:
|
||||
def _service(self) -> SelectionService:
|
||||
return SelectionService(_MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()))
|
||||
|
||||
def test_score_mode_candidates_carry_name(self) -> None:
|
||||
result = self._service().select(
|
||||
SelectionQuery(
|
||||
method="score",
|
||||
factors=[{"name": "momentum_20", "weight": 1.0}],
|
||||
top_n=2,
|
||||
as_of=date(2024, 12, 31),
|
||||
universe=UniverseSpec(exclude_st=False, min_listing_days=0),
|
||||
)
|
||||
)
|
||||
assert result.candidates
|
||||
for cand in result.candidates:
|
||||
assert cand.name == _NAMES[cand.symbol]
|
||||
|
||||
def test_condition_mode_candidates_carry_name(self) -> None:
|
||||
result = self._service().select(
|
||||
SelectionQuery(
|
||||
method="condition",
|
||||
conditions=[{"field": "close", "op": "gte", "ref": "ma20"}],
|
||||
as_of=date(2024, 12, 31),
|
||||
universe=UniverseSpec(exclude_st=False, min_listing_days=0),
|
||||
)
|
||||
)
|
||||
assert result.candidates
|
||||
for cand in result.candidates:
|
||||
assert cand.name == _NAMES[cand.symbol]
|
||||
|
||||
|
||||
class TestQlibEngineNameFill:
|
||||
def test_qlib_engine_path_also_fills_names(self, tmp_path) -> None:
|
||||
"""名称回填写在服务层,因此 LocalEngine / QlibEngine 两条路径都必须生效。
|
||||
|
||||
这条用例走真实的 Qlib 数据管线(落盘 bin → D.features 读回),
|
||||
而不是只测本地引擎后「假设」另一条路径也覆盖到了。
|
||||
"""
|
||||
pytest.importorskip("qlib") # 未安装 qlib 的环境跳过,不假装覆盖
|
||||
from app.quant.qlib_adapter.engine import QlibEngine
|
||||
|
||||
service = ResearchService(
|
||||
_MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), QlibEngine(qlib_dir=tmp_path)
|
||||
)
|
||||
result = service.run_backtest(_backtest_spec())
|
||||
assert result.symbol_curves, "qlib 路径也应产生个股曲线"
|
||||
for curve in result.symbol_curves:
|
||||
assert curve.name == _NAMES[curve.symbol]
|
||||
for pos in result.positions:
|
||||
assert pos.name == _NAMES[pos.symbol]
|
||||
for pick in result.selection_history:
|
||||
assert pick.name == _NAMES[pick.symbol]
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def client() -> TestClient:
|
||||
"""只覆盖 stock 仓储(这三个接口只用 StockRepoDep,不需要 DB session)。"""
|
||||
stocks = [*_stocks(), Stock(symbol="600519.SH", name="贵州茅台", list_date=date(2001, 8, 27))]
|
||||
app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(stocks) # noqa: SLF001
|
||||
with TestClient(app) as c:
|
||||
yield c
|
||||
app.dependency_overrides.clear()
|
||||
|
||||
|
||||
class TestStockNamesApi:
|
||||
def test_names_returns_symbol_to_name_dict(self, client: TestClient) -> None:
|
||||
resp = client.get("/api/stocks/names")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
# 契约:dict(不是数组),前端按 Record<string, string> 消费
|
||||
assert isinstance(body, dict)
|
||||
assert body["600519.SH"] == "贵州茅台"
|
||||
assert body["600000.SH"] == _NAMES["600000.SH"]
|
||||
assert set(body) == {"600000.SH", "600001.SH", "600002.SH", "600519.SH"}
|
||||
|
||||
def test_symbol_path_not_swallowed_by_names_route(self, client: TestClient) -> None:
|
||||
"""路由顺序回归:/{symbol} 仍按代码查询(证明 /names 未被参数吞掉、也没吞掉它)。"""
|
||||
one = client.get("/api/stocks/600519.SH")
|
||||
assert one.status_code == 200
|
||||
assert one.json()["symbol"] == "600519.SH"
|
||||
assert one.json()["name"] == "贵州茅台"
|
||||
|
||||
assert client.get("/api/stocks/999999.SZ").status_code == 404
|
||||
assert client.get("/api/stocks?q=茅台").json()[0]["symbol"] == "600519.SH"
|
||||
@@ -20,10 +20,14 @@ class FakePro:
|
||||
self.payload = payload or []
|
||||
self.error = error
|
||||
self.calls: list[str] = []
|
||||
# 关键:必须记录关键字参数,否则「list_status 是否真的透传给 tushare」无法被断言
|
||||
# (退市股同步完全依赖该参数,见 TestDelistedStocks)
|
||||
self.call_kwargs: list[dict] = []
|
||||
|
||||
def __getattr__(self, api: str):
|
||||
def _run(**kwargs):
|
||||
self.calls.append(api)
|
||||
self.call_kwargs.append(kwargs)
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
return self.payload
|
||||
@@ -257,3 +261,68 @@ class TestIndexWeight:
|
||||
rows = p.get_index_weight("000300.SH")
|
||||
assert pro.calls == ["index_weight"]
|
||||
assert len(rows) == 1 and rows[0].symbol == "600519.SH"
|
||||
|
||||
|
||||
class TestDelistedStocks:
|
||||
"""退市股(幸存者偏差修正):list_status 透传、NaN 归一、代码规范过滤。
|
||||
|
||||
实测背景:tushare `stock_basic` 不带 list_status 时**只返回在市股票**,
|
||||
退市股整体缺失(230 只 2019-12 后退市)→ 回测系统性高估收益;
|
||||
而退市股记录的 industry/area 是 NaN、status 为空,且含 T 前缀异常代码。
|
||||
"""
|
||||
|
||||
def test_list_status_passed_to_api(self) -> None:
|
||||
pro = FakePro(payload=[])
|
||||
p = TushareProvider(token="t", pro=pro)
|
||||
p.get_stock_basic("D")
|
||||
assert pro.calls == ["stock_basic"]
|
||||
# tushare 不带 list_status 时只返回在市股票 → 退市股同步必须真的传 "D"
|
||||
assert pro.call_kwargs[0]["list_status"] == "D"
|
||||
|
||||
def test_default_list_status_is_listed_only(self) -> None:
|
||||
pro = FakePro(payload=[])
|
||||
p = TushareProvider(token="t", pro=pro)
|
||||
p.get_stock_basic()
|
||||
assert pro.calls == ["stock_basic"] # 默认 "L",保持既有行为
|
||||
assert pro.call_kwargs[0]["list_status"] == "L"
|
||||
|
||||
def test_nan_optional_fields_become_none(self) -> None:
|
||||
stocks = TushareProvider.normalize_stock(
|
||||
[
|
||||
{
|
||||
"ts_code": "000005.SZ",
|
||||
"name": "ST星源(退)",
|
||||
"area": float("nan"),
|
||||
"industry": float("nan"),
|
||||
"market": float("nan"),
|
||||
"exchange": "SZSE",
|
||||
"list_date": "19901210",
|
||||
"delist_date": "20240426",
|
||||
}
|
||||
]
|
||||
)
|
||||
s = stocks[0]
|
||||
assert s.industry is None and s.area is None and s.market is None
|
||||
assert s.exchange == "SZSE"
|
||||
assert s.delist_date == date(2024, 4, 26)
|
||||
|
||||
def test_missing_status_falls_back_to_query_status(self) -> None:
|
||||
"""退市表 status 为空:必须按查询的 list_status 兜底,不能一律标成 L。"""
|
||||
rec = [{"ts_code": "000005.SZ", "name": "ST星源(退)", "list_date": "19901210",
|
||||
"delist_date": "20240426", "status": None}]
|
||||
assert TushareProvider.normalize_stock(rec)[0].status == "L" # 既有默认
|
||||
assert TushareProvider.normalize_stock(rec, default_status="D")[0].status == "D"
|
||||
|
||||
def test_abnormal_code_skipped_not_fatal(self) -> None:
|
||||
"""'T600018.SH'(上港集箱(退),2006 退市)不得让整批退市列表拉取失败。"""
|
||||
pro = FakePro(
|
||||
payload=[
|
||||
{"ts_code": "000005.SZ", "name": "ST星源(退)", "list_date": "19901210",
|
||||
"delist_date": "20240426"},
|
||||
{"ts_code": "T600018.SH", "name": "上港集箱(退)", "list_date": "19960101",
|
||||
"delist_date": "20061020"},
|
||||
]
|
||||
)
|
||||
p = TushareProvider(token="t", pro=pro)
|
||||
stocks = p.get_stock_basic("D")
|
||||
assert [s.symbol for s in stocks] == ["000005.SZ"]
|
||||
|
||||
Reference in New Issue
Block a user