"""M8.3 策略测试:repo CRUD + /api/strategies + expand 为 ResearchSpec。""" 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.infrastructure.persistence.sqlalchemy.base import Base from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import ( SqlAlchemyStrategyRepository, ) from app.main import app from fastapi.testclient import TestClient from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker def _st() -> StrategyDefinition: return StrategyDefinition( name="质量成长动量", description="ROE+动量(演示)", factors=[{"name": "momentum_60", "weight": 1.0}], selection=SelectionSpec(top_n=10), ) @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.selection.top_n == 10 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_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" @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_expand(self, client) -> None: body = { "name": "演示策略", "description": "动量", "factors": [{"name": "momentum_60", "weight": 1}], "selection": {"top_n": 10}, } 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"] == "演示策略" 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 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: 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( f"/api/strategies/{sid}/expand", json={"period": ["2024-06-01", "2024-01-01"]}, ) assert bad.status_code == 400