diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index d1e8987..df6f264 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -22,6 +22,7 @@ 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.composite_impl import ( SqlAlchemyCompositeRepository, ) @@ -39,6 +40,9 @@ from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl impor from app.infrastructure.persistence.sqlalchemy.repositories.signal_impl import ( SqlAlchemySignalRepository, ) +from app.infrastructure.persistence.sqlalchemy.repositories.strategy_impl import ( + SqlAlchemyStrategyRepository, +) from app.infrastructure.persistence.sqlalchemy.session import get_session from app.quant.engine import LocalEngine, QuantEngine from app.quant.service import ResearchService @@ -102,6 +106,10 @@ def _signal_repo_factory(session: DbSession) -> SignalRepository: return SqlAlchemySignalRepository(session) +def _strategy_repo_factory(session: DbSession) -> StrategyRepository: + return SqlAlchemyStrategyRepository(session) + + StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)] DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)] EngineDep = Annotated[QuantEngine, Depends(_engine_factory)] @@ -112,6 +120,7 @@ FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)] CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)] SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)] SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)] +StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)] def _job_repo_factory(session: DbSession): diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 321c020..23469eb 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -19,6 +19,7 @@ from app.api import ( selections, signals, stocks, + strategies, ) api_router = APIRouter() @@ -29,6 +30,7 @@ api_router.include_router(composites.router) api_router.include_router(research.router) api_router.include_router(selections.router) api_router.include_router(signals.router) +api_router.include_router(strategies.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 new file mode 100644 index 0000000..b6b8310 --- /dev/null +++ b/backend/app/api/strategies.py @@ -0,0 +1,80 @@ +"""策略 API(M8.3):/api/strategies CRUD + 展开为 ResearchSpec。 + +POST /api/strategies 保存策略(name 唯一) +GET /api/strategies 列表 +GET /api/strategies/{id} +DELETE /api/strategies/{id} +POST /api/strategies/{id}/expand body: {period:[start,end], initial_capital?} → ResearchSpec +""" + +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 + +router = APIRouter(prefix="/strategies", tags=["strategies"]) + + +class ExpandRequest(BaseModel): + period: tuple[date, date] + initial_capital: float = Field(default=1_000_000.0, gt=0) + + +@router.post("", response_model=StrategyDefinition, summary="保存策略") +def create_strategy( + definition: StrategyDefinition, + strategy_repo: StrategyRepoDep, + session: DbSession, +) -> StrategyDefinition: + try: + saved = strategy_repo.save(definition.model_copy(update={"id": new_id("STG")})) + session.commit() + return saved + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + +@router.get("", response_model=list[StrategyDefinition], summary="策略列表") +def list_strategies(strategy_repo: StrategyRepoDep) -> list[StrategyDefinition]: + return strategy_repo.list() + + +@router.get("/{strategy_id}", response_model=StrategyDefinition, summary="读取策略") +def get_strategy(strategy_id: str, strategy_repo: StrategyRepoDep) -> StrategyDefinition: + row = strategy_repo.get(strategy_id) + if row is None: + raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在") + return row + + +@router.delete("/{strategy_id}", summary="删除策略") +def delete_strategy( + strategy_id: str, + strategy_repo: StrategyRepoDep, + session: DbSession, +) -> dict: + if not strategy_repo.delete(strategy_id): + raise HTTPException(status_code=404, detail=f"策略 {strategy_id} 不存在") + 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/domain/entities/strategy.py b/backend/app/domain/entities/strategy.py new file mode 100644 index 0000000..728ea69 --- /dev/null +++ b/backend/app/domain/entities/strategy.py @@ -0,0 +1,53 @@ +"""策略领域实体(M8.3,v2 §17/§5.3)。 + +Strategy = 完整策略定义(universe + factors + selection + rebalance + costs + +portfolio,除回测区间 period 外),保存为命名资产;回测时补 period 展开为 +ResearchSpec(v2 §18 Research Specification 为统一契约,策略是其持久化形态)。 +""" + +from __future__ import annotations + +from datetime import date, datetime + +from pydantic import BaseModel, Field + +from app.domain.entities.research import ( + CostSpec, + FactorSpec, + PortfolioSpec, + SelectionSpec, + UniverseSpec, +) + + +class StrategyDefinition(BaseModel): + id: str = "" + name: str = Field(min_length=1, max_length=64) + description: str = "" + spec_type: str = Field(default="backtest", pattern="^(backtest|factor_test)$") + universe: UniverseSpec = UniverseSpec() + price_adjustment: str = Field(default="none", pattern="^(none|qfq)$") + factors: list[FactorSpec] = Field(min_length=1) + selection: SelectionSpec = SelectionSpec() + rebalance: str = Field(default="monthly", pattern="^(weekly|monthly)$") + 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 + + return ResearchSpec( + type=self.spec_type, + universe=self.universe, + price_adjustment=self.price_adjustment, + factors=self.factors, + selection=self.selection, + rebalance=self.rebalance, + period=period, + costs=self.costs, + portfolio=self.portfolio, + initial_capital=initial_capital if initial_capital else 1_000_000.0, + ) diff --git a/backend/app/domain/repositories/strategy.py b/backend/app/domain/repositories/strategy.py new file mode 100644 index 0000000..7571693 --- /dev/null +++ b/backend/app/domain/repositories/strategy.py @@ -0,0 +1,20 @@ +"""策略 Repository Protocol(M8.3)。""" + +from __future__ import annotations + +from typing import Protocol + +from app.domain.entities.strategy import StrategyDefinition + + +class StrategyRepository(Protocol): + def save(self, definition: StrategyDefinition) -> StrategyDefinition: + """新建(name 冲突抛 ValueError)。""" + + def get(self, strategy_id: str) -> StrategyDefinition | None: ... + + def get_by_name(self, name: str) -> StrategyDefinition | None: ... + + def list(self) -> list[StrategyDefinition]: ... + + def delete(self, strategy_id: str) -> bool: ... diff --git a/backend/app/infrastructure/persistence/migrations/versions/e1f2a3b4c5d6_strategy_table.py b/backend/app/infrastructure/persistence/migrations/versions/e1f2a3b4c5d6_strategy_table.py new file mode 100644 index 0000000..3b25f2c --- /dev/null +++ b/backend/app/infrastructure/persistence/migrations/versions/e1f2a3b4c5d6_strategy_table.py @@ -0,0 +1,38 @@ +"""strategy 表(M8.3 策略持久化) + +Revision ID: e1f2a3b4c5d6 +Revises: d8e0b2f3c4d5 +Create Date: 2026-09-09 + +""" + +from __future__ import annotations + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "e1f2a3b4c5d6" +down_revision: str | None = "d8e0b2f3c4d5" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "strategy", + 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), + sa.Column("spec_type", sa.String(length=16), nullable=False), + sa.Column("config_json", sa.Text(), nullable=False), + sa.Column("version", sa.String(length=16), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("name", name="uq_strategy_name"), + ) + + +def downgrade() -> None: + op.drop_table("strategy") diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py index e031465..914cb34 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py @@ -30,3 +30,6 @@ from app.infrastructure.persistence.sqlalchemy.models.signal import ( # noqa: F SignalEventModel, SignalSnapshotModel, ) +from app.infrastructure.persistence.sqlalchemy.models.strategy import ( # noqa: F401 + StrategyModel, +) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/strategy.py b/backend/app/infrastructure/persistence/sqlalchemy/models/strategy.py new file mode 100644 index 0000000..80e23e1 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/strategy.py @@ -0,0 +1,22 @@ +"""strategy 表(M8.3)。""" + +from __future__ import annotations + +from datetime import datetime + +from sqlalchemy import DateTime, String, Text +from sqlalchemy.orm import Mapped, mapped_column + +from app.infrastructure.persistence.sqlalchemy.base import Base + + +class StrategyModel(Base): + __tablename__ = "strategy" + + 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="") + spec_type: Mapped[str] = mapped_column(String(16), default="backtest") + config_json: Mapped[str] = mapped_column(Text) + 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/strategy_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/strategy_impl.py new file mode 100644 index 0000000..f0a98a4 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/strategy_impl.py @@ -0,0 +1,86 @@ +"""策略 Repository 的 SQLAlchemy 实现(M8.3)。 + +config 以 JSON 存(StrategyDefinition.model_dump);读取时重建实体。 +""" + +from __future__ import annotations + +import json +from datetime import datetime + +from sqlalchemy import select +from sqlalchemy.orm import Session + +from app.domain.entities.strategy import StrategyDefinition +from app.infrastructure.persistence.sqlalchemy.models.strategy import StrategyModel + + +class SqlAlchemyStrategyRepository: + def __init__(self, session: Session) -> None: + self._session = session + + def save(self, definition: StrategyDefinition) -> StrategyDefinition: + if not definition.id: + raise ValueError("需要 id(由调用方生成)") + dup = self._session.scalar( + select(StrategyModel).where(StrategyModel.name == definition.name).limit(1) + ) + if dup is not None and dup.id != definition.id: + raise ValueError(f"策略名已存在:{definition.name}") + now = definition.created_at or datetime.now() + row = self._session.get(StrategyModel, definition.id) + config_json = json.dumps(definition.model_dump(exclude={"id", "created_at"}), ensure_ascii=False) + if row is None: + self._session.add( + StrategyModel( + id=definition.id, + name=definition.name, + description=definition.description, + spec_type=definition.spec_type, + config_json=config_json, + version=definition.version, + created_at=now, + ) + ) + else: + row.name = definition.name + row.description = definition.description + row.spec_type = definition.spec_type + row.config_json = config_json + row.version = definition.version + self._session.flush() + return definition + + def get(self, strategy_id: str) -> StrategyDefinition | 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: + 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]: + rows = self._session.scalars( + select(StrategyModel).order_by(StrategyModel.name) + ).all() + return [_to_entity(r) for r in rows] + + def delete(self, strategy_id: str) -> bool: + row = self._session.get(StrategyModel, strategy_id) + if row is None: + return False + self._session.delete(row) + return True + + +def _to_entity(row: StrategyModel) -> StrategyDefinition: + data = json.loads(row.config_json) + # 列字段由 DB 行回填,避免与 config_json 重复 + for key in ("name", "version", "description", "spec_type"): + data.pop(key, None) + return StrategyDefinition( + id=row.id, name=row.name, version=row.version, + created_at=row.created_at, **data, + ) diff --git a/backend/tests/test_strategies.py b/backend/tests/test_strategies.py new file mode 100644 index 0000000..03effe2 --- /dev/null +++ b/backend/tests/test_strategies.py @@ -0,0 +1,123 @@ +"""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