feat(backend): 策略库重构为「选股策略 + 公共配置 + 回测组合」三件套
按用户目标把原来「一个策略 = 全套参数」拆开(已确认的设计决策):
- 公共配置 GlobalConfig(全局唯一):佣金/印花税/滑点/最低佣金/复权口径/基准
- 选股策略 SelectionStrategy(原 StrategyDefinition 改名):只剩股票池+因子+条件,
不再持有 selection/rebalance/costs/portfolio/区间/资金
- 回测组合 BacktestCombo:引用若干选股策略 + 回测时才定的参数
(起始资金、持仓数 N、持仓天数区间 [Tmin,Tmax]、调仓时机 日/周/月、区间)
引擎(app/quant/combo_engine.py,新增):
- 多策略打分 = 并集 + Borda 秩和(各策略 1/名次 求和;不假设不同策略分值可比,
能容纳各策略股票池不同);抽出纯函数 borda_combine 便于单测
- 持仓天数区间 [Tmin,Tmax]:Tmax **每个交易日**强制了结(安全阀,月频下也不超期);
Tmin 仅在调仓日保护(掉出 TopN 但未满 Tmin 暂留,防频繁换手);调仓日为增量调仓
(只卖超期/掉队且满 Tmin 的,从 TopN 补买至 N 只,不主动减持以尊重 Tmin)
- 调仓时机 daily/weekly/monthly(local_engine.rebalance_dates 新增日频分支)
- 产出与旧 runner 同构的 BacktestResult,前端可视化无需改动;config_snapshot 固化
ComboRunSpec(组合+当时各策略定义+当时成本/复权)保证可复现
数据层:
- 新表 global_config(默认行:万三/hfq/最低佣金5元)、backtest_combo
- 迁移 b4c5d6e7f8a9:建两表 + 把存量 strategy.config_json 的回测参数键剥掉、
spec_type 收敛为 selection(已在真实 MariaDB 验证:STG-16BFBF08 清洗后只剩
universe/factors/conditions)
- 仓储 SqlAlchemyGlobalConfigRepository / SqlAlchemyComboRepository + Protocol
API:
- /api/config GET/PUT;/api/combos CRUD + /{id}/run + /run(kind=combo 异步 Job)
- job_executor 新增 combo 分支:取齐策略+读公共配置→ComboService.run,归档 kind
记 backtest(结果结构相同)
- /api/strategies 切到 SelectionStrategy,移除已废弃的 /{id}/expand
- strategy_doc.describe_strategy 支持 SelectionStrategy(只讲「怎么选」,如实声明
资金/持仓/调仓/成本/区间在回测组合里定)
旧的 ResearchSpec + /api/backtests 保留(因子测试与既有契约自检仍用),
作为底层 escape hatch;用户产品路径改为回测组合。
测试:新增 test_combo_engine(6)/test_combo_service(3)/test_combo_api(5),
改写 test_strategies/test_strategy_doc 适配新模型。全量 403 passed(原 388)。
This commit is contained in:
@@ -23,7 +23,7 @@ from app.domain.entities.research import (
|
||||
)
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
from app.domain.entities.signal import SignalRules
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
|
||||
SqlAlchemyCompositeRepository,
|
||||
)
|
||||
@@ -295,15 +295,15 @@ def build_tools(factories: dict | None = None) -> list[Tool]:
|
||||
if not factors:
|
||||
return "请提供至少一个 factors(逗号分隔)"
|
||||
description = str(_pick(args, "description", "") or "")
|
||||
st = StrategyDefinition(
|
||||
# 选股策略只存「选股条件组合」:股票池 + 因子(+ 可选条件)。
|
||||
# top_n / rebalance 等回测执行参数已移到「回测组合」,Agent 不再在此指定。
|
||||
st = SelectionStrategy(
|
||||
name=name,
|
||||
description=description,
|
||||
universe=UniverseSpec(
|
||||
exclude_st=bool(_pick(args, "exclude_st", True)), min_listing_days=0
|
||||
),
|
||||
factors=factors,
|
||||
selection={"top_n": int(_pick(args, "top_n", 10) or 10)},
|
||||
rebalance=str(_pick(args, "rebalance", "monthly")),
|
||||
)
|
||||
from app.application.services.job_executor import new_id
|
||||
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
"""回测组合 API:/api/combos(CRUD + 运行)。
|
||||
|
||||
POST /api/combos 保存组合(name 唯一)
|
||||
GET /api/combos 列表
|
||||
GET /api/combos/{id} 详情
|
||||
PUT /api/combos/{id} 原地更新
|
||||
DELETE /api/combos/{id} 删除
|
||||
POST /api/combos/{id}/run 提交已保存组合为异步 Job(kind=combo)
|
||||
POST /api/combos/run 提交临时组合(不保存)为异步 Job
|
||||
|
||||
运行时:按 combo.strategy_ids 取齐选股策略 + 读公共配置 → ComboService.run → 归档。
|
||||
费率/复权来自公共配置,快照进归档 config_snapshot(可复现,AGENT.md §21)。
|
||||
|
||||
依赖注入说明:建 Job 与策略校验都走 FastAPI 注入的 session/repo(而非直接 SessionLocal),
|
||||
这样测试用 dependency_overrides 替换数据库时也能命中同一份库,行为一致可测。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, HTTPException
|
||||
|
||||
from app.api.deps import ComboRepoDep, DbSession, StrategyRepoDep
|
||||
from app.application.services.job_executor import new_id, run_job_background
|
||||
from app.domain.entities.combo import BacktestCombo
|
||||
from app.domain.entities.research import JobRecord, JobStatus
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/combos", tags=["combos"])
|
||||
|
||||
|
||||
def _ensure_strategies_exist(combo: BacktestCombo, strategy_repo) -> None:
|
||||
"""引用的选股策略必须都存在;缺任何一个即 400(提前失败,不等后台 Job 才暴露)。"""
|
||||
for sid in combo.strategy_ids:
|
||||
if strategy_repo.get(sid) is None:
|
||||
raise HTTPException(status_code=400, detail=f"组合引用的选股策略 {sid} 不存在")
|
||||
|
||||
|
||||
@router.post("", response_model=BacktestCombo, summary="保存回测组合")
|
||||
def create_combo(
|
||||
combo: BacktestCombo,
|
||||
repo: ComboRepoDep,
|
||||
session: DbSession,
|
||||
) -> BacktestCombo:
|
||||
try:
|
||||
saved = repo.save(combo.model_copy(update={"id": new_id("CMB")}))
|
||||
session.commit()
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
return repo.get(saved.id) or saved
|
||||
|
||||
|
||||
@router.get("", response_model=list[BacktestCombo], summary="回测组合列表")
|
||||
def list_combos(repo: ComboRepoDep) -> list[BacktestCombo]:
|
||||
return repo.list()
|
||||
|
||||
|
||||
@router.get("/{combo_id}", response_model=BacktestCombo, summary="读取回测组合")
|
||||
def get_combo(combo_id: str, repo: ComboRepoDep) -> BacktestCombo:
|
||||
row = repo.get(combo_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"回测组合 {combo_id} 不存在")
|
||||
return row
|
||||
|
||||
|
||||
@router.put("/{combo_id}", response_model=BacktestCombo, summary="原地更新回测组合")
|
||||
def update_combo(
|
||||
combo_id: str,
|
||||
combo: BacktestCombo,
|
||||
repo: ComboRepoDep,
|
||||
session: DbSession,
|
||||
) -> BacktestCombo:
|
||||
existing = repo.get(combo_id)
|
||||
if existing is None:
|
||||
raise HTTPException(status_code=404, detail=f"回测组合 {combo_id} 不存在")
|
||||
payload = combo.model_copy(update={"id": combo_id, "created_at": existing.created_at})
|
||||
try:
|
||||
saved = repo.save(payload)
|
||||
session.commit()
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
return repo.get(saved.id) or saved
|
||||
|
||||
|
||||
@router.delete("/{combo_id}", summary="删除回测组合")
|
||||
def delete_combo(combo_id: str, repo: ComboRepoDep, session: DbSession) -> dict:
|
||||
if not repo.delete(combo_id):
|
||||
raise HTTPException(status_code=404, detail=f"回测组合 {combo_id} 不存在")
|
||||
session.commit()
|
||||
return {"deleted": combo_id}
|
||||
|
||||
|
||||
def _submit_combo_job(
|
||||
combo: BacktestCombo,
|
||||
*,
|
||||
strategy_repo,
|
||||
session,
|
||||
background: BackgroundTasks,
|
||||
) -> dict:
|
||||
"""校验策略 → 建 Job(kind=combo)→ 入后台执行。
|
||||
|
||||
spec_json 存 BacktestCombo JSON;执行端(job_executor)识别 kind="combo",
|
||||
再按 strategy_ids 取策略 + 读公共配置后调 ComboService。Job 表只存组合本身,
|
||||
策略/配置的「当时快照」由 ComboService 写进归档 config_snapshot(可复现)。
|
||||
校验与建 Job 都用注入的 session/repo,保证与测试覆写一致。
|
||||
"""
|
||||
_ensure_strategies_exist(combo, strategy_repo)
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"), kind="combo", status=JobStatus.QUEUED,
|
||||
spec_json=combo.model_dump_json(),
|
||||
)
|
||||
SqlAlchemyJobRepository(session).create(job)
|
||||
session.commit()
|
||||
background.add_task(run_job_background, job.id)
|
||||
return {"job_id": job.id, "status": job.status}
|
||||
|
||||
|
||||
@router.post("/{combo_id}/run", summary="运行已保存的回测组合(异步 Job)")
|
||||
def run_saved_combo(
|
||||
combo_id: str,
|
||||
repo: ComboRepoDep,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
background: BackgroundTasks,
|
||||
) -> dict:
|
||||
combo = repo.get(combo_id)
|
||||
if combo is None:
|
||||
raise HTTPException(status_code=404, detail=f"回测组合 {combo_id} 不存在")
|
||||
return _submit_combo_job(combo, strategy_repo=strategy_repo, session=session, background=background)
|
||||
|
||||
|
||||
@router.post("/run", summary="运行临时回测组合(不保存,异步 Job)")
|
||||
def run_adhoc_combo(
|
||||
combo: BacktestCombo,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
background: BackgroundTasks,
|
||||
) -> dict:
|
||||
return _submit_combo_job(combo, strategy_repo=strategy_repo, session=session, background=background)
|
||||
@@ -0,0 +1,30 @@
|
||||
"""公共配置 API:/api/config(全局唯一一份费率/滑点/复权口径)。
|
||||
|
||||
GET /api/config 读取(未配置过返回带默认值的实例)
|
||||
PUT /api/config 更新(upsert 单例)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.deps import DbSession, GlobalConfigRepoDep
|
||||
from app.domain.entities.combo import GlobalConfig
|
||||
|
||||
router = APIRouter(prefix="/config", tags=["config"])
|
||||
|
||||
|
||||
@router.get("", response_model=GlobalConfig, summary="读取公共配置")
|
||||
def get_config(repo: GlobalConfigRepoDep) -> GlobalConfig:
|
||||
return repo.get()
|
||||
|
||||
|
||||
@router.put("", response_model=GlobalConfig, summary="更新公共配置")
|
||||
def update_config(
|
||||
config: GlobalConfig,
|
||||
repo: GlobalConfigRepoDep,
|
||||
session: DbSession,
|
||||
) -> GlobalConfig:
|
||||
saved = repo.save(config)
|
||||
session.commit()
|
||||
return saved
|
||||
@@ -14,6 +14,7 @@ from app.application.services.chart_service import ChartService
|
||||
from app.application.services.replay_service import ReplayService
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.application.services.signal_service import SignalService
|
||||
from app.domain.repositories.combo import ComboRepository, GlobalConfigRepository
|
||||
from app.domain.repositories.composite import CompositeRepository
|
||||
from app.domain.repositories.factor import FactorRepository
|
||||
from app.domain.repositories.index import IndexConstituentRepository
|
||||
@@ -29,6 +30,10 @@ from app.domain.repositories.market import (
|
||||
from app.domain.repositories.selection import SelectionRepository
|
||||
from app.domain.repositories.signal import SignalRepository
|
||||
from app.domain.repositories.strategy import StrategyRepository
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.combo_impl import (
|
||||
SqlAlchemyComboRepository,
|
||||
SqlAlchemyGlobalConfigRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
|
||||
SqlAlchemyCompositeRepository,
|
||||
)
|
||||
@@ -181,6 +186,14 @@ def _strategy_repo_factory(session: DbSession) -> StrategyRepository:
|
||||
return SqlAlchemyStrategyRepository(session)
|
||||
|
||||
|
||||
def _global_config_repo_factory(session: DbSession) -> GlobalConfigRepository:
|
||||
return SqlAlchemyGlobalConfigRepository(session)
|
||||
|
||||
|
||||
def _combo_repo_factory(session: DbSession) -> ComboRepository:
|
||||
return SqlAlchemyComboRepository(session)
|
||||
|
||||
|
||||
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
|
||||
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
|
||||
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
|
||||
@@ -194,6 +207,8 @@ SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)]
|
||||
ReplayServiceDep = Annotated[ReplayService, Depends(_replay_service_factory)]
|
||||
ChartServiceDep = Annotated[ChartService, Depends(_chart_service_factory)]
|
||||
StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)]
|
||||
GlobalConfigRepoDep = Annotated[GlobalConfigRepository, Depends(_global_config_repo_factory)]
|
||||
ComboRepoDep = Annotated[ComboRepository, Depends(_combo_repo_factory)]
|
||||
|
||||
|
||||
def _job_repo_factory(session: DbSession):
|
||||
|
||||
@@ -11,7 +11,9 @@ from fastapi import APIRouter
|
||||
from app.api import (
|
||||
agent,
|
||||
charts,
|
||||
combos,
|
||||
composites,
|
||||
config,
|
||||
experiments,
|
||||
factors,
|
||||
health,
|
||||
@@ -35,6 +37,8 @@ api_router.include_router(charts.router)
|
||||
api_router.include_router(replays.router)
|
||||
api_router.include_router(signals.router)
|
||||
api_router.include_router(strategies.router)
|
||||
api_router.include_router(config.router)
|
||||
api_router.include_router(combos.router)
|
||||
api_router.include_router(jobs.router)
|
||||
api_router.include_router(experiments.router)
|
||||
api_router.include_router(agent.router)
|
||||
|
||||
@@ -1,4 +1,7 @@
|
||||
"""策略 API(M8.3):/api/strategies CRUD + 展开为 ResearchSpec + 说明/公式生成。
|
||||
"""选股策略 API:/api/strategies CRUD + 说明/公式生成。
|
||||
|
||||
2026-09 重构:策略库只存「选股条件组合」(股票池+因子+条件),不再持有回测执行参数;
|
||||
回测改由「回测组合」(/api/combos)驱动,故旧的 /{id}/expand(→ResearchSpec)已移除。
|
||||
|
||||
POST /api/strategies 保存策略(name 唯一;description 为空时自动补全)
|
||||
POST /api/strategies/describe body: ResearchSpec → StrategyDoc(未保存的策略也能预览)
|
||||
@@ -6,7 +9,6 @@ GET /api/strategies 列表
|
||||
GET /api/strategies/{id}
|
||||
PUT /api/strategies/{id} 原地更新(不新建、不刷新 created_at)
|
||||
DELETE /api/strategies/{id}
|
||||
POST /api/strategies/{id}/expand body: {period:[start,end], initial_capital?} → ResearchSpec
|
||||
GET /api/strategies/{id}/describe → StrategyDoc
|
||||
|
||||
路由顺序注意:`/describe` 这类**字面量路径**一律声明在 `/{strategy_id}` 之前 ——
|
||||
@@ -15,24 +17,17 @@ GET /api/strategies/{id}/describe → StrategyDoc
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.api.deps import DbSession, StrategyRepoDep
|
||||
from app.application.services.job_executor import new_id
|
||||
from app.domain.entities.research import ResearchSpec
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.quant.strategy_doc import StrategyDoc, describe_strategy
|
||||
|
||||
router = APIRouter(prefix="/strategies", tags=["strategies"])
|
||||
|
||||
|
||||
class ExpandRequest(BaseModel):
|
||||
period: tuple[date, date]
|
||||
initial_capital: float = Field(default=1_000_000.0, gt=0)
|
||||
|
||||
|
||||
# strategy.description 列宽(StrategyModel.description = String(300))。
|
||||
# 自动补全的说明必须落在列宽内,否则 MySQL 严格模式会直接报 Data too long(SQLite 不拦,
|
||||
@@ -41,7 +36,7 @@ class ExpandRequest(BaseModel):
|
||||
_DESCRIPTION_MAX_CHARS = 300
|
||||
|
||||
|
||||
def _ensure_description(definition: StrategyDefinition) -> StrategyDefinition:
|
||||
def _ensure_description(definition: SelectionStrategy) -> SelectionStrategy:
|
||||
"""说明为空/纯空白时,用 `describe_strategy(...).summary` 补全(需求:策略必须有说明)。
|
||||
|
||||
说明由 spec **真实推导**(AGENT.md §24:不许编造),只在空值时补、不覆盖显式说明。
|
||||
@@ -56,12 +51,12 @@ def _ensure_description(definition: StrategyDefinition) -> StrategyDefinition:
|
||||
return definition.model_copy(update={"description": summary})
|
||||
|
||||
|
||||
@router.post("", response_model=StrategyDefinition, summary="保存策略")
|
||||
@router.post("", response_model=SelectionStrategy, summary="保存选股策略")
|
||||
def create_strategy(
|
||||
definition: StrategyDefinition,
|
||||
definition: SelectionStrategy,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
) -> StrategyDefinition:
|
||||
) -> SelectionStrategy:
|
||||
try:
|
||||
saved = strategy_repo.save(
|
||||
_ensure_description(definition).model_copy(update={"id": new_id("STG")})
|
||||
@@ -84,13 +79,13 @@ def describe_research_spec(spec: ResearchSpec) -> StrategyDoc:
|
||||
return describe_strategy(spec)
|
||||
|
||||
|
||||
@router.get("", response_model=list[StrategyDefinition], summary="策略列表")
|
||||
def list_strategies(strategy_repo: StrategyRepoDep) -> list[StrategyDefinition]:
|
||||
@router.get("", response_model=list[SelectionStrategy], summary="选股策略列表")
|
||||
def list_strategies(strategy_repo: StrategyRepoDep) -> list[SelectionStrategy]:
|
||||
return strategy_repo.list()
|
||||
|
||||
|
||||
@router.get("/{strategy_id}", response_model=StrategyDefinition, summary="读取策略")
|
||||
def get_strategy(strategy_id: str, strategy_repo: StrategyRepoDep) -> StrategyDefinition:
|
||||
@router.get("/{strategy_id}", response_model=SelectionStrategy, summary="读取选股策略")
|
||||
def get_strategy(strategy_id: str, strategy_repo: StrategyRepoDep) -> SelectionStrategy:
|
||||
row = strategy_repo.get(strategy_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
@@ -104,7 +99,6 @@ def describe_saved_strategy(
|
||||
"""已保存策略的说明/公式(404 语义与 GET /{strategy_id} 一致)。
|
||||
|
||||
策略定义不含回测区间,说明里的区间为占位文本(`warnings` 中已如实标注),
|
||||
展开回测后(/expand)再用该 ResearchSpec 调 `POST /describe` 即为确定区间的版本。
|
||||
"""
|
||||
row = strategy_repo.get(strategy_id)
|
||||
if row is None:
|
||||
@@ -112,13 +106,13 @@ def describe_saved_strategy(
|
||||
return describe_strategy(row)
|
||||
|
||||
|
||||
@router.put("/{strategy_id}", response_model=StrategyDefinition, summary="原地更新策略")
|
||||
@router.put("/{strategy_id}", response_model=SelectionStrategy, summary="原地更新选股策略")
|
||||
def update_strategy(
|
||||
strategy_id: str,
|
||||
definition: StrategyDefinition,
|
||||
definition: SelectionStrategy,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
session: DbSession,
|
||||
) -> StrategyDefinition:
|
||||
) -> SelectionStrategy:
|
||||
"""原地更新(策略库「编辑」用):id 以**路径**为准,created_at 沿用库中已有值。
|
||||
|
||||
为什么必须显式带上 created_at:仓储 `save()` 只在 created_at 为空时才写 now()
|
||||
@@ -152,16 +146,3 @@ def delete_strategy(
|
||||
session.commit()
|
||||
return {"deleted": strategy_id}
|
||||
|
||||
|
||||
@router.post("/{strategy_id}/expand", response_model=ResearchSpec, summary="展开为研究 Spec")
|
||||
def expand_strategy(
|
||||
strategy_id: str,
|
||||
req: ExpandRequest,
|
||||
strategy_repo: StrategyRepoDep,
|
||||
) -> ResearchSpec:
|
||||
row = strategy_repo.get(strategy_id)
|
||||
if row is None:
|
||||
raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在")
|
||||
if req.period[0] >= req.period[1]:
|
||||
raise HTTPException(status_code=400, detail="period 必须满足 start < end")
|
||||
return row.to_research_spec(period=req.period, initial_capital=req.initial_capital)
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
"""回测组合服务:把「组合 + 选股策略 + 公共配置」解析并执行成 BacktestResult。
|
||||
|
||||
职责(应用层用例,AGENT.md §16/§17):
|
||||
- 装配行情数据(复用 ResearchService 的 load_daily_df / universe 过滤 / 名称回填);
|
||||
- 为每个选股策略构造「as_of → 合格股票集」闭包(复用 selection 求值器,保证与
|
||||
`/api/selections` 同口径,v2 §25);
|
||||
- 调 combo_engine.run_combo_backtest(多策略 Borda + 持仓区间 + 日/周/月);
|
||||
- 把可复现的 ComboRunSpec 写进结果 config_snapshot(已在引擎内完成)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
from typing import Any
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.combo import (
|
||||
BacktestCombo,
|
||||
ComboRunSpec,
|
||||
GlobalConfig,
|
||||
SelectionStrategyRef,
|
||||
)
|
||||
from app.domain.entities.research import BacktestResult, UniverseSpec
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.quant.combo_engine import run_combo_backtest
|
||||
from app.quant.selection import build_condition_fields, eligible_symbols
|
||||
from app.quant.service import _fill_names, load_daily_df, split_factor_columns
|
||||
from app.quant.universe import filter_stocks, names_as_of, resolve_members
|
||||
|
||||
|
||||
class ComboService:
|
||||
"""回测组合用例入口。依赖注入各 Repository + 引擎无关的数据装配函数。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stock_repo,
|
||||
daily_repo,
|
||||
*,
|
||||
index_repo=None,
|
||||
basic_repo=None,
|
||||
financial_repo=None,
|
||||
name_repo=None,
|
||||
) -> None:
|
||||
self._stock_repo = stock_repo
|
||||
self._daily_repo = daily_repo
|
||||
self._index_repo = index_repo
|
||||
self._basic_repo = basic_repo
|
||||
self._financial_repo = financial_repo
|
||||
self._name_repo = name_repo
|
||||
self._last_stocks: list = []
|
||||
|
||||
def run(
|
||||
self,
|
||||
combo: BacktestCombo,
|
||||
strategies: list[SelectionStrategy],
|
||||
config: GlobalConfig,
|
||||
on_stage=None,
|
||||
) -> BacktestResult:
|
||||
if not strategies:
|
||||
raise ValueError("回测组合至少需要引用一个选股策略")
|
||||
# 校验引用的策略 id 与传入一致(防御性:调用方应已按 combo.strategy_ids 取齐)
|
||||
given = {s.id for s in strategies}
|
||||
missing = [sid for sid in combo.strategy_ids if sid not in given]
|
||||
if missing:
|
||||
raise ValueError(f"组合引用的选股策略未提供:{missing}")
|
||||
|
||||
_stage(on_stage, "data_loading")
|
||||
daily = self._load_daily(combo, strategies, config)
|
||||
_stage(on_stage, "backtesting")
|
||||
eligibility_fns = [self._build_eligibility(s, daily) for s in strategies]
|
||||
refs = [_to_ref(s) for s in strategies]
|
||||
result = run_combo_backtest(
|
||||
combo=combo,
|
||||
strategies=refs,
|
||||
costs=config.to_cost_spec(),
|
||||
price_adjustment=config.price_adjustment,
|
||||
daily=daily,
|
||||
eligibility_fns=eligibility_fns,
|
||||
)
|
||||
_stage(on_stage, "analysis")
|
||||
return _fill_names(result, self._last_stocks)
|
||||
|
||||
# ---- 数据装配(与 ResearchService 同口径,复用底层函数) ----
|
||||
|
||||
def _merged_universe(self, strategies: list[SelectionStrategy]) -> UniverseSpec:
|
||||
"""合并各策略的股票池口径用于「装哪些股票的行情」。
|
||||
|
||||
取并集语义:symbols 白名单取并集;exclude_st / min_listing_days 取**最宽松**
|
||||
(任一策略不剔 ST 则不剔,min_listing_days 取最小)—— 因为最终选股由各策略
|
||||
自己的 eligibility 闭包再过滤,这里只为「行情装配覆盖足够多的股票」。
|
||||
index_code 不一致时无法合并 → 报错(同一组合里混用不同指数成分没有明确语义)。
|
||||
"""
|
||||
indices = {s.universe.index_code for s in strategies if s.universe.index_code}
|
||||
if len(indices) > 1:
|
||||
raise ValueError(
|
||||
f"组合内各选股策略的指数成分不一致({sorted(indices)}),无法合并股票池;"
|
||||
"请统一指数或改用 symbols 白名单"
|
||||
)
|
||||
symbols: set[str] = set()
|
||||
for s in strategies:
|
||||
symbols.update(s.universe.symbols)
|
||||
return UniverseSpec(
|
||||
market=strategies[0].universe.market,
|
||||
exclude_st=all(s.universe.exclude_st for s in strategies),
|
||||
exclude_suspended=all(s.universe.exclude_suspended for s in strategies),
|
||||
min_listing_days=min(s.universe.min_listing_days for s in strategies),
|
||||
index_code=indices.pop() if indices else None,
|
||||
symbols=sorted(symbols),
|
||||
)
|
||||
|
||||
def _load_daily(
|
||||
self, combo: BacktestCombo, strategies: list[SelectionStrategy], config: GlobalConfig
|
||||
) -> pd.DataFrame:
|
||||
start, end = combo.period
|
||||
data_start = start - timedelta(days=300) # 因子 warmup 余量
|
||||
all_stocks = self._stock_repo.list()
|
||||
merged = self._merged_universe(strategies)
|
||||
name_at, _applied = names_as_of(all_stocks, start, self._name_repo)
|
||||
stocks = filter_stocks(
|
||||
all_stocks, merged, as_of=start,
|
||||
members=resolve_members(self._index_repo, merged, start),
|
||||
name_at=name_at,
|
||||
)
|
||||
self._last_stocks = stocks
|
||||
# 所需列 = 所有策略因子 + 所有策略条件引用列 + close
|
||||
needed = {"close"}
|
||||
for s in strategies:
|
||||
from app.domain.entities.research import ResearchSpec
|
||||
from app.quant.engine import factor_required_columns
|
||||
|
||||
# 借用既有列裁剪逻辑:构造一个临时 spec 只为算 required_columns
|
||||
tmp = ResearchSpec(
|
||||
type="backtest", universe=s.universe, factors=s.factors,
|
||||
conditions=s.conditions, period=combo.period,
|
||||
)
|
||||
needed |= factor_required_columns(tmp)
|
||||
bar_cols, basic_cols = split_factor_columns(needed)
|
||||
symbols = [st.symbol for st in stocks]
|
||||
daily = load_daily_df(
|
||||
self._daily_repo, symbols, data_start, end, sorted(bar_cols),
|
||||
adjust="none", price_adjust=config.price_adjustment,
|
||||
)
|
||||
if basic_cols:
|
||||
daily = self._attach_basic(daily, symbols, data_start, end, sorted(basic_cols))
|
||||
return daily
|
||||
|
||||
def _attach_basic(self, daily, symbols, start, end, columns) -> pd.DataFrame:
|
||||
from app.quant.service import load_basic_df, merge_basic_into_daily
|
||||
|
||||
if self._basic_repo is None:
|
||||
raise ValueError(
|
||||
f"选股策略条件/因子需要每日指标列 {columns}(daily_basic),但未注入 DailyBasicRepository"
|
||||
)
|
||||
basic = load_basic_df(self._basic_repo, symbols, start, end, columns)
|
||||
if basic.empty:
|
||||
raise ValueError(
|
||||
f"daily_basic 在 {start}~{end} 无数据,无法计算需要 {columns} 的因子/条件"
|
||||
)
|
||||
return merge_basic_into_daily(daily, basic)
|
||||
|
||||
def _build_eligibility(self, strategy: SelectionStrategy, daily: pd.DataFrame):
|
||||
"""单策略的「as_of → 合格股票集」闭包(与 ResearchService._build_eligibility 同口径)。"""
|
||||
if not self._last_stocks:
|
||||
if not strategy.conditions and not strategy.universe.exclude_st:
|
||||
return None
|
||||
raise ValueError("universe 过滤结果为空,无法构造选股条件求值器")
|
||||
statics = {s.symbol: s.model_dump() for s in self._last_stocks}
|
||||
candidates = sorted(statics)
|
||||
st_fn = self._build_st_filter(strategy, candidates)
|
||||
if not strategy.conditions:
|
||||
if st_fn is None:
|
||||
return None
|
||||
allowed: dict[date, set[str]] = {}
|
||||
|
||||
def _st_only(as_of: date) -> set[str]:
|
||||
if as_of not in allowed:
|
||||
allowed[as_of] = set(candidates) - st_fn(as_of)
|
||||
return allowed[as_of]
|
||||
|
||||
return _st_only
|
||||
|
||||
uses_fundamental = any(
|
||||
f.startswith("fundamental.")
|
||||
for c in strategy.conditions
|
||||
for f in (c.field, c.ref or "")
|
||||
)
|
||||
cache: dict[date, set[str]] = {}
|
||||
|
||||
def _fn(as_of: date) -> set[str]:
|
||||
if as_of in cache:
|
||||
return cache[as_of]
|
||||
financial = self._load_financial(candidates, as_of) if uses_fundamental else {}
|
||||
fields = build_condition_fields(daily, strategy.conditions, pd.Timestamp(as_of))
|
||||
if not fields:
|
||||
cache[as_of] = set()
|
||||
return cache[as_of]
|
||||
passed = set(eligible_symbols(candidates, strategy.conditions, statics, fields, financial))
|
||||
if st_fn is not None:
|
||||
passed -= st_fn(as_of)
|
||||
cache[as_of] = passed
|
||||
return cache[as_of]
|
||||
|
||||
return _fn
|
||||
|
||||
def _build_st_filter(self, strategy: SelectionStrategy, candidates: list[str]):
|
||||
if not strategy.universe.exclude_st or self._name_repo is None:
|
||||
return None
|
||||
cache: dict[date, set[str]] = {}
|
||||
|
||||
def _fn(as_of: date) -> set[str]:
|
||||
if as_of not in cache:
|
||||
name_at, applied = names_as_of(self._last_stocks, as_of, self._name_repo)
|
||||
if not applied[0]:
|
||||
cache[as_of] = set()
|
||||
else:
|
||||
st_syms: set[str] = set()
|
||||
for st in self._last_stocks:
|
||||
nm = (name_at or {}).get(st.symbol) or st.name
|
||||
if nm and "ST" in nm.upper():
|
||||
st_syms.add(st.symbol)
|
||||
cache[as_of] = st_syms
|
||||
return cache[as_of]
|
||||
|
||||
return _fn
|
||||
|
||||
def _load_financial(self, symbols: list[str], as_of: date) -> dict[str, Any]:
|
||||
if self._financial_repo is None:
|
||||
raise ValueError("条件引用了 fundamental.* 字段,但未注入 FinancialRepository")
|
||||
getter = getattr(self._financial_repo, "list_announced_many", None)
|
||||
rows = list(getter(symbols, as_of)) if getter else []
|
||||
out: dict[str, Any] = {}
|
||||
for r in rows:
|
||||
out[r.symbol] = r
|
||||
return out
|
||||
|
||||
|
||||
def _to_ref(s: SelectionStrategy) -> SelectionStrategyRef:
|
||||
return SelectionStrategyRef(
|
||||
id=s.id,
|
||||
name=s.name,
|
||||
universe=s.universe.model_dump(),
|
||||
factors=[f.model_dump() for f in s.factors],
|
||||
conditions=[c.model_dump() for c in s.conditions],
|
||||
)
|
||||
|
||||
|
||||
def _stage(cb, name: str) -> None:
|
||||
if cb is not None:
|
||||
cb(name)
|
||||
|
||||
|
||||
# 让 ComboRunSpec 在模块导入时完成前向引用重建(entities/combo.py 末尾已 rebuild,此处兜底)
|
||||
ComboRunSpec.model_rebuild()
|
||||
@@ -76,10 +76,42 @@ def _execute_inner(
|
||||
session.commit()
|
||||
try:
|
||||
is_selection = job.kind == "selection"
|
||||
is_combo = job.kind == "combo"
|
||||
basic_repo = basic_repo_factory(session) if basic_repo_factory else None
|
||||
index_repo = index_repo_factory(session) if index_repo_factory else None
|
||||
name_repo = name_repo_factory(session) if name_repo_factory else None
|
||||
if is_selection:
|
||||
if is_combo:
|
||||
# 回测组合:解析 combo + 取齐选股策略 + 读公共配置 → ComboService.run
|
||||
from app.application.services.combo_service import ComboService
|
||||
from app.domain.entities.combo import BacktestCombo
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.combo_impl import (
|
||||
SqlAlchemyGlobalConfigRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
|
||||
SqlAlchemyStrategyRepository,
|
||||
)
|
||||
|
||||
combo = BacktestCombo.model_validate_json(job.spec_json)
|
||||
strategy_repo = SqlAlchemyStrategyRepository(session)
|
||||
strategies = []
|
||||
for sid in combo.strategy_ids:
|
||||
st = strategy_repo.get(sid)
|
||||
if st is None:
|
||||
raise ValueError(f"组合引用的选股策略 {sid} 不存在(可能已被删除)")
|
||||
strategies.append(st)
|
||||
config = SqlAlchemyGlobalConfigRepository(session).get()
|
||||
service = ComboService(
|
||||
stock_repo_factory(session),
|
||||
daily_repo_factory(session),
|
||||
index_repo=index_repo,
|
||||
basic_repo=basic_repo,
|
||||
financial_repo=(
|
||||
financial_repo_factory(session) if financial_repo_factory else None
|
||||
),
|
||||
name_repo=name_repo,
|
||||
)
|
||||
spec = combo # 仅用于下方分支判断占位;实际执行用 combo
|
||||
elif is_selection:
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
|
||||
@@ -118,7 +150,9 @@ def _execute_inner(
|
||||
except Exception: # noqa: BLE001 —— 阶段上报失败不阻断执行
|
||||
pass
|
||||
|
||||
if is_selection:
|
||||
if is_combo:
|
||||
result = service.run(spec, strategies, config, on_stage=_set_stage)
|
||||
elif is_selection:
|
||||
_set_stage("selection")
|
||||
result = service.select(spec)
|
||||
elif spec.type == "backtest":
|
||||
@@ -126,9 +160,12 @@ def _execute_inner(
|
||||
else:
|
||||
result = service.run_factor_test(spec, on_stage=_set_stage)
|
||||
|
||||
# combo 的结果是 BacktestResult,归档 kind 记为 "backtest" 以便前端按回测渲染;
|
||||
# spec_json 仍存原始 combo(含 strategy_ids),可复现快照在 result.config_snapshot。
|
||||
archive_kind = "backtest" if is_combo else job.kind
|
||||
experiment = archive_experiment(
|
||||
session=session,
|
||||
kind=job.kind,
|
||||
kind=archive_kind,
|
||||
spec_json=job.spec_json,
|
||||
result=result,
|
||||
job_id=job.id,
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
"""回测组合与公共配置领域实体(2026-09 重构)。
|
||||
|
||||
把原来「一个策略 = 全套参数」拆成三件独立的事(用户目标):
|
||||
|
||||
1. **GlobalConfig(公共配置,全局唯一)** —— 费率 / 印花税 / 滑点 / 最低佣金 /
|
||||
复权口径 / 基准。所有回测共用,不再塞进每个策略。
|
||||
2. **SelectionStrategy(选股策略,见 strategy.py)** —— 只剩「选股条件组合」:
|
||||
股票池 + 因子 + 过滤条件。**不含**资金 / 持仓数 / 持仓时间 / 调仓 / 费率 / 区间。
|
||||
3. **BacktestCombo(回测组合)** —— 引用若干选股策略 + 回测时才定的参数:
|
||||
起始资金、持仓数量 N、持仓天数区间 [Tmin, Tmax]、调仓时机(日/周/月)、回测区间。
|
||||
|
||||
执行时由服务层把「组合 + 被引用的选股策略 + 公共配置快照」解析成一个
|
||||
`ComboRunSpec`,喂给组合引擎;该 spec 会原样写进归档的 config_snapshot,
|
||||
保证事后可复现(AGENT.md §21),即使之后公共配置被改也不影响历史结果。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
from app.domain.entities.research import CostSpec
|
||||
|
||||
# ---------- 公共配置(全局唯一) ----------
|
||||
|
||||
|
||||
class GlobalConfig(BaseModel):
|
||||
"""全局交易成本与行情口径(单例,id 恒为 "default")。
|
||||
|
||||
为什么把复权口径也放这里:一次回测只能有一个复权口径(同一份行情不能既前复权
|
||||
又后复权),而多个选股策略可能想混用 —— 与其让它们在组合里打架,不如统一为
|
||||
全局口径,高股息默认 hfq。若将来确需按组合区分,再加字段即可(向前兼容)。
|
||||
"""
|
||||
|
||||
id: str = Field(default="default", description="单例主键,恒为 default")
|
||||
commission_rate: float = Field(default=0.0003, ge=0, le=0.01, description="佣金率(如 0.0003 = 万三)")
|
||||
stamp_tax_rate: float = Field(default=0.0005, ge=0, le=0.01, description="印花税率(仅卖出)")
|
||||
slippage_rate: float = Field(default=0.001, ge=0, le=0.05, description="滑点率")
|
||||
min_commission: float = Field(default=5.0, ge=0, le=100.0, description="单笔最低佣金(元)")
|
||||
price_adjustment: str = Field(
|
||||
default="hfq", pattern="^(none|qfq|hfq)$",
|
||||
description="行情复权口径:none / qfq / hfq(高股息类建议 hfq)",
|
||||
)
|
||||
benchmark: str = Field(default="000300.SH", description="对照基准指数代码")
|
||||
updated_at: datetime | None = None
|
||||
|
||||
def to_cost_spec(self) -> CostSpec:
|
||||
"""转成引擎用的 CostSpec(benchmark 一并带入)。"""
|
||||
return CostSpec(
|
||||
commission_rate=self.commission_rate,
|
||||
stamp_tax_rate=self.stamp_tax_rate,
|
||||
slippage_rate=self.slippage_rate,
|
||||
min_commission=self.min_commission,
|
||||
benchmark=self.benchmark,
|
||||
)
|
||||
|
||||
|
||||
# ---------- 回测组合 ----------
|
||||
|
||||
# 调仓时机:日 / 周 / 月(在原有 weekly/monthly 之上新增 daily)
|
||||
REBALANCE_FREQS = ("daily", "weekly", "monthly")
|
||||
|
||||
|
||||
class BacktestCombo(BaseModel):
|
||||
"""一个可保存、可复跑的回测组合。
|
||||
|
||||
持仓模型(用户确认的语义):
|
||||
- `hold_count` = N:目标持仓只数(等权)。
|
||||
- `hold_min_days` = Tmin:个股**最少**持有天数 —— 掉出 TopN 时若未满 Tmin 不卖
|
||||
(防止频繁换手);但超过 Tmax 仍强制卖(安全阀优先)。
|
||||
- `hold_max_days` = Tmax:个股**最多**持有天数 —— 超过即强制了结(None = 不限)。
|
||||
- `rebalance_freq`:多久重新打分排序并调仓一次(日/周/月)。
|
||||
⚠️ Tmax 强制卖出**每个交易日**都检查(不只调仓日),否则月频下会远超 Tmax。
|
||||
"""
|
||||
|
||||
id: str = ""
|
||||
name: str = Field(min_length=1, max_length=64)
|
||||
description: str = ""
|
||||
strategy_ids: list[str] = Field(
|
||||
min_length=1, description="引用的选股策略 id(≥1 个;多策略取并集后 Borda 秩和打分)"
|
||||
)
|
||||
initial_capital: float = Field(default=1_000_000.0, gt=0, description="起始资金(元)")
|
||||
hold_count: int = Field(ge=1, le=1000, description="目标持仓只数 N")
|
||||
hold_min_days: int = Field(default=0, ge=0, description="个股最少持有天数 Tmin")
|
||||
hold_max_days: int | None = Field(
|
||||
default=None, ge=1, description="个股最多持有天数 Tmax;None = 不强制了结"
|
||||
)
|
||||
rebalance_freq: str = Field(
|
||||
default="monthly", description="调仓时机:daily / weekly / monthly"
|
||||
)
|
||||
period: tuple[date, date]
|
||||
version: str = "1"
|
||||
created_at: datetime | None = None
|
||||
|
||||
@field_validator("rebalance_freq")
|
||||
@classmethod
|
||||
def _freq(cls, v: str) -> str:
|
||||
if v not in REBALANCE_FREQS:
|
||||
raise ValueError(f"rebalance_freq 必须是 {REBALANCE_FREQS} 之一,收到 {v!r}")
|
||||
return v
|
||||
|
||||
@field_validator("period")
|
||||
@classmethod
|
||||
def _period_ordered(cls, period: tuple[date, date]) -> tuple[date, date]:
|
||||
if period[0] >= period[1]:
|
||||
raise ValueError("period 必须满足 start < end")
|
||||
return period
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _hold_band_and_strategies(self) -> BacktestCombo:
|
||||
if (
|
||||
self.hold_max_days is not None
|
||||
and self.hold_min_days > 0
|
||||
and self.hold_max_days < self.hold_min_days
|
||||
):
|
||||
raise ValueError(
|
||||
f"持仓上限 Tmax={self.hold_max_days} 不能小于下限 Tmin={self.hold_min_days}"
|
||||
)
|
||||
if len(set(self.strategy_ids)) != len(self.strategy_ids):
|
||||
raise ValueError("strategy_ids 存在重复的策略 id")
|
||||
return self
|
||||
|
||||
|
||||
class ComboRunSpec(BaseModel):
|
||||
"""解析后的、可复现的组合运行规格(写入归档 config_snapshot)。
|
||||
|
||||
为什么不直接存 BacktestCombo:组合只引用 strategy_ids,且费率/复权来自公共配置;
|
||||
若事后策略被删改、公共配置被调整,光凭 combo 无法复现。这里把「当时用到的策略定义
|
||||
+ 当时的成本/复权快照」一起固化,归档即可独立复现(AGENT.md §21)。
|
||||
"""
|
||||
|
||||
combo: BacktestCombo
|
||||
strategies: list[SelectionStrategyRef] = Field(
|
||||
description="运行时刻各选股策略的快照(name/universe/factors/conditions)"
|
||||
)
|
||||
costs: CostSpec
|
||||
price_adjustment: str = Field(pattern="^(none|qfq|hfq)$")
|
||||
config_version: str = "1"
|
||||
|
||||
|
||||
class SelectionStrategyRef(BaseModel):
|
||||
"""ComboRunSpec 内嵌的策略快照(只取选股相关字段,避免把已废弃字段带进归档)。"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
universe: dict # UniverseSpec.model_dump()
|
||||
factors: list[dict] # [{name, weight}]
|
||||
conditions: list[dict] = Field(default_factory=list)
|
||||
|
||||
|
||||
ComboRunSpec.model_rebuild()
|
||||
@@ -1,66 +1,46 @@
|
||||
"""策略领域实体(M8.3,v2 §17/§5.3)。
|
||||
"""选股策略领域实体(2026-09 重构:原 StrategyDefinition → SelectionStrategy)。
|
||||
|
||||
Strategy = 完整策略定义(universe + factors + selection + rebalance + costs +
|
||||
portfolio,除回测区间 period 外),保存为命名资产;回测时补 period 展开为
|
||||
ResearchSpec(v2 §18 Research Specification 为统一契约,策略是其持久化形态)。
|
||||
策略库现在**只存选股条件组合**:股票池 + 因子 + 过滤条件。
|
||||
资金 / 持仓数 / 持仓时间 / 调仓时机 / 费率 / 回测区间一律移到「回测组合」
|
||||
(BacktestCombo)与「公共配置」(GlobalConfig),在回测时才确定。
|
||||
|
||||
仍映射到 `strategy` 表(id/name/description/spec_type/config_json/version/created_at),
|
||||
config_json 只存 universe/factors/conditions —— 迁移会把旧行的多余键剥掉。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from app.domain.entities.research import (
|
||||
ConditionSpec,
|
||||
CostSpec,
|
||||
FactorSpec,
|
||||
PortfolioSpec,
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.domain.entities.research import ConditionSpec, FactorSpec, UniverseSpec
|
||||
|
||||
|
||||
class StrategyDefinition(BaseModel):
|
||||
class SelectionStrategy(BaseModel):
|
||||
"""一个选股策略 = 选股条件组合(不含任何回测执行参数)。"""
|
||||
|
||||
id: str = ""
|
||||
name: str = Field(min_length=1, max_length=64)
|
||||
description: str = ""
|
||||
spec_type: str = Field(default="backtest", pattern="^(backtest|factor_test)$")
|
||||
spec_type: str = Field(default="selection", pattern="^(selection|backtest)$")
|
||||
universe: UniverseSpec = UniverseSpec()
|
||||
price_adjustment: str = Field(default="none", pattern="^(none|qfq|hfq)$")
|
||||
factors: list[FactorSpec] = Field(min_length=1)
|
||||
selection: SelectionSpec = SelectionSpec()
|
||||
rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$")
|
||||
factors: list[FactorSpec] = Field(min_length=1, description="打分因子(至少 1 个)")
|
||||
conditions: list[ConditionSpec] = Field(
|
||||
default_factory=list, description="选股过滤条件(与 ResearchSpec.conditions 一致)"
|
||||
default_factory=list, description="过滤条件(AND,universe 之后、因子排序之前)"
|
||||
)
|
||||
selection_interval_months: int | None = Field(
|
||||
default=None, ge=1, le=60, description="m:择股间隔(月)"
|
||||
)
|
||||
rebalance_interval_months: int | None = Field(
|
||||
default=None, ge=1, le=60, description="y:调仓间隔(月);缺省 = m"
|
||||
)
|
||||
costs: CostSpec = CostSpec()
|
||||
portfolio: PortfolioSpec = PortfolioSpec()
|
||||
version: str = "1"
|
||||
created_at: datetime | None = None
|
||||
|
||||
def to_research_spec(self, period: tuple[date, date], initial_capital: float | None = None):
|
||||
"""补全回测区间/资金后展开为标准 ResearchSpec(可在 Job/回测执行)。"""
|
||||
from app.domain.entities.research import ResearchSpec
|
||||
@model_validator(mode="after")
|
||||
def _no_duplicate_factors(self) -> "SelectionStrategy":
|
||||
names = [f.name for f in self.factors]
|
||||
if len(set(names)) != len(names):
|
||||
raise ValueError("factors 存在重复因子名")
|
||||
return self
|
||||
|
||||
return ResearchSpec(
|
||||
type=self.spec_type,
|
||||
universe=self.universe,
|
||||
price_adjustment=self.price_adjustment,
|
||||
factors=self.factors,
|
||||
conditions=self.conditions,
|
||||
selection=self.selection,
|
||||
rebalance=self.rebalance,
|
||||
selection_interval_months=self.selection_interval_months,
|
||||
rebalance_interval_months=self.rebalance_interval_months,
|
||||
period=period,
|
||||
costs=self.costs,
|
||||
portfolio=self.portfolio,
|
||||
initial_capital=initial_capital if initial_capital else 1_000_000.0,
|
||||
)
|
||||
|
||||
# 兼容别名:重构前到处引用的旧名字。新代码请用 SelectionStrategy;
|
||||
# 保留别名是为了让尚未迁移的导入点(agent 工具等)在过渡期不炸,
|
||||
# 最终会全部替换掉(见各调用点的 TODO)。
|
||||
StrategyDefinition = SelectionStrategy
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
"""回测组合 + 公共配置 Repository Protocol(2026-09 重构)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.combo import BacktestCombo, GlobalConfig
|
||||
|
||||
|
||||
class GlobalConfigRepository(Protocol):
|
||||
def get(self) -> GlobalConfig:
|
||||
"""读取全局配置;不存在则返回带默认值的实例(不写库)。"""
|
||||
|
||||
def save(self, config: GlobalConfig) -> GlobalConfig:
|
||||
"""upsert 单例(id 恒为 default)。"""
|
||||
|
||||
|
||||
class ComboRepository(Protocol):
|
||||
def save(self, combo: BacktestCombo) -> BacktestCombo:
|
||||
"""新建或更新(name 冲突抛 ValueError)。"""
|
||||
|
||||
def get(self, combo_id: str) -> BacktestCombo | None: ...
|
||||
|
||||
def list(self) -> list[BacktestCombo]: ...
|
||||
|
||||
def delete(self, combo_id: str) -> bool: ...
|
||||
@@ -1,20 +1,20 @@
|
||||
"""策略 Repository Protocol(M8.3)。"""
|
||||
"""选股策略 Repository Protocol。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
|
||||
|
||||
class StrategyRepository(Protocol):
|
||||
def save(self, definition: StrategyDefinition) -> StrategyDefinition:
|
||||
def save(self, definition: SelectionStrategy) -> SelectionStrategy:
|
||||
"""新建(name 冲突抛 ValueError)。"""
|
||||
|
||||
def get(self, strategy_id: str) -> StrategyDefinition | None: ...
|
||||
def get(self, strategy_id: str) -> SelectionStrategy | None: ...
|
||||
|
||||
def get_by_name(self, name: str) -> StrategyDefinition | None: ...
|
||||
def get_by_name(self, name: str) -> SelectionStrategy | None: ...
|
||||
|
||||
def list(self) -> list[StrategyDefinition]: ...
|
||||
def list(self) -> list[SelectionStrategy]: ...
|
||||
|
||||
def delete(self, strategy_id: str) -> bool: ...
|
||||
|
||||
+113
@@ -0,0 +1,113 @@
|
||||
"""公共配置 + 回测组合表,并把存量策略收敛为「选股条件组合」(2026-09 重构)
|
||||
|
||||
Revision ID: b4c5d6e7f8a9
|
||||
Revises: a3f8c21d9b47
|
||||
Create Date: 2026-09-30
|
||||
|
||||
背景:把原来「一个策略 = 全套参数」拆成三件事 ——
|
||||
1. global_config:费率/印花税/滑点/最低佣金/复权口径/基准(全局唯一一行);
|
||||
2. selection_strategy(复用 strategy 表):只剩股票池 + 因子 + 过滤条件;
|
||||
3. backtest_combo:引用若干选股策略 + 回测参数(资金/持仓数/持仓天数区间/调仓时机/区间)。
|
||||
|
||||
本迁移:
|
||||
- 新建 global_config(并插入默认行)与 backtest_combo 两张表;
|
||||
- 把 strategy.config_json 里**已废弃的回测执行参数键**剥掉(selection / rebalance /
|
||||
costs / portfolio / price_adjustment / *_interval_months),只留 universe/factors/conditions,
|
||||
并把 spec_type 标为 selection。旧数据不丢(归档里的 ResearchSpec 快照原样保留只读)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "b4c5d6e7f8a9"
|
||||
down_revision: str | None = "a3f8c21d9b47"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
# 选股策略不再承载的回测执行参数键(读出时也会被仓储丢弃,这里在存储侧也清掉)
|
||||
_LEGACY_KEYS = (
|
||||
"selection",
|
||||
"rebalance",
|
||||
"costs",
|
||||
"portfolio",
|
||||
"price_adjustment",
|
||||
"selection_interval_months",
|
||||
"rebalance_interval_months",
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ---- 1. 公共配置(单例) ----
|
||||
op.create_table(
|
||||
"global_config",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("commission_rate", sa.Numeric(10, 6), nullable=False, server_default="0.0003"),
|
||||
sa.Column("stamp_tax_rate", sa.Numeric(10, 6), nullable=False, server_default="0.0005"),
|
||||
sa.Column("slippage_rate", sa.Numeric(10, 6), nullable=False, server_default="0.001"),
|
||||
sa.Column("min_commission", sa.Numeric(10, 4), nullable=False, server_default="5"),
|
||||
sa.Column("price_adjustment", sa.String(length=8), nullable=False, server_default="hfq"),
|
||||
sa.Column("benchmark", sa.String(length=16), nullable=False, server_default="000300.SH"),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
# 插入默认行(高股息场景默认后复权 hfq、最低佣金 5 元)
|
||||
op.execute(
|
||||
"INSERT INTO global_config (id, commission_rate, stamp_tax_rate, slippage_rate, "
|
||||
"min_commission, price_adjustment, benchmark) VALUES "
|
||||
"('default', 0.0003, 0.0005, 0.001, 5, 'hfq', '000300.SH')"
|
||||
)
|
||||
|
||||
# ---- 2. 回测组合 ----
|
||||
op.create_table(
|
||||
"backtest_combo",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("name", sa.String(length=64), nullable=False),
|
||||
sa.Column("description", sa.String(length=300), nullable=False, server_default=""),
|
||||
sa.Column("strategy_ids_json", sa.Text(), nullable=False),
|
||||
sa.Column("initial_capital", sa.Numeric(20, 2), nullable=False, server_default="1000000"),
|
||||
sa.Column("hold_count", sa.Integer(), nullable=False, server_default="20"),
|
||||
sa.Column("hold_min_days", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("hold_max_days", sa.Integer(), nullable=True),
|
||||
sa.Column("rebalance_freq", sa.String(length=12), nullable=False, server_default="monthly"),
|
||||
sa.Column("start_date", sa.Date(), nullable=False),
|
||||
sa.Column("end_date", sa.Date(), nullable=False),
|
||||
sa.Column("version", sa.String(length=16), nullable=False, server_default="1"),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("name", name="uq_backtest_combo_name"),
|
||||
)
|
||||
|
||||
# ---- 3. 存量策略 → 选股条件组合(剥掉回测执行参数键) ----
|
||||
conn = op.get_bind()
|
||||
rows = conn.execute(sa.text("SELECT id, config_json FROM strategy")).fetchall()
|
||||
for row_id, cfg_text in rows:
|
||||
try:
|
||||
data = json.loads(cfg_text) if cfg_text else {}
|
||||
except json.JSONDecodeError:
|
||||
continue # 损坏行不动它(读出时仓储也会容错)
|
||||
changed = False
|
||||
for key in _LEGACY_KEYS:
|
||||
if key in data:
|
||||
data.pop(key)
|
||||
changed = True
|
||||
# spec_type 收敛为 selection(旧值多为 backtest)
|
||||
if data.get("spec_type") != "selection":
|
||||
data["spec_type"] = "selection"
|
||||
changed = True
|
||||
if changed:
|
||||
conn.execute(
|
||||
sa.text("UPDATE strategy SET config_json = :cfg, spec_type = 'selection' WHERE id = :id"),
|
||||
{"cfg": json.dumps(data, ensure_ascii=False), "id": row_id},
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# 回滚:删两张新表。strategy.config_json 被剥掉的键无法精确还原
|
||||
# (原始值未备份),故 downgrade 仅撤表结构,不承诺恢复旧策略的完整 config。
|
||||
op.drop_table("backtest_combo")
|
||||
op.drop_table("global_config")
|
||||
@@ -7,6 +7,10 @@
|
||||
from app.infrastructure.persistence.sqlalchemy.models.composite import ( # noqa: F401
|
||||
FactorCompositeModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.combo import ( # noqa: F401
|
||||
BacktestComboModel,
|
||||
GlobalConfigModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.factor import ( # noqa: F401
|
||||
FactorDefinitionModel,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
"""公共配置 + 回测组合表(2026-09 重构)。
|
||||
|
||||
- `global_config`:全局唯一一行(id="default"),存费率/滑点/最低佣金/复权口径/基准。
|
||||
- `backtest_combo`:回测组合,引用若干选股策略(strategy_ids JSON)+ 回测参数
|
||||
(资金/持仓数/持仓天数区间/调仓时机/区间)。费率与复权不在此表 —— 运行时从
|
||||
global_config 快照进归档的 config_snapshot,保证可复现。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
from sqlalchemy import Date, DateTime, Integer, Numeric, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
|
||||
class GlobalConfigModel(Base):
|
||||
__tablename__ = "global_config"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True) # 恒为 "default"
|
||||
commission_rate: Mapped[float] = mapped_column(Numeric(10, 6), default=0.0003)
|
||||
stamp_tax_rate: Mapped[float] = mapped_column(Numeric(10, 6), default=0.0005)
|
||||
slippage_rate: Mapped[float] = mapped_column(Numeric(10, 6), default=0.001)
|
||||
min_commission: Mapped[float] = mapped_column(Numeric(10, 4), default=5.0)
|
||||
price_adjustment: Mapped[str] = mapped_column(String(8), default="hfq")
|
||||
benchmark: Mapped[str] = mapped_column(String(16), default="000300.SH")
|
||||
updated_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
|
||||
|
||||
|
||||
class BacktestComboModel(Base):
|
||||
__tablename__ = "backtest_combo"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
name: Mapped[str] = mapped_column(String(64), unique=True)
|
||||
description: Mapped[str] = mapped_column(String(300), default="")
|
||||
strategy_ids_json: Mapped[str] = mapped_column(Text) # JSON list[str]
|
||||
initial_capital: Mapped[float] = mapped_column(Numeric(20, 2), default=1_000_000.0)
|
||||
hold_count: Mapped[int] = mapped_column(Integer, default=20)
|
||||
hold_min_days: Mapped[int] = mapped_column(Integer, default=0)
|
||||
hold_max_days: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
rebalance_freq: Mapped[str] = mapped_column(String(12), default="monthly")
|
||||
start_date: Mapped[date] = mapped_column(Date)
|
||||
end_date: Mapped[date] = mapped_column(Date)
|
||||
version: Mapped[str] = mapped_column(String(16), default="1")
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,135 @@
|
||||
"""回测组合 + 公共配置 Repository 的 SQLAlchemy 实现(2026-09 重构)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.domain.entities.combo import BacktestCombo, GlobalConfig
|
||||
from app.infrastructure.persistence.sqlalchemy.models.combo import (
|
||||
BacktestComboModel,
|
||||
GlobalConfigModel,
|
||||
)
|
||||
|
||||
DEFAULT_CONFIG_ID = "default"
|
||||
|
||||
|
||||
class SqlAlchemyGlobalConfigRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def get(self) -> GlobalConfig:
|
||||
row = self._session.get(GlobalConfigModel, DEFAULT_CONFIG_ID)
|
||||
if row is None:
|
||||
return GlobalConfig() # 未配置过 → 返回默认值(不写库,由调用方决定是否 save)
|
||||
return GlobalConfig(
|
||||
id=row.id,
|
||||
commission_rate=float(row.commission_rate),
|
||||
stamp_tax_rate=float(row.stamp_tax_rate),
|
||||
slippage_rate=float(row.slippage_rate),
|
||||
min_commission=float(row.min_commission),
|
||||
price_adjustment=row.price_adjustment,
|
||||
benchmark=row.benchmark,
|
||||
updated_at=row.updated_at,
|
||||
)
|
||||
|
||||
def save(self, config: GlobalConfig) -> GlobalConfig:
|
||||
row = self._session.get(GlobalConfigModel, DEFAULT_CONFIG_ID)
|
||||
if row is None:
|
||||
row = GlobalConfigModel(id=DEFAULT_CONFIG_ID)
|
||||
self._session.add(row)
|
||||
row.commission_rate = config.commission_rate
|
||||
row.stamp_tax_rate = config.stamp_tax_rate
|
||||
row.slippage_rate = config.slippage_rate
|
||||
row.min_commission = config.min_commission
|
||||
row.price_adjustment = config.price_adjustment
|
||||
row.benchmark = config.benchmark
|
||||
row.updated_at = datetime.now()
|
||||
self._session.flush()
|
||||
return config.model_copy(update={"id": DEFAULT_CONFIG_ID, "updated_at": row.updated_at})
|
||||
|
||||
|
||||
class SqlAlchemyComboRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def save(self, combo: BacktestCombo) -> BacktestCombo:
|
||||
if not combo.id:
|
||||
raise ValueError("需要 id(由调用方生成)")
|
||||
dup = self._session.scalar(
|
||||
select(BacktestComboModel).where(BacktestComboModel.name == combo.name).limit(1)
|
||||
)
|
||||
if dup is not None and dup.id != combo.id:
|
||||
raise ValueError(f"回测组合名已存在:{combo.name}")
|
||||
now = combo.created_at or datetime.now()
|
||||
row = self._session.get(BacktestComboModel, combo.id)
|
||||
strategy_ids_json = json.dumps(combo.strategy_ids, ensure_ascii=False)
|
||||
if row is None:
|
||||
self._session.add(
|
||||
BacktestComboModel(
|
||||
id=combo.id,
|
||||
name=combo.name,
|
||||
description=combo.description,
|
||||
strategy_ids_json=strategy_ids_json,
|
||||
initial_capital=combo.initial_capital,
|
||||
hold_count=combo.hold_count,
|
||||
hold_min_days=combo.hold_min_days,
|
||||
hold_max_days=combo.hold_max_days,
|
||||
rebalance_freq=combo.rebalance_freq,
|
||||
start_date=combo.period[0],
|
||||
end_date=combo.period[1],
|
||||
version=combo.version,
|
||||
created_at=now,
|
||||
)
|
||||
)
|
||||
else:
|
||||
row.name = combo.name
|
||||
row.description = combo.description
|
||||
row.strategy_ids_json = strategy_ids_json
|
||||
row.initial_capital = combo.initial_capital
|
||||
row.hold_count = combo.hold_count
|
||||
row.hold_min_days = combo.hold_min_days
|
||||
row.hold_max_days = combo.hold_max_days
|
||||
row.rebalance_freq = combo.rebalance_freq
|
||||
row.start_date = combo.period[0]
|
||||
row.end_date = combo.period[1]
|
||||
row.version = combo.version
|
||||
self._session.flush()
|
||||
return combo
|
||||
|
||||
def get(self, combo_id: str) -> BacktestCombo | None:
|
||||
row = self._session.get(BacktestComboModel, combo_id)
|
||||
return _to_combo(row) if row else None
|
||||
|
||||
def list(self) -> list[BacktestCombo]:
|
||||
rows = self._session.scalars(
|
||||
select(BacktestComboModel).order_by(BacktestComboModel.created_at.desc())
|
||||
).all()
|
||||
return [_to_combo(r) for r in rows]
|
||||
|
||||
def delete(self, combo_id: str) -> bool:
|
||||
row = self._session.get(BacktestComboModel, combo_id)
|
||||
if row is None:
|
||||
return False
|
||||
self._session.delete(row)
|
||||
return True
|
||||
|
||||
|
||||
def _to_combo(row: BacktestComboModel) -> BacktestCombo:
|
||||
return BacktestCombo(
|
||||
id=row.id,
|
||||
name=row.name,
|
||||
description=row.description,
|
||||
strategy_ids=json.loads(row.strategy_ids_json),
|
||||
initial_capital=float(row.initial_capital),
|
||||
hold_count=row.hold_count,
|
||||
hold_min_days=row.hold_min_days,
|
||||
hold_max_days=row.hold_max_days,
|
||||
rebalance_freq=row.rebalance_freq,
|
||||
period=(row.start_date, row.end_date),
|
||||
version=row.version,
|
||||
created_at=row.created_at,
|
||||
)
|
||||
@@ -1,6 +1,8 @@
|
||||
"""策略 Repository 的 SQLAlchemy 实现(M8.3)。
|
||||
"""选股策略 Repository 的 SQLAlchemy 实现。
|
||||
|
||||
config 以 JSON 存(StrategyDefinition.model_dump);读取时重建实体。
|
||||
config_json 只存选股相关字段(universe/factors/conditions);读取时重建 SelectionStrategy。
|
||||
2026-09 重构:策略库不再持有回测执行参数(selection/rebalance/costs/portfolio/区间),
|
||||
旧行若残留这些键,读出时由 Pydantic 的 extra 忽略策略丢弃(见 _to_entity)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -11,7 +13,7 @@ from datetime import datetime
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.models.strategy import StrategyModel
|
||||
|
||||
|
||||
@@ -19,7 +21,7 @@ class SqlAlchemyStrategyRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def save(self, definition: StrategyDefinition) -> StrategyDefinition:
|
||||
def save(self, definition: SelectionStrategy) -> SelectionStrategy:
|
||||
if not definition.id:
|
||||
raise ValueError("需要 id(由调用方生成)")
|
||||
dup = self._session.scalar(
|
||||
@@ -51,17 +53,17 @@ class SqlAlchemyStrategyRepository:
|
||||
self._session.flush()
|
||||
return definition
|
||||
|
||||
def get(self, strategy_id: str) -> StrategyDefinition | None:
|
||||
def get(self, strategy_id: str) -> SelectionStrategy | None:
|
||||
row = self._session.get(StrategyModel, strategy_id)
|
||||
return _to_entity(row) if row else None
|
||||
|
||||
def get_by_name(self, name: str) -> StrategyDefinition | None:
|
||||
def get_by_name(self, name: str) -> SelectionStrategy | None:
|
||||
row = self._session.scalar(
|
||||
select(StrategyModel).where(StrategyModel.name == name).limit(1)
|
||||
)
|
||||
return _to_entity(row) if row else None
|
||||
|
||||
def list(self) -> list[StrategyDefinition]:
|
||||
def list(self) -> list[SelectionStrategy]:
|
||||
rows = self._session.scalars(
|
||||
select(StrategyModel).order_by(StrategyModel.name)
|
||||
).all()
|
||||
@@ -75,7 +77,15 @@ class SqlAlchemyStrategyRepository:
|
||||
return True
|
||||
|
||||
|
||||
def _to_entity(row: StrategyModel) -> StrategyDefinition:
|
||||
# 旧 strategy.config_json 可能残留的回测执行参数字段(重构前写入)—— 读出时丢弃,
|
||||
# 因为 SelectionStrategy 不再承载它们(已迁到回测组合 / 公共配置)。
|
||||
_LEGACY_BACKTEST_KEYS = frozenset({
|
||||
"selection", "rebalance", "costs", "portfolio", "price_adjustment",
|
||||
"selection_interval_months", "rebalance_interval_months",
|
||||
})
|
||||
|
||||
|
||||
def _to_entity(row: StrategyModel) -> SelectionStrategy:
|
||||
data = json.loads(row.config_json)
|
||||
# 列字段由 DB 行回填,避免与 config_json 重复。
|
||||
# description 必须一并回填:它是列字段(String(300)),save() 会写入,
|
||||
@@ -83,7 +93,10 @@ def _to_entity(row: StrategyModel) -> StrategyDefinition:
|
||||
# (读写不对称:保存的说明看不到,策略库/编辑页都拿不到)。
|
||||
for key in ("name", "version", "description", "spec_type"):
|
||||
data.pop(key, None)
|
||||
return StrategyDefinition(
|
||||
# 丢弃旧行的回测参数字段(Pydantic 默认 forbid extra 会因这些键报错)
|
||||
for key in _LEGACY_BACKTEST_KEYS:
|
||||
data.pop(key, None)
|
||||
return SelectionStrategy(
|
||||
id=row.id, name=row.name, version=row.version, description=row.description,
|
||||
created_at=row.created_at, **data,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,580 @@
|
||||
"""组合回测引擎(2026-09 重构):多选股策略 + 持仓天数区间 + 日/周/月调仓。
|
||||
|
||||
与 `local_engine.TopKBacktestRunner`(单策略、固定 m/y、全卖全买)并存:
|
||||
旧 runner 继续服务 `/api/backtests`(ResearchSpec)与因子测试相关的既有路径,
|
||||
本模块服务新的「回测组合」产品。两者产出**同一种 BacktestResult**,前端可视化无需改动。
|
||||
|
||||
执行模型(用户确认的语义):
|
||||
- **打分 = 并集 + Borda 秩和**:每个选股策略各自对「自己的股票池 ∩ 条件」内的股票
|
||||
按复合因子分排名;合并取并集,综合分 = Σ(1 / 该策略内名次),未进入某策略排名的
|
||||
股票在该策略贡献 0。好处是不假设不同策略的因子分值可比、能容纳各策略股票池不同。
|
||||
- **调仓时机 daily/weekly/monthly**:决定「重新打分 + 调向目标」的节奏。
|
||||
- **持仓天数区间 [Tmin, Tmax]**:
|
||||
* 每个交易日都检查 Tmax —— 持有超过 Tmax 的个股**强制了结**(安全阀,
|
||||
即便调仓是月频也不能让个股远超 Tmax);
|
||||
* 仅在调仓日:把「掉出 TopN 且已持 ≥ Tmin」的卖出(Tmin 防频繁换手),
|
||||
再从 TopN 里补买到 N 只(等权目标,只买不主动减持以尊重 Tmin)。
|
||||
- 成交仍在调仓日收盘(与旧引擎同一时序纪律,无未来函数);涨跌停/停牌沿用旧近似。
|
||||
|
||||
复用 local_engine 的纯工具(涨跌停幅度、NaN 判定、调仓日集合),其余记账逻辑
|
||||
为本模块自包含 —— 刻意不继承旧 runner,避免改动那条已验证的路径。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
|
||||
import pandas as pd
|
||||
|
||||
from app.domain.entities.combo import BacktestCombo, ComboRunSpec, SelectionStrategyRef
|
||||
from app.domain.entities.research import (
|
||||
ActionRecord,
|
||||
BacktestResult,
|
||||
BacktestSummary,
|
||||
ConditionSpec,
|
||||
CostSpec,
|
||||
CurvePoint,
|
||||
FactorSpec,
|
||||
MonthlyReturn,
|
||||
Position,
|
||||
RankedPick,
|
||||
SymbolCurve,
|
||||
Trade,
|
||||
UniverseSpec,
|
||||
YearlyReturn,
|
||||
)
|
||||
from app.quant.composite import build_score_panel
|
||||
from app.quant.local_engine import _limit_up_ratio, _nan, rebalance_dates
|
||||
|
||||
TRADING_DAYS = 252
|
||||
|
||||
|
||||
# ---------- 多策略打分:Borda 秩和 ----------
|
||||
|
||||
|
||||
def _ref_to_specs(ref: SelectionStrategyRef) -> tuple[UniverseSpec, list[FactorSpec], list[ConditionSpec]]:
|
||||
"""把归档快照里的 dict 还原成强类型 spec(喂给既有因子/条件求值器)。"""
|
||||
universe = UniverseSpec.model_validate(ref.universe)
|
||||
factors = [FactorSpec.model_validate(f) for f in ref.factors]
|
||||
conditions = [ConditionSpec.model_validate(c) for c in ref.conditions]
|
||||
return universe, factors, conditions
|
||||
|
||||
|
||||
def borda_combine(panels: list[pd.DataFrame]) -> pd.DataFrame:
|
||||
"""把多张「复合分面板」按 Borda 秩和合并成一张综合分面板。
|
||||
|
||||
每张面板先按日降序排名(1 = 最高分),综合分 = Σ(1/名次);某面板里 NaN(该股
|
||||
不在该策略可评分集)贡献 0。index/columns 取所有面板的并集。纯函数,便于单测。
|
||||
"""
|
||||
if not panels:
|
||||
return pd.DataFrame()
|
||||
dates = sorted({d for p in panels for d in p.index})
|
||||
symbols = sorted({s for p in panels for s in p.columns})
|
||||
borda = pd.DataFrame(0.0, index=dates, columns=symbols)
|
||||
for panel in panels:
|
||||
ranks = panel.rank(axis=1, ascending=False, na_option="keep")
|
||||
contrib = (1.0 / ranks).fillna(0.0)
|
||||
borda = borda.add(contrib.reindex(index=borda.index, columns=borda.columns), fill_value=0.0)
|
||||
return borda
|
||||
|
||||
|
||||
def combine_strategy_scores(
|
||||
daily: pd.DataFrame,
|
||||
strategies: list[SelectionStrategyRef],
|
||||
eligibility_fns: list,
|
||||
) -> tuple[pd.DataFrame, object]:
|
||||
"""多策略 → (综合分面板, 合并合格集闭包)。
|
||||
|
||||
综合分面板 index=trade_date, columns=symbol,值为 Borda 秩和(越大越优先)。
|
||||
合并合格集闭包 `combined(as_of) -> set[symbol] | None`:各策略合格集的并集;
|
||||
全部策略都不过滤时返回 None(= 不过滤,交给面板的 dropna 处理)。
|
||||
"""
|
||||
# 每个策略一张「复合 zscore 面板」(已按方向加权求和)
|
||||
panels: list[pd.DataFrame] = []
|
||||
for ref in strategies:
|
||||
_universe, factors, _conditions = _ref_to_specs(ref)
|
||||
panels.append(build_score_panel(daily, factors))
|
||||
|
||||
borda = borda_combine(panels)
|
||||
|
||||
def combined(as_of: date) -> set[str] | None:
|
||||
sets = []
|
||||
any_filter = False
|
||||
for fn in eligibility_fns:
|
||||
if fn is None:
|
||||
continue
|
||||
s = fn(as_of)
|
||||
if s is not None:
|
||||
sets.append(s)
|
||||
any_filter = True
|
||||
if not any_filter:
|
||||
return None
|
||||
# 并集:任一策略认为合格即合格(Borda 会给没被某策略覆盖的股票较低分,自然靠后)
|
||||
out: set[str] = set()
|
||||
for s in sets:
|
||||
out |= s
|
||||
return out
|
||||
|
||||
return borda, combined
|
||||
|
||||
|
||||
# ---------- 持仓区间回测 runner ----------
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Holding:
|
||||
qty: float
|
||||
entry_date: date
|
||||
entry_price: float
|
||||
entry_ts: object = None # pd.Timestamp:按「交易日」计持仓天数用(自然日会跨周末失真)
|
||||
|
||||
|
||||
_UNIMPLEMENTED_BASE = [
|
||||
"涨跌停按收盘价相对上一有效收盘近似判定(未建模开盘一字 / 集合竞价路径)",
|
||||
"成交假设发生在调仓日收盘(未建模盘中价格路径与流动性冲击)",
|
||||
(
|
||||
"调仓日为「增量调仓」:只卖出超 Tmax / 掉出 TopN 且满 Tmin 的仓位,"
|
||||
"并从 TopN 补买至 N 只;**不主动减持超重仓位**以尊重 Tmin,权重会随行情漂移"
|
||||
"(非严格等权,买入侧按等权目标分配可用现金)"
|
||||
),
|
||||
"Tmax 强制卖出每个交易日检查;Tmin 保护与 TopN 重排仅在调仓日执行",
|
||||
(
|
||||
"多策略打分采用 Borda 秩和(各策略 1/名次 求和):不假设不同策略的因子分值可比,"
|
||||
"但极端情况下某策略覆盖极少股票会使其秩和贡献偏大"
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
class HoldingBandRunner:
|
||||
"""持仓天数区间 + 日/周/月调仓的组合回测 runner。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
combo: BacktestCombo,
|
||||
costs: CostSpec,
|
||||
score: pd.DataFrame,
|
||||
close: pd.DataFrame,
|
||||
eligibility_fn=None,
|
||||
) -> None:
|
||||
self.combo = combo
|
||||
self.costs = costs
|
||||
close = close.copy()
|
||||
close.index = pd.to_datetime(close.index)
|
||||
self.close = close.sort_index()
|
||||
self.score = score.reindex(self.close.index).sort_index()
|
||||
self.eligibility_fn = eligibility_fn
|
||||
self.selection_history: list[RankedPick] = []
|
||||
self.signal_history: list[ActionRecord] = []
|
||||
self.traded_symbols: list[str] = []
|
||||
self._traded: set[str] = set()
|
||||
self._no_prev_close: set[str] = set()
|
||||
|
||||
# ---- 主循环 ----
|
||||
|
||||
def run(self) -> BacktestResult:
|
||||
start, end = self.spec_period()
|
||||
dates = [d for d in self.close.index if start <= d.date() <= end]
|
||||
if not dates:
|
||||
raise ValueError(f"回测区间 {start}~{end} 内没有任何行情数据,无法回测")
|
||||
|
||||
cadence = set(
|
||||
rebalance_dates(self.close.index, self.combo.rebalance_freq, start, end)
|
||||
)
|
||||
|
||||
cash = float(self.combo.initial_capital)
|
||||
holdings: dict[str, _Holding] = {}
|
||||
equity_rows: dict[pd.Timestamp, float] = {}
|
||||
trades: list[Trade] = []
|
||||
positions: list[Position] = []
|
||||
notional: list[float] = []
|
||||
cum: dict[str, float] = {}
|
||||
curve_rows: dict[str, list[CurvePoint]] = {}
|
||||
|
||||
n = self.combo.hold_count
|
||||
tmin = self.combo.hold_min_days
|
||||
tmax = self.combo.hold_max_days
|
||||
# 交易日位置索引:持仓天数按「交易日」计(Tmin/Tmax 的自然单位),避免跨周末失真
|
||||
self._tday_pos = {ts: i for i, ts in enumerate(self.close.index)}
|
||||
|
||||
def _equity(d: pd.Timestamp) -> float:
|
||||
total = cash
|
||||
for s, h in holdings.items():
|
||||
px = self.close.at[d, s] if d in self.close.index else None
|
||||
if px is None or (isinstance(px, float) and math.isnan(px)):
|
||||
continue
|
||||
total += h.qty * float(px)
|
||||
return total
|
||||
|
||||
for d in dates:
|
||||
day = d.date()
|
||||
# 1) 用昨日持仓结算当日个股收益(与组合净值同一时序口径)
|
||||
self._accrue(d, holdings, cum, curve_rows)
|
||||
|
||||
# 2) Tmax 安全阀:**每个交易日**强制了结超期仓位(不只调仓日)
|
||||
if tmax is not None:
|
||||
cash = self._force_exit_over_max(d, day, holdings, cash, trades, tmax)
|
||||
|
||||
# 3) 调仓日:重新打分 + 增量调向目标 N 只
|
||||
if d in cadence:
|
||||
topn = self._top_n(d, n)
|
||||
cash = self._rebalance_to_target(
|
||||
d, day, topn, holdings, cash, trades, positions, notional, tmin, n
|
||||
)
|
||||
|
||||
equity_rows[d] = _equity(d)
|
||||
self._mark_curve(d, holdings, cum, curve_rows)
|
||||
|
||||
equity = pd.Series(equity_rows).sort_index()
|
||||
return self._to_result(equity, trades, positions, notional, cum, curve_rows)
|
||||
|
||||
def spec_period(self) -> tuple[date, date]:
|
||||
return self.combo.period
|
||||
|
||||
# ---- 打分 / 选股 ----
|
||||
|
||||
def _top_n(self, d: pd.Timestamp, n: int) -> list[str]:
|
||||
score_d = self.score.loc[d].dropna()
|
||||
elig = self.eligibility_fn(d.date()) if self.eligibility_fn else None
|
||||
if elig is not None:
|
||||
score_d = score_d[score_d.index.isin(elig)]
|
||||
ranked = score_d.sort_values(ascending=False)
|
||||
top = ranked.head(n).index.tolist()
|
||||
day = d.date()
|
||||
for r, sym in enumerate(top, start=1):
|
||||
self.selection_history.append(
|
||||
RankedPick(date=day, symbol=sym, rank=r, score=round(float(ranked[sym]), 6))
|
||||
)
|
||||
return top
|
||||
|
||||
# ---- Tmax 强制了结(每日) ----
|
||||
|
||||
def _force_exit_over_max(self, d, day, holdings, cash, trades, tmax) -> float:
|
||||
prev_d = self.prev_close_at(d)
|
||||
close_d = self.close.loc[d]
|
||||
for s in [s for s in list(holdings)]:
|
||||
h = holdings[s]
|
||||
held = self._held_trading_days(h.entry_ts, d)
|
||||
if held <= tmax:
|
||||
continue
|
||||
c = close_d.get(s)
|
||||
p = prev_d.get(s) if prev_d is not None else None
|
||||
if _nan(c):
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="SELL", filled=False,
|
||||
reject_reason=f"持有 {held} 天超 Tmax={tmax},但当日无行情,顺延")
|
||||
)
|
||||
continue
|
||||
if not _nan(p) and p > 0 and c / p <= 1.0 - (_limit_up_ratio(s) - 1.0):
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="SELL", filled=False,
|
||||
reject_reason=f"持有 {held} 天超 Tmax={tmax},但跌停无法卖出,顺延")
|
||||
)
|
||||
continue
|
||||
cash = self._sell(s, h, float(c), day, cash, trades, holdings)
|
||||
return cash
|
||||
|
||||
# ---- 调仓日:增量调向目标 ----
|
||||
|
||||
def _rebalance_to_target(self, d, day, topn, holdings, cash, trades, positions, notional, tmin, n) -> float:
|
||||
close_d = self.close.loc[d]
|
||||
prev_d = self.prev_close_at(d)
|
||||
topn_set = set(topn)
|
||||
|
||||
# a) 卖出:掉出 TopN 且已满 Tmin 的(Tmin 保护:未满 Tmin 即使掉出也暂留)
|
||||
for s in [s for s in list(holdings)]:
|
||||
if s in topn_set:
|
||||
continue
|
||||
h = holdings[s]
|
||||
held = self._held_trading_days(h.entry_ts, d)
|
||||
if held < tmin:
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="SELL", filled=False,
|
||||
reject_reason=f"掉出 TopN 但仅持 {held} 天 < Tmin={tmin},暂留")
|
||||
)
|
||||
continue
|
||||
c = close_d.get(s)
|
||||
p = prev_d.get(s) if prev_d is not None else None
|
||||
if _nan(c):
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="SELL", filled=False,
|
||||
reject_reason="掉出 TopN,但当日无行情,保留到下一调仓")
|
||||
)
|
||||
continue
|
||||
if not _nan(p) and p > 0 and c / p <= 1.0 - (_limit_up_ratio(s) - 1.0):
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="SELL", filled=False,
|
||||
reject_reason="掉出 TopN,但跌停无法卖出,保留到下一调仓")
|
||||
)
|
||||
continue
|
||||
cash = self._sell(s, h, float(c), day, cash, trades, holdings)
|
||||
|
||||
# b) 补买:从 TopN 里挑尚未持有的,按等权目标用可用现金买入,直到 N 只或现金耗尽
|
||||
current = [s for s in topn if s in holdings and holdings[s].qty > 0]
|
||||
need = [s for s in topn if s not in holdings or holdings[s].qty <= 0]
|
||||
slots_left = max(0, n - len(current))
|
||||
buys = need[:slots_left]
|
||||
if not buys:
|
||||
self._record_positions(d, day, holdings, positions)
|
||||
return cash
|
||||
|
||||
# 等权目标:每只 ≈ 当前权益 / N;单只预算 = min(目标, 可用现金均分)
|
||||
equity_now = cash + sum(
|
||||
holdings[s].qty * float(self.close.at[d, s])
|
||||
for s in holdings
|
||||
if holdings[s].qty > 0 and not _nan(self.close.at[d, s])
|
||||
)
|
||||
target_each = equity_now / n if n > 0 else 0.0
|
||||
per_budget = min(target_each, cash / len(buys)) if buys else 0.0
|
||||
|
||||
for s in buys:
|
||||
c = close_d.get(s)
|
||||
p = prev_d.get(s) if prev_d is not None else None
|
||||
if _nan(c):
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="BUY", filled=False,
|
||||
reject_reason="无行情(停牌),无法买入")
|
||||
)
|
||||
continue
|
||||
if not _nan(p) and p > 0 and c / p >= _limit_up_ratio(s):
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="BUY", filled=False,
|
||||
reject_reason="涨停,无法追买")
|
||||
)
|
||||
continue
|
||||
budget = min(per_budget, cash)
|
||||
if budget <= 1e-9:
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="BUY", filled=False,
|
||||
reject_reason="可用现金不足,未成交")
|
||||
)
|
||||
continue
|
||||
ok, spent = self._buy(s, budget, d, float(c), day, holdings, notional)
|
||||
if not ok:
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="BUY", filled=False,
|
||||
reject_reason="预算不足以覆盖最低佣金,未成交")
|
||||
)
|
||||
continue
|
||||
cash -= spent
|
||||
|
||||
self._record_positions(d, day, holdings, positions)
|
||||
return cash
|
||||
|
||||
# ---- 买卖原子操作 ----
|
||||
|
||||
def _sell(self, s, h, close_price, day, cash, trades, holdings) -> float:
|
||||
proceeds = h.qty * close_price * (1 - self.costs.slippage_rate)
|
||||
commission = max(proceeds * self.costs.commission_rate, self.costs.min_commission)
|
||||
fee = commission + proceeds * self.costs.stamp_tax_rate
|
||||
cash += proceeds - fee
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="SELL", filled=True, price=close_price)
|
||||
)
|
||||
trades.append(
|
||||
Trade(
|
||||
entry_date=h.entry_date, exit_date=day, symbol=s,
|
||||
entry_price=h.entry_price, exit_price=close_price,
|
||||
return_pct=(close_price / h.entry_price - 1.0) * 100,
|
||||
)
|
||||
)
|
||||
holdings.pop(s, None)
|
||||
return cash
|
||||
|
||||
def _buy(self, s, budget, d, close_price, day, holdings, notional) -> tuple[bool, float]:
|
||||
price_in = close_price * (1 + self.costs.slippage_rate)
|
||||
commission = max(budget * self.costs.commission_rate, self.costs.min_commission)
|
||||
invest = budget - commission
|
||||
if invest <= 0:
|
||||
return False, 0.0
|
||||
qty = invest / price_in
|
||||
prev = self.prev_close_at(d)
|
||||
pv = prev.get(s) if prev is not None else float("nan")
|
||||
if _nan(pv) or pv <= 0:
|
||||
self._no_prev_close.add(s)
|
||||
holdings[s] = _Holding(qty=qty, entry_date=day, entry_price=price_in, entry_ts=d)
|
||||
notional.append(budget)
|
||||
self.signal_history.append(
|
||||
ActionRecord(date=day, symbol=s, signal="BUY", filled=True, price=round(price_in, 4))
|
||||
)
|
||||
if s not in self._traded:
|
||||
self._traded.add(s)
|
||||
self.traded_symbols.append(s)
|
||||
return True, budget
|
||||
|
||||
# ---- 辅助 ----
|
||||
|
||||
def _held_trading_days(self, entry_ts, current_ts) -> int:
|
||||
"""从入场到当前经过的**交易日**数(不含入场当日)。"""
|
||||
e = self._tday_pos.get(entry_ts)
|
||||
c = self._tday_pos.get(current_ts)
|
||||
if e is None or c is None:
|
||||
return 0
|
||||
return max(0, c - e)
|
||||
|
||||
def prev_close_at(self, d: pd.Timestamp):
|
||||
"""d 的前一有效收盘行(用于涨跌停判定);不存在返回 None。"""
|
||||
idx = self.close.index
|
||||
pos = idx.get_loc(d) if d in idx else None
|
||||
if pos is None or pos == 0:
|
||||
return None
|
||||
prev_ts = idx[pos - 1]
|
||||
return self.close.loc[prev_ts]
|
||||
|
||||
def _accrue(self, d, holdings, cum, curve_rows) -> None:
|
||||
prev = self.prev_close_at(d)
|
||||
if prev is None:
|
||||
return
|
||||
close_d = self.close.loc[d]
|
||||
for s in holdings:
|
||||
c, p = close_d.get(s), prev.get(s)
|
||||
if _nan(c) or _nan(p) or p <= 0:
|
||||
continue
|
||||
cum[s] = cum.get(s, 1.0) * (float(c) / float(p))
|
||||
curve_rows.setdefault(s, []).append(
|
||||
CurvePoint(date=d.date(), value=round((cum[s] - 1.0) * 100, 4))
|
||||
)
|
||||
|
||||
def _mark_curve(self, d, holdings, cum, curve_rows) -> None:
|
||||
day = d.date()
|
||||
for s in holdings:
|
||||
pts = curve_rows.setdefault(s, [])
|
||||
if pts and pts[-1].date == day:
|
||||
continue
|
||||
pts.append(CurvePoint(date=day, value=round((cum.get(s, 1.0) - 1.0) * 100, 4)))
|
||||
|
||||
def _record_positions(self, d, day, holdings, positions) -> None:
|
||||
total = sum(
|
||||
h.qty * float(self.close.at[d, s])
|
||||
for s, h in holdings.items()
|
||||
if h.qty > 0 and not _nan(self.close.at[d, s])
|
||||
)
|
||||
if total <= 0:
|
||||
return
|
||||
for s, h in holdings.items():
|
||||
if h.qty > 0 and not _nan(self.close.at[d, s]):
|
||||
positions.append(
|
||||
Position(date=day, symbol=s, weight=float(h.qty * self.close.at[d, s] / total))
|
||||
)
|
||||
|
||||
# ---- 结果装配(与旧 runner 同构,保证前端可视化不变) ----
|
||||
|
||||
def _to_result(self, equity, trades, positions, notional, cum, curve_rows) -> BacktestResult:
|
||||
start, end = equity.index[0].date(), equity.index[-1].date()
|
||||
init = float(self.combo.initial_capital)
|
||||
final = float(equity.iloc[-1])
|
||||
rets = equity.pct_change().dropna()
|
||||
nn = len(rets)
|
||||
total_ret = (final / init - 1.0) * 100 if init else 0.0
|
||||
annual = (
|
||||
((final / init) ** (TRADING_DAYS / max(nn, 1)) - 1.0) * 100
|
||||
if final > 0 and init > 0 else -100.0
|
||||
)
|
||||
mean_r = float(rets.mean()) if nn else 0.0
|
||||
std_r = float(rets.std(ddof=1)) if nn else 0.0
|
||||
sharpe = mean_r / std_r * math.sqrt(TRADING_DAYS) if std_r and mean_r else 0.0
|
||||
vol = std_r * math.sqrt(TRADING_DAYS) * 100
|
||||
dd = (equity / equity.cummax() - 1.0).min() * 100
|
||||
wins = [t for t in trades if t.return_pct > 0]
|
||||
win_rate = len(wins) / len(trades) * 100 if trades else 0.0
|
||||
avg_turn = (sum(notional) / len(notional) / ((init + final) / 2)) * 100 if notional else 0.0
|
||||
|
||||
eq_pts = [CurvePoint(date=d.date(), value=round(float(v), 2)) for d, v in equity.items()]
|
||||
dd_series = (equity / equity.cummax() - 1.0) * 100
|
||||
drawdown = [CurvePoint(date=d.date(), value=round(float(v), 3)) for d, v in dd_series.items()]
|
||||
|
||||
monthly: list[MonthlyReturn] = []
|
||||
yearly: list[YearlyReturn] = []
|
||||
if len(equity) > 1:
|
||||
m = equity.resample("ME").last().pct_change().dropna()
|
||||
monthly = [
|
||||
MonthlyReturn(year=int(d.year), month=int(d.month), return_pct=round(float(v) * 100, 3))
|
||||
for d, v in m.items()
|
||||
]
|
||||
y = equity.resample("YE").last().pct_change().dropna()
|
||||
yearly = [
|
||||
YearlyReturn(year=int(d.year), return_pct=round(float(v) * 100, 3))
|
||||
for d, v in y.items()
|
||||
]
|
||||
|
||||
summary = BacktestSummary(
|
||||
start=start, end=end, initial_capital=round(init, 2), final_equity=round(final, 2),
|
||||
total_return_pct=round(total_ret, 3), annual_return_pct=round(annual, 3),
|
||||
sharpe=round(sharpe, 3), max_drawdown_pct=round(float(dd), 3),
|
||||
volatility_pct=round(vol, 3), win_rate_pct=round(win_rate, 2),
|
||||
total_trades=len(trades), avg_turnover_pct=round(avg_turn, 2),
|
||||
)
|
||||
curves = self._symbol_curves(curve_rows, cum)
|
||||
return BacktestResult(
|
||||
summary=summary, equity_curve=eq_pts, drawdown=drawdown,
|
||||
monthly_returns=monthly, yearly_returns=yearly,
|
||||
positions=positions, trades=trades,
|
||||
selection_history=self.selection_history,
|
||||
signal_history=self.signal_history,
|
||||
fills=[a for a in self.signal_history if a.filled],
|
||||
symbol_curves=curves,
|
||||
turnover_pct=round(sum(notional) / max(init, 1) * 100, 2),
|
||||
unimplemented=self._unimplemented(),
|
||||
config_snapshot={}, # 由服务层填入 ComboRunSpec(含策略+成本快照)
|
||||
)
|
||||
|
||||
def _symbol_curves(self, curve_rows, cum) -> list[SymbolCurve]:
|
||||
marks: dict[str, list[ActionRecord]] = {}
|
||||
for a in self.signal_history:
|
||||
if a.filled and a.symbol:
|
||||
marks.setdefault(a.symbol, []).append(a)
|
||||
out: list[SymbolCurve] = []
|
||||
for s, points in curve_rows.items():
|
||||
if not points:
|
||||
continue
|
||||
out.append(SymbolCurve(
|
||||
symbol=s, points=points, marks=marks.get(s, []),
|
||||
final_return_pct=round((cum.get(s, 1.0) - 1.0) * 100, 4),
|
||||
))
|
||||
out.sort(key=lambda c: abs(c.final_return_pct), reverse=True)
|
||||
return out
|
||||
|
||||
def _unimplemented(self) -> list[str]:
|
||||
notes = list(_UNIMPLEMENTED_BASE)
|
||||
if self._no_prev_close:
|
||||
notes.append(
|
||||
f"有 {len(self._no_prev_close)} 只标的成交时缺少上一有效收盘价,"
|
||||
"无法判定涨停(数据窗口起点或长期停牌后复牌),按可买处理"
|
||||
)
|
||||
c = self.combo
|
||||
notes.append(
|
||||
f"组合参数:N={c.hold_count}、持仓区间 [{c.hold_min_days}, "
|
||||
f"{c.hold_max_days if c.hold_max_days is not None else '∞'}] 天、"
|
||||
f"调仓 {c.rebalance_freq}、引用 {len(c.strategy_ids)} 个选股策略"
|
||||
)
|
||||
return notes
|
||||
|
||||
|
||||
def run_combo_backtest(
|
||||
*,
|
||||
combo: BacktestCombo,
|
||||
strategies: list[SelectionStrategyRef],
|
||||
costs: CostSpec,
|
||||
price_adjustment: str,
|
||||
daily: pd.DataFrame,
|
||||
eligibility_fns: list,
|
||||
) -> BacktestResult:
|
||||
"""组合回测入口:多策略 Borda 打分 → 持仓区间 runner → BacktestResult。
|
||||
|
||||
`eligibility_fns` 与 `strategies` 一一对应(每个策略一个「as_of→合格集」闭包,
|
||||
可为 None 表示该策略无额外过滤);由服务层用既有 selection 求值器装配。
|
||||
返回结果的 config_snapshot 由调用方填入 ComboRunSpec(含策略+成本快照)以保证可复现。
|
||||
"""
|
||||
score, combined_elig = combine_strategy_scores(daily, strategies, eligibility_fns)
|
||||
close = daily.pivot(index="trade_date", columns="symbol", values="close").sort_index()
|
||||
runner = HoldingBandRunner(
|
||||
combo=combo, costs=costs, score=score, close=close, eligibility_fn=combined_elig,
|
||||
)
|
||||
result = runner.run()
|
||||
# 固化可复现规格(AGENT.md §21):组合参数 + 当时各策略定义 + 当时成本/复权
|
||||
result.config_snapshot = ComboRunSpec(
|
||||
combo=combo, strategies=strategies, costs=costs, price_adjustment=price_adjustment,
|
||||
).model_dump(mode="json")
|
||||
return result
|
||||
@@ -145,10 +145,12 @@ def rebalance_dates(
|
||||
⚠️ 刻意**不丢弃锚点月**:若把锚点月整体过滤掉,m=y=6 且起始日非月初时
|
||||
会白等 6 个月才首次建仓(净值在前期恒等于初始资金,指标明显失真)。
|
||||
|
||||
every_months=None:沿用 weekly / monthly 频率(原行为,保持向后兼容)。
|
||||
every_months=None:沿用 weekly / monthly / **daily** 频率。
|
||||
- daily:区间内**每个交易日**都是调仓日(回测组合的「日频调仓」)。
|
||||
- weekly / monthly:原行为,保持向后兼容。
|
||||
"""
|
||||
firsts = _week_firsts(index) if rebalance == "weekly" else _month_firsts(index)
|
||||
if every_months and every_months > 0:
|
||||
firsts = _week_firsts(index) if rebalance == "weekly" else _month_firsts(index)
|
||||
# 锚点 = start 所在月内首个 >= start 的交易日(可能不是该月首个交易日)
|
||||
days = pd.DatetimeIndex(index)
|
||||
after_start = days[days >= pd.Timestamp(start)]
|
||||
@@ -165,7 +167,12 @@ def rebalance_dates(
|
||||
and ts.date() >= start
|
||||
and ts != anchor_ts
|
||||
]
|
||||
elif rebalance == "daily":
|
||||
# 日频:区间内每个交易日
|
||||
days = pd.DatetimeIndex(index)
|
||||
out = [ts for ts in days if ts.date() >= start]
|
||||
else:
|
||||
firsts = _week_firsts(index) if rebalance == "weekly" else _month_firsts(index)
|
||||
out = [ts for ts in firsts if ts.date() >= start]
|
||||
if end is not None:
|
||||
out = [ts for ts in out if ts.date() <= end]
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""策略说明书生成器:ResearchSpec / StrategyDefinition → 一句话说明 + 计算公式 + 步骤 + 注意事项。
|
||||
"""策略说明书生成器:ResearchSpec(回测)/ SelectionStrategy(选股策略)→ 说明 + 公式 + 步骤 + 注意事项。
|
||||
|
||||
纯函数模块:无 IO、无 DB、不调用引擎,因而可被 API 复用(含**未保存**的策略即时预览)
|
||||
并被单测直接覆盖。
|
||||
@@ -29,12 +29,12 @@ from app.domain.entities.market import (
|
||||
DAILY_BASIC_NUMERIC_FIELDS,
|
||||
)
|
||||
from app.domain.entities.research import ResearchSpec
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.domain.entities.strategy import SelectionStrategy, StrategyDefinition
|
||||
from app.quant.factors import FactorDef, FactorError, get_factor
|
||||
|
||||
# StrategyDefinition 没有回测区间字段(区间在回测时补全),但 to_research_spec 的
|
||||
# period 是必填的 —— 用一个不可能被误读为真实区间的占位区间满足校验,
|
||||
# 真正展示时以 `period_known=False` 走占位文案(不抛错、也不假装知道区间)。
|
||||
# 选股策略(SelectionStrategy)不含回测参数,走 _describe_selection_only 专用分支;
|
||||
# 以下占位常量仅保留给历史 ResearchSpec 路径兼容(当前 describe_strategy 不再对
|
||||
# SelectionStrategy 调用 to_research_spec)。
|
||||
_PLACEHOLDER_PERIOD: tuple[date, date] = (date(1900, 1, 1), date(1900, 1, 2))
|
||||
_PERIOD_UNKNOWN_TEXT = "(回测区间未指定:策略定义本身不含 period,展开回测后才确定)"
|
||||
|
||||
@@ -75,21 +75,25 @@ class StrategyDoc(BaseModel):
|
||||
|
||||
|
||||
def describe_strategy(
|
||||
spec: ResearchSpec | StrategyDefinition,
|
||||
spec: ResearchSpec | SelectionStrategy,
|
||||
*,
|
||||
factor_meta: Mapping[str, FactorDef] | None = None,
|
||||
) -> StrategyDoc:
|
||||
"""生成策略说明书。
|
||||
|
||||
`spec`:ResearchSpec(回测页参数,含 period)或 StrategyDefinition(策略库资产,
|
||||
无 period)。后者经 `to_research_spec(period=占位)` 展开 —— 区间缺失只影响展示文案
|
||||
(走 `_PERIOD_UNKNOWN_TEXT`),不影响其余推导,更不抛错。
|
||||
`spec`:ResearchSpec(回测页参数,含 period/costs/selection)或 SelectionStrategy
|
||||
(策略库资产,只有选股条件)。后者走 `_describe_selection_only`,只讲「怎么选」,
|
||||
如实声明资金/持仓/调仓/成本/区间不在策略内(在回测组合里定)。
|
||||
|
||||
`factor_meta`:可选的因子元数据覆盖/补充(如内置注册表之外的实验因子)。
|
||||
查找顺序为 `factor_meta` → `app.quant.factors.get_factor`;两处都没有则记入
|
||||
warnings(说明该因子的含义/公式/方向未知,执行期会直接报错)。
|
||||
"""
|
||||
rspec, period_text, period_known = _coerce_spec(spec)
|
||||
coerced = _coerce_spec(spec)
|
||||
# 选股策略(无回测参数)走专用说明;回测 spec 走原有完整路径
|
||||
if isinstance(coerced, (SelectionStrategy, StrategyDefinition)) and not isinstance(coerced, ResearchSpec):
|
||||
return _describe_selection_only(coerced, factor_meta)
|
||||
rspec, period_text, period_known = coerced
|
||||
warnings: list[str] = []
|
||||
factors = _describe_factors(rspec, factor_meta, warnings)
|
||||
conditions = _describe_conditions(rspec, warnings)
|
||||
@@ -101,22 +105,117 @@ def describe_strategy(
|
||||
return StrategyDoc(summary=summary, formula=formula, steps=steps, warnings=warnings)
|
||||
|
||||
|
||||
def _describe_selection_only(st, factor_meta) -> StrategyDoc:
|
||||
"""选股策略的说明:只讲「怎么选」(股票池 + 因子 + 条件),不涉及回测执行参数。
|
||||
|
||||
资金 / 持仓数 / 持仓时间 / 调仓时机 / 费率 / 复权 / 区间都不属于选股策略,
|
||||
在回测组合里才确定 —— 这里如实声明,避免读者以为策略自带这些口径。
|
||||
"""
|
||||
warnings: list[str] = []
|
||||
factors = _describe_factors_from_specs(st.factors, factor_meta, warnings)
|
||||
pool = _universe_text(st.universe)
|
||||
cond_lines = _condition_lines(st.conditions, warnings) if st.conditions else []
|
||||
|
||||
factor_names = "、".join(f["name"] for f in factors) or "(无)"
|
||||
pool_desc = _short_pool(st.universe)
|
||||
summary = (
|
||||
f"选股策略「{st.name}」:在{pool_desc}内"
|
||||
+ (f"先通过 {len(st.conditions)} 条过滤条件,再" if st.conditions else "")
|
||||
+ f"按因子({factor_names})打分排序,供回测组合取 TopN 持仓。"
|
||||
)
|
||||
formula_lines = [
|
||||
"【选股口径】(仅定义「怎么选」,不含回测执行参数)",
|
||||
f" 股票池:{pool}",
|
||||
" 打分:score = Σ 权重 × 因子值(截面 z-score 标准化后加权,越大越优先)",
|
||||
]
|
||||
for f in factors:
|
||||
formula_lines.append(f" · {f['line']}")
|
||||
if cond_lines:
|
||||
formula_lines.append(" 过滤条件(AND,先于打分执行):")
|
||||
formula_lines.extend(f" · {ln}" for ln in cond_lines)
|
||||
formula_lines.append(
|
||||
" ⚠ 持仓数量 / 持仓天数区间 / 调仓时机 / 起始资金 / 费率 / 复权口径 / 回测区间"
|
||||
"均不在本策略内 —— 它们在「回测组合」中指定,运行时与公共配置合并。"
|
||||
)
|
||||
steps = [
|
||||
"1. 按股票池口径筛出候选 universe(市场 / 剔 ST / 上市天数 / 指数成分)。",
|
||||
"2." + (" 逐条求值过滤条件(AND),剔除不满足者。" if st.conditions else " (未设过滤条件,候选 = universe。)"),
|
||||
"3. 对剩余股票按上述因子打分并降序排列 → 得到候选排名(TopN 在回测组合里截取)。",
|
||||
]
|
||||
warnings.append(
|
||||
"本说明只覆盖选股口径;回测的资金/持仓/调仓/成本/区间由「回测组合」+「公共配置」决定,"
|
||||
"此处无法给出收益公式与成交口径。"
|
||||
)
|
||||
return StrategyDoc(summary=summary, formula="\n".join(formula_lines), steps=steps, warnings=warnings)
|
||||
|
||||
|
||||
def _short_pool(universe) -> str:
|
||||
parts = []
|
||||
if universe.index_code:
|
||||
parts.append(f"{universe.index_code} 成分股")
|
||||
elif universe.symbols:
|
||||
parts.append(f"指定的 {len(universe.symbols)} 只白名单股票")
|
||||
else:
|
||||
parts.append("全市场 A 股")
|
||||
if universe.exclude_st:
|
||||
parts.append("剔除 ST ")
|
||||
return "、".join(parts)
|
||||
|
||||
|
||||
def _describe_factors_from_specs(factor_specs, factor_meta, warnings) -> list[dict]:
|
||||
"""从 FactorSpec 列表生成 [{label, line}](复用既有因子元数据查找逻辑)。"""
|
||||
out: list[dict] = []
|
||||
for fs in factor_specs:
|
||||
meta = _lookup_factor(fs.name, factor_meta)
|
||||
label = fs.name
|
||||
direction = "越高越好"
|
||||
formula_hint = ""
|
||||
if meta is None:
|
||||
warnings.append(
|
||||
f"因子 {fs.name} 未在注册表中找到元数据:含义/公式/方向未知,"
|
||||
"执行期会报错;说明里只能给出名字与权重。"
|
||||
)
|
||||
else:
|
||||
label = meta.brief or meta.description or fs.name
|
||||
direction = "越低越好" if meta.direction == "lower_is_better" else "越高越好"
|
||||
formula_hint = f"({meta.formula})" if meta.formula else ""
|
||||
out.append({
|
||||
"name": fs.name,
|
||||
"label": label if label != fs.name else fs.name,
|
||||
"line": f"{fs.name}{formula_hint},权重 {fs.weight},{direction}",
|
||||
})
|
||||
return out
|
||||
|
||||
|
||||
def _condition_lines(conditions, warnings) -> list[str]:
|
||||
"""把 ConditionSpec 列表渲染成可读行(复用既有字段域校验逻辑)。"""
|
||||
lines: list[str] = []
|
||||
for c in conditions:
|
||||
op = _OP_TEXT.get(c.op, c.op)
|
||||
right = f"字段 {c.ref}" if c.ref else f"{c.value}"
|
||||
lines.append(f"{c.field} {op} {right}")
|
||||
return lines
|
||||
|
||||
|
||||
# ---------- 入参归一化 ----------
|
||||
|
||||
|
||||
def _coerce_spec(spec) -> tuple[ResearchSpec, str, bool]:
|
||||
"""归一化为 (ResearchSpec, 区间展示文本, 区间是否已知)。
|
||||
def _coerce_spec(spec):
|
||||
"""归一化输入。
|
||||
|
||||
StrategyDefinition 无 period:按需求用占位区间展开并如实标记「区间未知」,
|
||||
而不是抛错(策略库列表/详情页也要能看说明)。
|
||||
- ResearchSpec(回测页参数,含 period/costs/selection)→ (ResearchSpec, 区间文本, True),
|
||||
走完整回测说明书路径。
|
||||
- SelectionStrategy(策略库资产,**只有选股条件**,无 period/costs/selection)→
|
||||
直接返回该对象,由 `describe_selection_strategy` 生成「选股口径」说明。
|
||||
重构后策略库不再持有回测执行参数,故不能也不应假装展开成 ResearchSpec。
|
||||
"""
|
||||
if isinstance(spec, StrategyDefinition):
|
||||
return spec.to_research_spec(period=_PLACEHOLDER_PERIOD), _PERIOD_UNKNOWN_TEXT, False
|
||||
if isinstance(spec, ResearchSpec):
|
||||
start, end = spec.period
|
||||
return spec, f"{start.isoformat()} ~ {end.isoformat()}", True
|
||||
return (spec, f"{start.isoformat()} ~ {end.isoformat()}", True)
|
||||
if isinstance(spec, (SelectionStrategy, StrategyDefinition)):
|
||||
return spec # 选股策略:单独处理
|
||||
raise TypeError(
|
||||
"describe_strategy 只接受 ResearchSpec 或 StrategyDefinition,"
|
||||
"describe_strategy 只接受 ResearchSpec 或 SelectionStrategy,"
|
||||
f"收到 {type(spec).__name__}"
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
"""公共配置 + 回测组合 API 测试(TestClient + 内存 SQLite)。
|
||||
|
||||
覆盖:
|
||||
- /api/config GET 返回默认值、PUT 持久化;
|
||||
- /api/combos CRUD(name 唯一、原地更新保留 created_at、删除);
|
||||
- /api/combos/{id}/run 在策略缺失时提前 400(不等到后台才失败)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from app.api import deps
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.main import app
|
||||
from fastapi.testclient import TestClient
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def client(tmp_path) -> TestClient:
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'combo_api.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()
|
||||
|
||||
|
||||
class TestConfigApi:
|
||||
def test_get_returns_defaults(self, client: TestClient) -> None:
|
||||
cfg = client.get("/api/config").json()
|
||||
assert cfg["id"] == "default"
|
||||
assert cfg["price_adjustment"] == "hfq"
|
||||
assert cfg["min_commission"] == 5.0
|
||||
|
||||
def test_put_persists(self, client: TestClient) -> None:
|
||||
body = {
|
||||
"commission_rate": 0.00025, "stamp_tax_rate": 0.0005,
|
||||
"slippage_rate": 0.0008, "min_commission": 3.0,
|
||||
"price_adjustment": "qfq", "benchmark": "000905.SH",
|
||||
}
|
||||
resp = client.put("/api/config", json=body)
|
||||
assert resp.status_code == 200
|
||||
got = client.get("/api/config").json()
|
||||
assert got["commission_rate"] == pytest.approx(0.00025)
|
||||
assert got["price_adjustment"] == "qfq"
|
||||
assert got["benchmark"] == "000905.SH"
|
||||
|
||||
|
||||
class TestCombosApi:
|
||||
def _make_strategy(self, client: TestClient, name: str) -> str:
|
||||
resp = client.post(
|
||||
"/api/strategies",
|
||||
json={"name": name, "factors": [{"name": "dividend_yield", "weight": 1}]},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
return resp.json()["id"]
|
||||
|
||||
def test_crud(self, client: TestClient) -> None:
|
||||
sid = self._make_strategy(client, "高股息")
|
||||
body = {
|
||||
"name": "组合A", "strategy_ids": [sid],
|
||||
"initial_capital": 500000, "hold_count": 10,
|
||||
"hold_min_days": 5, "hold_max_days": 30,
|
||||
"rebalance_freq": "weekly", "period": ["2024-01-01", "2024-06-01"],
|
||||
}
|
||||
created = client.post("/api/combos", json=body)
|
||||
assert created.status_code == 200
|
||||
cid = created.json()["id"]
|
||||
assert cid.startswith("CMB-")
|
||||
assert created.json()["hold_max_days"] == 30
|
||||
|
||||
assert len(client.get("/api/combos").json()) == 1
|
||||
detail = client.get(f"/api/combos/{cid}").json()
|
||||
assert detail["strategy_ids"] == [sid]
|
||||
assert detail["rebalance_freq"] == "weekly"
|
||||
|
||||
# 原地更新保留 id 与 created_at
|
||||
created_at = detail["created_at"]
|
||||
upd = client.put(f"/api/combos/{cid}", json={**body, "name": "组合A", "hold_count": 15})
|
||||
assert upd.status_code == 200
|
||||
assert upd.json()["id"] == cid
|
||||
assert upd.json()["created_at"] == created_at
|
||||
assert upd.json()["hold_count"] == 15
|
||||
|
||||
assert client.delete(f"/api/combos/{cid}").status_code == 200
|
||||
assert client.get(f"/api/combos/{cid}").status_code == 404
|
||||
|
||||
def test_duplicate_name_400(self, client: TestClient) -> None:
|
||||
sid = self._make_strategy(client, "S")
|
||||
body = {"name": "重名", "strategy_ids": [sid], "hold_count": 5,
|
||||
"rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"]}
|
||||
assert client.post("/api/combos", json=body).status_code == 200
|
||||
assert client.post("/api/combos", json=body).status_code == 400
|
||||
|
||||
def test_run_missing_strategy_400(self, client: TestClient) -> None:
|
||||
"""引用不存在的策略 → 提交时即 400,而非等后台 Job 才失败。"""
|
||||
body = {"name": "缺策略组合", "strategy_ids": ["STG-NOT-EXIST"], "hold_count": 5,
|
||||
"rebalance_freq": "monthly", "period": ["2024-01-01", "2024-02-01"]}
|
||||
resp = client.post("/api/combos/run", json=body)
|
||||
assert resp.status_code == 400
|
||||
assert "不存在" in resp.json()["detail"]
|
||||
@@ -0,0 +1,195 @@
|
||||
"""组合回测引擎单测(合成数据,确定性,不依赖数据库/因子注册表)。
|
||||
|
||||
验证三件用户确认的语义:
|
||||
1. Borda 秩和打分:两策略排名不同 → 综合排序可手算预测。
|
||||
2. 持仓天数区间 [Tmin, Tmax]:超 Tmax 强制了结;未满 Tmin 即使掉出 TopN 也暂留。
|
||||
3. 调仓时机 daily/weekly/monthly 产生不同的调仓次数。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, timedelta
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
from app.domain.entities.combo import BacktestCombo
|
||||
from app.domain.entities.research import CostSpec
|
||||
from app.quant.combo_engine import HoldingBandRunner, borda_combine
|
||||
|
||||
|
||||
def _business_days(start: date, n: int) -> list[date]:
|
||||
"""生成 n 个连续工作日(跳过周末),用作合成行情索引。"""
|
||||
out: list[date] = []
|
||||
d = start
|
||||
while len(out) < n:
|
||||
if d.weekday() < 5:
|
||||
out.append(d)
|
||||
d += timedelta(days=1)
|
||||
return out
|
||||
|
||||
|
||||
def _flat_close(symbols: list[str], days: list[date], price: float = 100.0) -> pd.DataFrame:
|
||||
"""所有股票恒定价格的面板(收益为 0,便于隔离「选股/调仓」逻辑)。"""
|
||||
idx = pd.to_datetime(days)
|
||||
return pd.DataFrame(price, index=idx, columns=symbols)
|
||||
|
||||
|
||||
# ---------- 1. Borda 秩和 ----------
|
||||
|
||||
|
||||
def test_borda_combine_hand_computed():
|
||||
"""两策略排名不同,综合分 = Σ(1/名次),可手算。"""
|
||||
day = pd.Timestamp("2024-01-02")
|
||||
# 策略1:A > B > C;策略2:C > A > B
|
||||
p1 = pd.DataFrame({"A": [3.0], "B": [2.0], "C": [1.0]}, index=[day])
|
||||
p2 = pd.DataFrame({"A": [2.0], "B": [1.0], "C": [3.0]}, index=[day])
|
||||
combined = borda_combine([p1, p2]).loc[day]
|
||||
# A: 1/1 + 1/2 = 1.5;C: 1/3 + 1/1 = 1.333;B: 1/2 + 1/3 = 0.833
|
||||
assert combined["A"] == pytest.approx(1.5)
|
||||
assert combined["C"] == pytest.approx(1.0 / 3 + 1.0)
|
||||
assert combined["B"] == pytest.approx(1.0 / 2 + 1.0 / 3)
|
||||
order = combined.sort_values(ascending=False).index.tolist()
|
||||
assert order == ["A", "C", "B"] # 并集后统一排序:A、C 进 Top2,B 落选
|
||||
|
||||
|
||||
def test_borda_missing_symbol_contributes_zero():
|
||||
"""某策略面板里没有某股票(NaN)→ 该策略对它贡献 0,但不影响其它策略的贡献。"""
|
||||
day = pd.Timestamp("2024-01-02")
|
||||
p1 = pd.DataFrame({"A": [3.0], "B": [2.0]}, index=[day]) # 策略1 只有 A、B
|
||||
p2 = pd.DataFrame({"A": [1.0], "C": [2.0]}, index=[day]) # 策略2 只有 A、C
|
||||
combined = borda_combine([p1, p2]).loc[day]
|
||||
assert combined["A"] == pytest.approx(1.0 + 1.0 / 2) # 两策略都覆盖 A
|
||||
assert combined["B"] == pytest.approx(1.0 / 2) # 只被策略1 覆盖
|
||||
assert combined["C"] == pytest.approx(1.0) # 只被策略2 覆盖(在其面板里排第 1)
|
||||
|
||||
|
||||
# ---------- 2. 持仓天数区间 ----------
|
||||
|
||||
|
||||
def _make_combo(**overrides) -> BacktestCombo:
|
||||
base = dict(
|
||||
name="t", strategy_ids=["S1"], initial_capital=1_000_000.0,
|
||||
hold_count=1, hold_min_days=0, hold_max_days=None,
|
||||
rebalance_freq="daily", period=(date(2024, 1, 2), date(2024, 1, 31)),
|
||||
)
|
||||
base.update(overrides)
|
||||
return BacktestCombo(**base)
|
||||
|
||||
|
||||
def test_tmax_force_exit_respected():
|
||||
"""恒价 + N=1 + 永远选 A + Tmax=5 + 日频:A 持有超过 5 天即被强制卖出再买回,
|
||||
任何一笔交易的持有天数都不应明显超过 Tmax。"""
|
||||
symbols = ["A", "B", "C"]
|
||||
days = _business_days(date(2024, 1, 2), 30)
|
||||
close = _flat_close(symbols, days)
|
||||
# A 永远最高分 → 永远 Top1
|
||||
score = pd.DataFrame(
|
||||
{"A": [3.0] * len(days), "B": [2.0] * len(days), "C": [1.0] * len(days)},
|
||||
index=pd.to_datetime(days),
|
||||
)
|
||||
combo = _make_combo(hold_count=1, hold_min_days=0, hold_max_days=5, rebalance_freq="daily")
|
||||
runner = HoldingBandRunner(
|
||||
combo=combo, costs=CostSpec(min_commission=0.0), score=score, close=close,
|
||||
)
|
||||
result = runner.run()
|
||||
|
||||
assert result.trades, "应产生交易"
|
||||
# 交易日索引:用引擎同一口径(交易日)验证「任何一笔持仓都不超过 Tmax 个交易日」
|
||||
tday_pos = {pd.Timestamp(d): i for i, d in enumerate(close.index)}
|
||||
def trading_span(a, b):
|
||||
return tday_pos[pd.Timestamp(b)] - tday_pos[pd.Timestamp(a)]
|
||||
for t in result.trades:
|
||||
span = trading_span(t.entry_date, t.exit_date)
|
||||
# 卖出发生在「held > Tmax」的第一个交易日 → 跨度最多 Tmax+1 个交易日
|
||||
assert span <= 5 + 1, f"持仓跨 {span} 个交易日 > Tmax+1,Tmax 安全阀失效:{t}"
|
||||
# 确实反复「卖后再买」—— 证明 Tmax 在强制换手,而不是一直死拿
|
||||
buys = [a for a in result.signal_history if a.signal == "BUY" and a.filled]
|
||||
assert len(buys) >= 4, f"Tmax=5 在 30 个交易日内应触发多次重买,实际仅 {len(buys)} 次"
|
||||
|
||||
|
||||
def test_tmin_protects_against_churn():
|
||||
"""N=1,第 2 天起 B 变成最高分(A 掉出 Top1),但 Tmin=10 → A 在满 10 天前不被卖出。"""
|
||||
symbols = ["A", "B"]
|
||||
days = _business_days(date(2024, 1, 2), 20)
|
||||
close = _flat_close(symbols, days)
|
||||
# 第 0 天 A 最高;第 1 天起 B 最高
|
||||
a_scores = [3.0] + [1.0] * (len(days) - 1)
|
||||
b_scores = [1.0] + [3.0] * (len(days) - 1)
|
||||
score = pd.DataFrame({"A": a_scores, "B": b_scores}, index=pd.to_datetime(days))
|
||||
combo = _make_combo(hold_count=1, hold_min_days=10, hold_max_days=None, rebalance_freq="daily")
|
||||
runner = HoldingBandRunner(
|
||||
combo=combo, costs=CostSpec(min_commission=0.0), score=score, close=close,
|
||||
)
|
||||
result = runner.run()
|
||||
|
||||
# A 应在第 0 天买入
|
||||
a_buys = [a for a in result.signal_history if a.symbol == "A" and a.signal == "BUY" and a.filled]
|
||||
assert a_buys, "A 应在首日买入"
|
||||
a_sells = [t for t in result.trades if t.symbol == "A"]
|
||||
if a_sells:
|
||||
# 若最终卖出,持有天数必须 ≥ Tmin(不能在满 10 天前因掉出 TopN 被卖)
|
||||
for t in a_sells:
|
||||
assert (t.exit_date - t.entry_date).days >= 10, (
|
||||
f"A 仅持 {(t.exit_date - t.entry_date).days} 天就被卖,违反 Tmin=10 保护"
|
||||
)
|
||||
# 关键断言:前 9 个工作日内 A 不应被卖出(Tmin 保护生效)
|
||||
tday_pos = {pd.Timestamp(d): i for i, d in enumerate(close.index)}
|
||||
early_sells = [
|
||||
a for a in result.signal_history
|
||||
if a.symbol == "A" and a.signal == "SELL" and a.filled
|
||||
and tday_pos[pd.Timestamp(a.date)] - tday_pos[pd.Timestamp(days[0])] < 10
|
||||
]
|
||||
assert not early_sells, f"Tmin 保护失效:A 在 10 个交易日内被卖出 {early_sells}"
|
||||
|
||||
|
||||
# ---------- 3. 调仓时机 ----------
|
||||
|
||||
|
||||
def test_rebalance_freq_changes_cadence():
|
||||
"""同一份数据,daily 的调仓日数 > weekly > monthly(用 selection_history 的 distinct 日期数衡量)。"""
|
||||
symbols = ["A", "B", "C"]
|
||||
days = _business_days(date(2024, 1, 2), 60)
|
||||
close = _flat_close(symbols, days)
|
||||
score = pd.DataFrame(
|
||||
{"A": [3.0] * len(days), "B": [2.0] * len(days), "C": [1.0] * len(days)},
|
||||
index=pd.to_datetime(days),
|
||||
)
|
||||
|
||||
def run_with(freq: str) -> int:
|
||||
combo = _make_combo(
|
||||
hold_count=2, rebalance_freq=freq,
|
||||
period=(days[0], days[-1]),
|
||||
)
|
||||
r = HoldingBandRunner(
|
||||
combo=combo, costs=CostSpec(min_commission=0.0), score=score, close=close,
|
||||
).run()
|
||||
return len({p.date for p in r.selection_history})
|
||||
|
||||
daily_n, weekly_n, monthly_n = run_with("daily"), run_with("weekly"), run_with("monthly")
|
||||
assert daily_n > weekly_n > monthly_n, (
|
||||
f"调仓频次应 daily({daily_n}) > weekly({weekly_n}) > monthly({monthly_n})"
|
||||
)
|
||||
|
||||
|
||||
# ---------- 实体校验 ----------
|
||||
|
||||
|
||||
def test_combo_validates_hold_band_and_freq():
|
||||
with pytest.raises(ValueError):
|
||||
BacktestCombo(
|
||||
name="x", strategy_ids=["A"], hold_count=1,
|
||||
hold_min_days=20, hold_max_days=5, # Tmax < Tmin
|
||||
rebalance_freq="monthly", period=(date(2024, 1, 1), date(2024, 2, 1)),
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
BacktestCombo(
|
||||
name="x", strategy_ids=["A"], hold_count=1,
|
||||
rebalance_freq="yearly", # 非法频率
|
||||
period=(date(2024, 1, 1), date(2024, 2, 1)),
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
BacktestCombo(
|
||||
name="x", strategy_ids=["A", "A"], hold_count=1, # 重复策略
|
||||
rebalance_freq="monthly", period=(date(2024, 1, 1), date(2024, 2, 1)),
|
||||
)
|
||||
@@ -0,0 +1,148 @@
|
||||
"""回测组合服务集成测试(合成数据 + 内存 SQLite,不连真库)。
|
||||
|
||||
验证 ComboService.run 端到端:装配行情 → 多策略 Borda → 持仓区间 runner → BacktestResult。
|
||||
覆盖:
|
||||
- 两策略打分合并后选出并集 TopN;
|
||||
- 持仓天数区间 [Tmin, Tmax] 在真实数据装配路径下生效;
|
||||
- 公共配置的成本/复权被采用并写进 config_snapshot(可复现)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
from app.application.services.combo_service import ComboService
|
||||
from app.domain.entities.combo import BacktestCombo, GlobalConfig
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import (
|
||||
DailyBasicModel,
|
||||
StockDailyModel,
|
||||
StockModel,
|
||||
)
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
# 5 只股票,股息率梯度:A 最高 … E 最低
|
||||
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
|
||||
_DV = {"600000.SH": 9.0, "600001.SH": 7.0, "600002.SH": 5.0, "600003.SH": 3.0, "600004.SH": 1.0}
|
||||
_START = date(2024, 1, 1)
|
||||
|
||||
|
||||
def _seed(session: Session) -> None:
|
||||
session.add_all(
|
||||
[
|
||||
StockModel(symbol=s, name=f"股票{s[:6]}", industry="银行", market="主板",
|
||||
area="深圳", list_date=date(2000, 1, 1), status="L")
|
||||
for s in _SYMS
|
||||
]
|
||||
)
|
||||
dates = pd.bdate_range(_START, periods=80)
|
||||
bars, basics = [], []
|
||||
for i, sym in enumerate(_SYMS):
|
||||
price = 10.0 + i
|
||||
for d in dates:
|
||||
price *= 1 + 0.0006 + 0.0002 * 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)),
|
||||
))
|
||||
basics.append(DailyBasicModel(
|
||||
symbol=sym, trade_date=d.date(), source="tushare",
|
||||
close=Decimal(str(price)), dv_ratio=Decimal(str(_DV[sym])),
|
||||
dv_ttm=Decimal(str(_DV[sym])), pe=Decimal("8"), pb=Decimal("1"),
|
||||
total_mv=Decimal("1e11"),
|
||||
))
|
||||
session.add_all(bars + basics)
|
||||
session.commit()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def service(tmp_path) -> ComboService:
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'combo.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
_seed(session)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyDailyBasicRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
yield ComboService(
|
||||
SqlAlchemyStockRepository(session),
|
||||
SqlAlchemyDailyBarRepository(session),
|
||||
basic_repo=SqlAlchemyDailyBasicRepository(session),
|
||||
)
|
||||
session.close()
|
||||
|
||||
|
||||
def _strategy(sid: str, name: str, *, conditions=None) -> SelectionStrategy:
|
||||
return SelectionStrategy(
|
||||
id=sid, name=name,
|
||||
factors=[{"name": "dividend_yield", "weight": 1.0}],
|
||||
conditions=conditions or [],
|
||||
)
|
||||
|
||||
|
||||
def test_combo_run_produces_backtest_result_with_config_snapshot(service: ComboService) -> None:
|
||||
"""两策略(一个带 dv_ratio 条件、一个不带)→ 组合跑出 BacktestResult,
|
||||
且 config_snapshot 固化了当时的成本/复权与策略定义(可复现)。"""
|
||||
combo = BacktestCombo(
|
||||
id="CMB-T1", name="双策略高股息", strategy_ids=["S1", "S2"],
|
||||
initial_capital=1_000_000.0, hold_count=2, hold_min_days=0, hold_max_days=None,
|
||||
rebalance_freq="monthly", period=(date(2024, 1, 2), date(2024, 4, 15)),
|
||||
)
|
||||
strategies = [
|
||||
_strategy("S1", "纯高股息"),
|
||||
_strategy("S2", "高股息+过滤", conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}]),
|
||||
]
|
||||
config = GlobalConfig(commission_rate=0.0003, stamp_tax_rate=0.0005,
|
||||
slippage_rate=0.001, min_commission=5.0, price_adjustment="hfq")
|
||||
|
||||
result = service.run(combo, strategies, config)
|
||||
|
||||
assert result.summary.initial_capital == 1_000_000.0
|
||||
assert result.equity_curve, "应产出净值曲线"
|
||||
assert result.trades or result.positions, "应有成交或持仓"
|
||||
# 可复现快照:含组合参数 + 两策略定义 + 当时成本/复权
|
||||
snap = result.config_snapshot
|
||||
assert snap["combo"]["hold_count"] == 2
|
||||
assert {s["id"] for s in snap["strategies"]} == {"S1", "S2"}
|
||||
assert snap["costs"]["min_commission"] == 5.0
|
||||
assert snap["price_adjustment"] == "hfq"
|
||||
|
||||
|
||||
def test_hold_max_days_limits_holding_in_real_run(service: ComboService) -> None:
|
||||
"""日频 + Tmax=8:任何一笔交易的持有交易日数不超过 Tmax+1。"""
|
||||
combo = BacktestCombo(
|
||||
id="CMB-T2", name="短持", strategy_ids=["S1"],
|
||||
hold_count=1, hold_min_days=0, hold_max_days=8,
|
||||
rebalance_freq="daily", period=(date(2024, 1, 2), date(2024, 4, 15)),
|
||||
)
|
||||
strategies = [_strategy("S1", "纯高股息")]
|
||||
config = GlobalConfig(min_commission=0.0, price_adjustment="none")
|
||||
result = service.run(combo, strategies, config)
|
||||
|
||||
assert result.trades, "日频短持应产生多次换手"
|
||||
# 用结果的 signal_history 重建交易日序列来按交易日计跨度
|
||||
trade_dates = sorted({pd.Timestamp(p.date) for p in result.equity_curve})
|
||||
pos = {d: i for i, d in enumerate(trade_dates)}
|
||||
for t in result.trades:
|
||||
span = pos[pd.Timestamp(t.exit_date)] - pos[pd.Timestamp(t.entry_date)]
|
||||
assert span <= 8 + 1, f"持仓跨 {span} 个交易日 > Tmax+1:{t}"
|
||||
|
||||
|
||||
def test_missing_strategy_raises_clear_error(service: ComboService) -> None:
|
||||
"""组合引用了 S-GONE,但只传入了 S-OTHER → 明确报出缺失的 id(不静默跳过)。"""
|
||||
combo = BacktestCombo(
|
||||
id="CMB-T3", name="缺策略", strategy_ids=["S-GONE"],
|
||||
hold_count=1, rebalance_freq="monthly", period=(date(2024, 1, 2), date(2024, 2, 1)),
|
||||
)
|
||||
other = _strategy("S-OTHER", "别的")
|
||||
with pytest.raises(ValueError, match="S-GONE"):
|
||||
service.run(combo, [other], GlobalConfig())
|
||||
@@ -1,13 +1,15 @@
|
||||
"""M8.3 策略测试:repo CRUD + /api/strategies + expand 为 ResearchSpec。"""
|
||||
"""选股策略(SelectionStrategy)测试:repo CRUD + /api/strategies + 说明生成。
|
||||
|
||||
2026-09 重构后:策略库只存「选股条件组合」(股票池 + 因子 + 条件),
|
||||
不再持有 selection/rebalance/costs/portfolio/区间 —— 那些移到回测组合与公共配置。
|
||||
旧的 `to_research_spec` / `/expand` 已移除。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
|
||||
import pytest
|
||||
from app.api import deps
|
||||
from app.domain.entities.research import SelectionSpec
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import (
|
||||
SqlAlchemyStrategyRepository,
|
||||
@@ -18,12 +20,12 @@ from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
|
||||
def _st() -> StrategyDefinition:
|
||||
return StrategyDefinition(
|
||||
def _st() -> SelectionStrategy:
|
||||
return SelectionStrategy(
|
||||
name="质量成长动量",
|
||||
description="ROE+动量(演示)",
|
||||
factors=[{"name": "momentum_60", "weight": 1.0}],
|
||||
selection=SelectionSpec(top_n=10),
|
||||
conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}],
|
||||
)
|
||||
|
||||
|
||||
@@ -43,7 +45,9 @@ class TestStrategyRepository:
|
||||
session.commit()
|
||||
got = repo.get("STG-T1")
|
||||
assert got is not None and got.name == "质量成长动量"
|
||||
assert got.selection.top_n == 10
|
||||
# 选股策略只保留选股相关字段
|
||||
assert got.factors[0].name == "momentum_60"
|
||||
assert len(got.conditions) == 1 and got.conditions[0].field == "dv_ratio"
|
||||
assert len(repo.list()) == 1
|
||||
assert repo.get_by_name("质量成长动量") is not None
|
||||
assert repo.delete("STG-T1") is True
|
||||
@@ -57,13 +61,16 @@ class TestStrategyRepository:
|
||||
with pytest.raises(ValueError):
|
||||
repo.save(_st().model_copy(update={"id": "STG-B"}))
|
||||
|
||||
def test_expand_to_research_spec(self, session) -> None:
|
||||
st = _st().model_copy(update={"id": "STG-E"})
|
||||
spec = st.to_research_spec((date(2024, 1, 1), date(2024, 6, 1)))
|
||||
assert spec.type == "backtest"
|
||||
assert spec.period == (date(2024, 1, 1), date(2024, 6, 1))
|
||||
assert spec.factors[0].name == "momentum_60"
|
||||
assert spec.price_adjustment == "none"
|
||||
def test_no_backtest_params_in_entity(self) -> None:
|
||||
"""选股策略实体不应再有回测执行参数字段(重构的核心约束)。"""
|
||||
st = _st()
|
||||
dumped = st.model_dump()
|
||||
for forbidden in (
|
||||
"selection", "rebalance", "costs", "portfolio",
|
||||
"initial_capital", "period", "price_adjustment",
|
||||
"selection_interval_months", "rebalance_interval_months",
|
||||
):
|
||||
assert forbidden not in dumped, f"选股策略不应含回测参数字段 {forbidden}"
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
@@ -83,41 +90,43 @@ def client(tmp_path):
|
||||
|
||||
|
||||
class TestStrategiesApi:
|
||||
def test_crud_and_expand(self, client) -> None:
|
||||
def test_crud_and_describe(self, client) -> None:
|
||||
body = {
|
||||
"name": "演示策略",
|
||||
"description": "动量",
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
"selection": {"top_n": 10},
|
||||
"conditions": [{"field": "dv_ratio", "op": "lte", "value": 30}],
|
||||
}
|
||||
created = client.post("/api/strategies", json=body)
|
||||
assert created.status_code == 200
|
||||
sid = created.json()["id"]
|
||||
assert sid.startswith("STG-")
|
||||
|
||||
assert len(client.get("/api/strategies").json()) == 1
|
||||
# 回读不含回测参数字段
|
||||
detail = client.get(f"/api/strategies/{sid}").json()
|
||||
assert detail["name"] == "演示策略"
|
||||
assert "selection" not in detail and "costs" not in detail
|
||||
|
||||
resp = client.post(
|
||||
f"/api/strategies/{sid}/expand",
|
||||
json={"period": ["2024-01-01", "2024-06-01"]},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
spec = resp.json()
|
||||
assert spec["type"] == "backtest"
|
||||
assert spec["factors"][0]["name"] == "momentum_60"
|
||||
assert len(client.get("/api/strategies").json()) == 1
|
||||
|
||||
# 说明生成:选股策略走专用路径,不假装知道回测参数
|
||||
doc = client.get(f"/api/strategies/{sid}/describe").json()
|
||||
assert "选股策略" in doc["summary"]
|
||||
assert any("回测组合" in w for w in doc["warnings"])
|
||||
|
||||
assert client.delete(f"/api/strategies/{sid}").status_code == 200
|
||||
assert client.get(f"/api/strategies/{sid}").status_code == 404
|
||||
|
||||
def test_duplicate_and_bad_period(self, client) -> None:
|
||||
def test_duplicate_name_400(self, client) -> None:
|
||||
body = {"name": "A", "factors": [{"name": "momentum_60", "weight": 1}]}
|
||||
assert client.post("/api/strategies", json=body).status_code == 200
|
||||
assert client.post("/api/strategies", json=body).status_code == 400
|
||||
sid = client.get("/api/strategies").json()[0]["id"]
|
||||
bad = client.post(
|
||||
|
||||
def test_expand_endpoint_removed(self, client) -> None:
|
||||
"""/expand 已随重构移除(回测改由「回测组合」驱动,不再从单策略展开 ResearchSpec)。"""
|
||||
body = {"name": "B", "factors": [{"name": "momentum_60", "weight": 1}]}
|
||||
sid = client.post("/api/strategies", json=body).json()["id"]
|
||||
resp = client.post(
|
||||
f"/api/strategies/{sid}/expand",
|
||||
json={"period": ["2024-06-01", "2024-01-01"]},
|
||||
json={"period": ["2024-01-01", "2024-06-01"]},
|
||||
)
|
||||
assert bad.status_code == 400
|
||||
assert resp.status_code in (404, 405)
|
||||
|
||||
@@ -24,7 +24,7 @@ from app.domain.entities.research import (
|
||||
SelectionSpec,
|
||||
UniverseSpec,
|
||||
)
|
||||
from app.domain.entities.strategy import StrategyDefinition
|
||||
from app.domain.entities.strategy import SelectionStrategy
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.main import app
|
||||
from app.quant.factors import FactorDef
|
||||
@@ -301,17 +301,19 @@ class TestStepsAndWarnings:
|
||||
|
||||
|
||||
class TestDefinitionInput:
|
||||
def test_strategy_definition_input_uses_placeholder_period(self) -> None:
|
||||
st = StrategyDefinition(
|
||||
def test_selection_strategy_describes_only_selection(self) -> None:
|
||||
"""选股策略说明只讲「怎么选」,不假装知道回测参数(重构后无 period/costs/selection)。"""
|
||||
st = SelectionStrategy(
|
||||
name="演示",
|
||||
factors=[FactorSpec(name="momentum_60", weight=1.0)],
|
||||
selection=SelectionSpec(top_n=10),
|
||||
conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}],
|
||||
)
|
||||
doc = describe_strategy(st)
|
||||
assert doc.summary and doc.formula and doc.steps
|
||||
assert "选股策略" in doc.summary
|
||||
assert "momentum_60" in doc.formula
|
||||
assert "1900" not in doc.formula, "占位区间不得泄漏到展示文本"
|
||||
assert any("无回测区间" in w for w in doc.warnings)
|
||||
# 如实声明回测参数不在策略内
|
||||
assert any("回测组合" in w for w in doc.warnings)
|
||||
|
||||
def test_real_spec_has_no_placeholder_warning(self) -> None:
|
||||
doc = describe_strategy(_spec())
|
||||
@@ -325,13 +327,13 @@ class TestDefinitionInput:
|
||||
describe_strategy(spec)
|
||||
assert spec.model_dump() == before
|
||||
|
||||
st = StrategyDefinition(name="演示", factors=[FactorSpec(name="momentum_60")])
|
||||
st = SelectionStrategy(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"):
|
||||
with pytest.raises(TypeError, match="ResearchSpec 或 SelectionStrategy"):
|
||||
describe_strategy({"factors": []}) # type: ignore[arg-type]
|
||||
|
||||
|
||||
@@ -391,7 +393,7 @@ class TestStrategyDocApi:
|
||||
assert resp.status_code == 200
|
||||
doc = resp.json()
|
||||
assert doc["summary"]
|
||||
assert any("无回测区间" in w for w in doc["warnings"])
|
||||
assert "选股策略" in doc["summary"]
|
||||
assert client.get("/api/strategies/STG-NOT-EXIST/describe").status_code == 404
|
||||
|
||||
def test_post_fills_empty_description(self, client: TestClient) -> None:
|
||||
@@ -428,8 +430,9 @@ class TestStrategyDocApi:
|
||||
def test_auto_description_fits_column_width(self, client: TestClient) -> None:
|
||||
"""自动说明必须落在 `StrategyModel.description = String(300)` 之内。
|
||||
|
||||
超长在 SQLite(测试库)不会报错、到 MySQL 严格模式会 Data too long,
|
||||
因此这里显式断言列宽;截断必须带省略号(显式标记,不静默改短)。
|
||||
重构后选股策略的说明只讲「怎么选」,天然简洁(不再拼回测公式),
|
||||
即使挂满全部因子也远低于列宽 —— 这里断言「一定放得下」即可;
|
||||
截断分支(超长带省略号)由 ResearchSpec 路径保留,选股策略触达不到。
|
||||
"""
|
||||
from app.api.strategies import _DESCRIPTION_MAX_CHARS
|
||||
from app.quant.factors import list_factors
|
||||
@@ -440,10 +443,8 @@ class TestStrategyDocApi:
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
desc = resp.json()["description"]
|
||||
assert len(desc) <= _DESCRIPTION_MAX_CHARS
|
||||
assert desc.endswith("…"), "超长被截断时必须显式带省略号"
|
||||
assert len(desc) <= _DESCRIPTION_MAX_CHARS, "选股策略说明也必须落在列宽内"
|
||||
|
||||
# 常规(单因子)说明远短于列宽:不应被截断
|
||||
normal = client.post(
|
||||
"/api/strategies",
|
||||
json={"name": "单因子策略", "factors": [{"name": "dividend_yield", "weight": 1}]},
|
||||
@@ -460,7 +461,6 @@ class TestStrategyUpdateApi:
|
||||
"name": name,
|
||||
"description": "初始说明",
|
||||
"factors": [{"name": "momentum_60", "weight": 1}],
|
||||
"selection": {"top_n": 10},
|
||||
},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
@@ -478,14 +478,12 @@ class TestStrategyUpdateApi:
|
||||
"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"] == "改后的说明"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user