From 40bd603b443c9d1995091abbb9235b64c7eadc35 Mon Sep 17 00:00:00 2001 From: Simon Date: Wed, 30 Sep 2026 21:43:28 +0800 Subject: [PATCH] =?UTF-8?q?feat(backend):=20=E7=AD=96=E7=95=A5=E5=BA=93?= =?UTF-8?q?=E9=87=8D=E6=9E=84=E4=B8=BA=E3=80=8C=E9=80=89=E8=82=A1=E7=AD=96?= =?UTF-8?q?=E7=95=A5=20+=20=E5=85=AC=E5=85=B1=E9=85=8D=E7=BD=AE=20+=20?= =?UTF-8?q?=E5=9B=9E=E6=B5=8B=E7=BB=84=E5=90=88=E3=80=8D=E4=B8=89=E4=BB=B6?= =?UTF-8?q?=E5=A5=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 按用户目标把原来「一个策略 = 全套参数」拆开(已确认的设计决策): - 公共配置 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)。 --- backend/app/agent/tools_impl.py | 8 +- backend/app/api/combos.py | 140 +++++ backend/app/api/config.py | 30 + backend/app/api/deps.py | 15 + backend/app/api/router.py | 4 + backend/app/api/strategies.py | 51 +- .../app/application/services/combo_service.py | 254 ++++++++ .../app/application/services/job_executor.py | 43 +- backend/app/domain/entities/combo.py | 152 +++++ backend/app/domain/entities/strategy.py | 74 +-- backend/app/domain/repositories/combo.py | 26 + backend/app/domain/repositories/strategy.py | 12 +- .../b4c5d6e7f8a9_combo_and_global_config.py | 113 ++++ .../persistence/sqlalchemy/models/__init__.py | 4 + .../persistence/sqlalchemy/models/combo.py | 47 ++ .../sqlalchemy/repositories/combo_impl.py | 135 ++++ .../sqlalchemy/repositories/strategy_impl.py | 31 +- backend/app/quant/combo_engine.py | 580 ++++++++++++++++++ backend/app/quant/local_engine.py | 11 +- backend/app/quant/strategy_doc.py | 135 +++- backend/tests/test_combo_api.py | 109 ++++ backend/tests/test_combo_engine.py | 195 ++++++ backend/tests/test_combo_service.py | 148 +++++ backend/tests/test_strategies.py | 75 ++- backend/tests/test_strategy_doc.py | 32 +- 25 files changed, 2250 insertions(+), 174 deletions(-) create mode 100644 backend/app/api/combos.py create mode 100644 backend/app/api/config.py create mode 100644 backend/app/application/services/combo_service.py create mode 100644 backend/app/domain/entities/combo.py create mode 100644 backend/app/domain/repositories/combo.py create mode 100644 backend/app/infrastructure/persistence/migrations/versions/b4c5d6e7f8a9_combo_and_global_config.py create mode 100644 backend/app/infrastructure/persistence/sqlalchemy/models/combo.py create mode 100644 backend/app/infrastructure/persistence/sqlalchemy/repositories/combo_impl.py create mode 100644 backend/app/quant/combo_engine.py create mode 100644 backend/tests/test_combo_api.py create mode 100644 backend/tests/test_combo_engine.py create mode 100644 backend/tests/test_combo_service.py diff --git a/backend/app/agent/tools_impl.py b/backend/app/agent/tools_impl.py index 40ca9fc..a6b2c9c 100644 --- a/backend/app/agent/tools_impl.py +++ b/backend/app/agent/tools_impl.py @@ -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 diff --git a/backend/app/api/combos.py b/backend/app/api/combos.py new file mode 100644 index 0000000..4e471d8 --- /dev/null +++ b/backend/app/api/combos.py @@ -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) diff --git a/backend/app/api/config.py b/backend/app/api/config.py new file mode 100644 index 0000000..cd24d74 --- /dev/null +++ b/backend/app/api/config.py @@ -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 diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index b4aa107..03f3faf 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -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): diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 6a7e19e..fd2610d 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -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) diff --git a/backend/app/api/strategies.py b/backend/app/api/strategies.py index c6a2d9a..faf63fa 100644 --- a/backend/app/api/strategies.py +++ b/backend/app/api/strategies.py @@ -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) diff --git a/backend/app/application/services/combo_service.py b/backend/app/application/services/combo_service.py new file mode 100644 index 0000000..a20e4fd --- /dev/null +++ b/backend/app/application/services/combo_service.py @@ -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() diff --git a/backend/app/application/services/job_executor.py b/backend/app/application/services/job_executor.py index 8be010e..2da9c5c 100644 --- a/backend/app/application/services/job_executor.py +++ b/backend/app/application/services/job_executor.py @@ -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, diff --git a/backend/app/domain/entities/combo.py b/backend/app/domain/entities/combo.py new file mode 100644 index 0000000..8e93e9f --- /dev/null +++ b/backend/app/domain/entities/combo.py @@ -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() diff --git a/backend/app/domain/entities/strategy.py b/backend/app/domain/entities/strategy.py index 7cba792..d0f7148 100644 --- a/backend/app/domain/entities/strategy.py +++ b/backend/app/domain/entities/strategy.py @@ -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 diff --git a/backend/app/domain/repositories/combo.py b/backend/app/domain/repositories/combo.py new file mode 100644 index 0000000..9b0f301 --- /dev/null +++ b/backend/app/domain/repositories/combo.py @@ -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: ... diff --git a/backend/app/domain/repositories/strategy.py b/backend/app/domain/repositories/strategy.py index 7571693..b0ac367 100644 --- a/backend/app/domain/repositories/strategy.py +++ b/backend/app/domain/repositories/strategy.py @@ -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: ... diff --git a/backend/app/infrastructure/persistence/migrations/versions/b4c5d6e7f8a9_combo_and_global_config.py b/backend/app/infrastructure/persistence/migrations/versions/b4c5d6e7f8a9_combo_and_global_config.py new file mode 100644 index 0000000..cb99330 --- /dev/null +++ b/backend/app/infrastructure/persistence/migrations/versions/b4c5d6e7f8a9_combo_and_global_config.py @@ -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") diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py index cb2254a..4a44645 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py @@ -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, ) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/combo.py b/backend/app/infrastructure/persistence/sqlalchemy/models/combo.py new file mode 100644 index 0000000..9863600 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/combo.py @@ -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) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/combo_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/combo_impl.py new file mode 100644 index 0000000..a0c782e --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/combo_impl.py @@ -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, + ) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/strategy_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/strategy_impl.py index b832762..1cd5ca2 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/repositories/strategy_impl.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/strategy_impl.py @@ -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, ) diff --git a/backend/app/quant/combo_engine.py b/backend/app/quant/combo_engine.py new file mode 100644 index 0000000..8e4328d --- /dev/null +++ b/backend/app/quant/combo_engine.py @@ -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 diff --git a/backend/app/quant/local_engine.py b/backend/app/quant/local_engine.py index f740911..915cc8d 100644 --- a/backend/app/quant/local_engine.py +++ b/backend/app/quant/local_engine.py @@ -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] diff --git a/backend/app/quant/strategy_doc.py b/backend/app/quant/strategy_doc.py index 32a9b55..a73d60e 100644 --- a/backend/app/quant/strategy_doc.py +++ b/backend/app/quant/strategy_doc.py @@ -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__}" ) diff --git a/backend/tests/test_combo_api.py b/backend/tests/test_combo_api.py new file mode 100644 index 0000000..7d5f352 --- /dev/null +++ b/backend/tests/test_combo_api.py @@ -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"] diff --git a/backend/tests/test_combo_engine.py b/backend/tests/test_combo_engine.py new file mode 100644 index 0000000..6ae0844 --- /dev/null +++ b/backend/tests/test_combo_engine.py @@ -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)), + ) diff --git a/backend/tests/test_combo_service.py b/backend/tests/test_combo_service.py new file mode 100644 index 0000000..bdf501a --- /dev/null +++ b/backend/tests/test_combo_service.py @@ -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()) diff --git a/backend/tests/test_strategies.py b/backend/tests/test_strategies.py index 03effe2..5398c6f 100644 --- a/backend/tests/test_strategies.py +++ b/backend/tests/test_strategies.py @@ -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) diff --git a/backend/tests/test_strategy_doc.py b/backend/tests/test_strategy_doc.py index 800a1ef..2b619b4 100644 --- a/backend/tests/test_strategy_doc.py +++ b/backend/tests/test_strategy_doc.py @@ -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"] == "改后的说明"