"""选股策略(SelectionStrategy)测试:repo CRUD + /api/strategies + 说明生成。 2026-09 重构后:策略库只存「选股条件组合」(股票池 + 因子 + 条件), 不再持有 selection/rebalance/costs/portfolio/区间 —— 那些移到回测组合与公共配置。 旧的 `to_research_spec` / `/expand` 已移除。 """ from __future__ import annotations import json from datetime import datetime import pytest from app.api import deps from app.domain.entities.strategy import SelectionStrategy from app.infrastructure.persistence.sqlalchemy.base import Base from app.infrastructure.persistence.sqlalchemy.models.strategy import StrategyModel from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import ( SqlAlchemyStrategyRepository, ) from app.main import app from fastapi.testclient import TestClient from pydantic import ValidationError from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker def _st() -> SelectionStrategy: return SelectionStrategy( name="质量成长动量", description="ROE+动量(演示)", factors=[{"name": "momentum_60", "weight": 1.0}], conditions=[{"field": "dv_ratio", "op": "lte", "value": 30}], ) @pytest.fixture() def session(tmp_path): engine = create_engine(f"sqlite:///{tmp_path / 'st.db'}", future=True) Base.metadata.create_all(engine) Session = sessionmaker(bind=engine, expire_on_commit=False) with Session() as s: yield s class TestStrategyRepository: def test_save_get_list_delete(self, session) -> None: repo = SqlAlchemyStrategyRepository(session) repo.save(_st().model_copy(update={"id": "STG-T1"})) session.commit() got = repo.get("STG-T1") assert got is not None and got.name == "质量成长动量" # 选股策略只保留选股相关字段 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 session.commit() assert repo.get("STG-T1") is None def test_duplicate_name(self, session) -> None: repo = SqlAlchemyStrategyRepository(session) repo.save(_st().model_copy(update={"id": "STG-A"})) session.commit() with pytest.raises(ValueError): repo.save(_st().model_copy(update={"id": "STG-B"})) 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}" def test_rejects_legacy_backtest_params(self) -> None: """混入旧版回测参数必须**报错**,不能静默丢弃(否则调用方以为设上了)。""" with pytest.raises(ValidationError) as exc: SelectionStrategy( name="带旧参数的策略", factors=[{"name": "momentum_60", "weight": 1}], costs={"commission_rate": 0.001}, rebalance="monthly", initial_capital=500_000, ) # 三个未知键都应被点名(便于调用方知道该搬去哪里) for key in ("costs", "rebalance", "initial_capital"): assert key in str(exc.value) def test_spec_type_only_selection(self) -> None: """spec_type 取值域收敛为 selection(本实体只表示选股策略)。""" with pytest.raises(ValidationError): SelectionStrategy( name="旧类型", spec_type="backtest", factors=[{"name": "momentum_60", "weight": 1}], ) assert _st().spec_type == "selection" def test_legacy_rows_with_extra_keys_still_readable(self, session) -> None: """历史行 config_json 残留旧键时仍能读出(仓储读出前剔除),forbid 不影响兼容。""" session.add( StrategyModel( id="STG-LEGACY", name="历史策略", description="历史说明", spec_type="backtest", config_json=json.dumps({ "universe": {"exclude_st": True}, "factors": [{"name": "dividend_yield", "weight": 1}], "conditions": [], "selection": {"top_n": 20}, "rebalance": "monthly", "costs": {"commission_rate": 0.0003}, }, ensure_ascii=False), version="1", created_at=datetime(2026, 1, 1), ) ) session.commit() got = SqlAlchemyStrategyRepository(session).get("STG-LEGACY") assert got is not None assert got.factors[0].name == "dividend_yield" assert got.spec_type == "selection" # 旧列值不参与实体(回落到默认) @pytest.fixture() def client(tmp_path): engine = create_engine(f"sqlite:///{tmp_path / '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 TestStrategiesApi: def test_crud_and_describe(self, client) -> None: body = { "name": "演示策略", "description": "动量", "factors": [{"name": "momentum_60", "weight": 1}], "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-") # 回读不含回测参数字段 detail = client.get(f"/api/strategies/{sid}").json() assert detail["name"] == "演示策略" assert "selection" not in detail and "costs" not in detail 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_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 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-01-01", "2024-06-01"]}, ) assert resp.status_code in (404, 405) def test_create_with_legacy_params_422(self, client) -> None: """旧调用方把回测参数塞进策略 → 422 并点名未知字段(而非 200 静默丢弃)。""" body = { "name": "旧调用方", "factors": [{"name": "momentum_60", "weight": 1}], "rebalance": "monthly", "costs": {"commission_rate": 0.001}, } resp = client.post("/api/strategies", json=body) assert resp.status_code == 422 assert "costs" in resp.text and "rebalance" in resp.text # 未落库:策略列表仍为空 assert client.get("/api/strategies").json() == []