diff --git a/backend/app/api/composites.py b/backend/app/api/composites.py new file mode 100644 index 0000000..53820db --- /dev/null +++ b/backend/app/api/composites.py @@ -0,0 +1,69 @@ +"""因子组合 API(M7.2b):/api/composites CRUD。 + +组合 = 可复用因子权重集;计算时方向取因子注册表,落库冗余快照。 +""" + +from __future__ import annotations + +from fastapi import APIRouter, HTTPException + +from app.api.deps import CompositeRepoDep, DbSession +from app.application.services.job_executor import new_id +from app.domain.entities.composite import CompositeComponent, CompositeDefinition +from app.quant.factors import FactorError, get_factor + +router = APIRouter(prefix="/composites", tags=["composites"]) + + +def _fill_direction(definition: CompositeDefinition) -> CompositeDefinition: + """以因子注册表元数据补齐/校正组件 direction(登记但不可计算的因子报错)。""" + out = [] + for c in definition.components: + try: + defn, _fn = get_factor(c.name) + except FactorError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + out.append(CompositeComponent(name=c.name, weight=c.weight, direction=defn.direction)) + return definition.model_copy(update={"components": out}) + + +@router.post("", summary="保存因子组合", response_model=CompositeDefinition) +def create_composite( + definition: CompositeDefinition, + composite_repo: CompositeRepoDep, + session: DbSession, +) -> CompositeDefinition: + prepared = _fill_direction(definition.model_copy(update={"id": ""})) + try: + saved = composite_repo.save( + prepared.model_copy(update={"id": new_id("CF")}) + ) + session.commit() + return saved + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + +@router.get("", summary="因子组合列表", response_model=list[CompositeDefinition]) +def list_composites(composite_repo: CompositeRepoDep) -> list[CompositeDefinition]: + return composite_repo.list() + + +@router.get("/{composite_id}", summary="读取因子组合", response_model=CompositeDefinition) +def get_composite(composite_id: str, composite_repo: CompositeRepoDep) -> CompositeDefinition: + row = composite_repo.get(composite_id) + if row is None: + raise HTTPException(status_code=404, detail=f"组合 {composite_id} 不存在") + return row + + +@router.delete("/{composite_id}", summary="删除因子组合") +def delete_composite( + composite_id: str, + composite_repo: CompositeRepoDep, + session: DbSession, +) -> dict: + if not composite_repo.delete(composite_id): + raise HTTPException(status_code=404, detail=f"组合 {composite_id} 不存在") + session.commit() + return {"deleted": composite_id} diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index fd52565..bb75a53 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -11,6 +11,7 @@ from fastapi import Depends from sqlalchemy.orm import Session from app.application.services.selection_service import SelectionService +from app.domain.repositories.composite import CompositeRepository from app.domain.repositories.factor import FactorRepository from app.domain.repositories.jobs import ExperimentRepository, JobRepository from app.domain.repositories.market import ( @@ -19,6 +20,9 @@ from app.domain.repositories.market import ( StockRepository, ) from app.domain.repositories.selection import SelectionRepository +from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import ( + SqlAlchemyCompositeRepository, +) from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import ( SqlAlchemyFactorRepository, ) @@ -78,6 +82,10 @@ def _factor_repo_factory(session: DbSession) -> FactorRepository: return SqlAlchemyFactorRepository(session) +def _composite_repo_factory(session: DbSession) -> CompositeRepository: + return SqlAlchemyCompositeRepository(session) + + StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)] DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)] EngineDep = Annotated[QuantEngine, Depends(_engine_factory)] @@ -85,6 +93,7 @@ ResearchServiceDep = Annotated[ResearchService, Depends(_service_factory)] SelectionServiceDep = Annotated[SelectionService, Depends(_selection_service_factory)] SelectionRepoDep = Annotated[SelectionRepository, Depends(_selection_repo_factory)] FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)] +CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)] def _job_repo_factory(session: DbSession): diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 94be253..81e4e74 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -8,12 +8,23 @@ from __future__ import annotations from fastapi import APIRouter -from app.api import agent, experiments, factors, health, jobs, research, selections, stocks +from app.api import ( + agent, + composites, + experiments, + factors, + health, + jobs, + research, + selections, + stocks, +) api_router = APIRouter() api_router.include_router(health.router) api_router.include_router(stocks.router) api_router.include_router(factors.router) +api_router.include_router(composites.router) api_router.include_router(research.router) api_router.include_router(selections.router) api_router.include_router(jobs.router) diff --git a/backend/app/domain/entities/composite.py b/backend/app/domain/entities/composite.py new file mode 100644 index 0000000..0aeb950 --- /dev/null +++ b/backend/app/domain/entities/composite.py @@ -0,0 +1,33 @@ +"""因子组合(Composite Factor)领域实体(M7.2b,v2 §13)。 + +组合 = 一组 {因子, 权重} + method(MVP fixed:截面 zscore×方向×权重求和; +方向在计算时取因子注册表元数据,落库时冗余快照以便列表展示)。 +""" + +from __future__ import annotations + +from datetime import datetime + +from pydantic import BaseModel, Field, model_validator + + +class CompositeComponent(BaseModel): + name: str + weight: float = Field(default=1.0, gt=0) + direction: str = Field(default="higher_is_better") + + +class CompositeDefinition(BaseModel): + id: str = "" + name: str = Field(min_length=1, max_length=64) + method: str = Field(default="fixed", pattern="^(fixed)$") + description: str = "" + components: list[CompositeComponent] = Field(min_length=1) + created_at: datetime | None = None + + @model_validator(mode="after") + def _no_duplicate(self) -> CompositeDefinition: + names = [c.name for c in self.components] + if len(set(names)) != len(names): + raise ValueError("components 存在重复因子名") + return self diff --git a/backend/app/domain/repositories/composite.py b/backend/app/domain/repositories/composite.py new file mode 100644 index 0000000..fbc3911 --- /dev/null +++ b/backend/app/domain/repositories/composite.py @@ -0,0 +1,19 @@ +"""Composite Factor Repository Protocol(M7.2b)。""" + +from __future__ import annotations + +from typing import Protocol + +from app.domain.entities.composite import CompositeDefinition + + +class CompositeRepository(Protocol): + def save(self, definition: CompositeDefinition) -> CompositeDefinition: + """新建组合(name 冲突抛 ValueError)。""" + + def get(self, composite_id: str) -> CompositeDefinition | None: ... + + def list(self) -> list[CompositeDefinition]: ... + + def delete(self, composite_id: str) -> bool: + """删除;不存在返回 False。""" diff --git a/backend/app/infrastructure/persistence/migrations/versions/c3e9a0d1f4b5_factor_composite_table.py b/backend/app/infrastructure/persistence/migrations/versions/c3e9a0d1f4b5_factor_composite_table.py new file mode 100644 index 0000000..9ae82c5 --- /dev/null +++ b/backend/app/infrastructure/persistence/migrations/versions/c3e9a0d1f4b5_factor_composite_table.py @@ -0,0 +1,37 @@ +"""factor_composite 表(M7.2b 因子组合保存/复用) + +Revision ID: c3e9a0d1f4b5 +Revises: b7f2a5e81c33 +Create Date: 2026-09-09 + +""" + +from __future__ import annotations + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "c3e9a0d1f4b5" +down_revision: str | None = "b7f2a5e81c33" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "factor_composite", + sa.Column("id", sa.String(length=32), nullable=False), + sa.Column("name", sa.String(length=64), nullable=False), + sa.Column("method", sa.String(length=16), nullable=False), + sa.Column("description", sa.String(length=300), nullable=False), + sa.Column("components_json", sa.Text(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("name", name="uq_factor_composite_name"), + ) + + +def downgrade() -> None: + op.drop_table("factor_composite") diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py index 345db4a..5311611 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py @@ -4,6 +4,9 @@ 模型统一继承 infra.persistence.sqlalchemy.base.Base。 """ +from app.infrastructure.persistence.sqlalchemy.models.composite import ( # noqa: F401 + FactorCompositeModel, +) from app.infrastructure.persistence.sqlalchemy.models.factor import ( # noqa: F401 FactorDefinitionModel, ) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/composite.py b/backend/app/infrastructure/persistence/sqlalchemy/models/composite.py new file mode 100644 index 0000000..f1d88fb --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/composite.py @@ -0,0 +1,25 @@ +"""因子组合表(M7.2b)。 + +factor_composite:可保存/复用的因子组合定义;components 以 JSON 存 +(组件方向冗余快照自因子注册表,计算时以注册表为准)。 +""" + +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 FactorCompositeModel(Base): + __tablename__ = "factor_composite" + + id: Mapped[str] = mapped_column(String(32), primary_key=True) + name: Mapped[str] = mapped_column(String(64), unique=True) + method: Mapped[str] = mapped_column(String(16), default="fixed") + description: Mapped[str] = mapped_column(String(300), default="") + components_json: Mapped[str] = mapped_column(Text) + created_at: Mapped[datetime] = mapped_column(DateTime) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/composite_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/composite_impl.py new file mode 100644 index 0000000..31ba0a8 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/composite_impl.py @@ -0,0 +1,80 @@ +"""因子组合 Repository 的 SQLAlchemy 实现(M7.2b)。""" + +from __future__ import annotations + +import json +from datetime import datetime + +from sqlalchemy import select +from sqlalchemy.orm import Session + +from app.domain.entities.composite import CompositeComponent, CompositeDefinition +from app.infrastructure.persistence.sqlalchemy.models.composite import FactorCompositeModel + + +class SqlAlchemyCompositeRepository: + def __init__(self, session: Session) -> None: + self._session = session + + def save(self, definition: CompositeDefinition) -> CompositeDefinition: + if not definition.id: + raise ValueError("需要 id(由调用方生成)") + exists = self._session.get(FactorCompositeModel, definition.id) + dup = self._session.scalar( + select(FactorCompositeModel) + .where(FactorCompositeModel.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() + if exists is None: + self._session.add( + FactorCompositeModel( + id=definition.id, + name=definition.name, + method=definition.method, + description=definition.description, + components_json=json.dumps( + [c.model_dump() for c in definition.components], ensure_ascii=False + ), + created_at=now, + ) + ) + else: + exists.name = definition.name + exists.method = definition.method + exists.description = definition.description + exists.components_json = json.dumps( + [c.model_dump() for c in definition.components], ensure_ascii=False + ) + self._session.flush() + return definition + + def get(self, composite_id: str) -> CompositeDefinition | None: + row = self._session.get(FactorCompositeModel, composite_id) + return _to_entity(row) if row else None + + def list(self) -> list[CompositeDefinition]: + rows = self._session.scalars( + select(FactorCompositeModel).order_by(FactorCompositeModel.name) + ).all() + return [_to_entity(r) for r in rows] + + def delete(self, composite_id: str) -> bool: + row = self._session.get(FactorCompositeModel, composite_id) + if row is None: + return False + self._session.delete(row) + return True + + +def _to_entity(row: FactorCompositeModel) -> CompositeDefinition: + return CompositeDefinition( + id=row.id, + name=row.name, + method=row.method, + description=row.description, + components=[CompositeComponent(**c) for c in json.loads(row.components_json)], + created_at=row.created_at, + ) diff --git a/backend/tests/test_composites_api.py b/backend/tests/test_composites_api.py new file mode 100644 index 0000000..06a04d8 --- /dev/null +++ b/backend/tests/test_composites_api.py @@ -0,0 +1,125 @@ +"""M7.2b 因子组合测试:repo CRUD(幂等/去重)+ /api/composites(含方向填充/404/删除)。""" + +from __future__ import annotations + +import pytest +from app.api import deps +from app.domain.entities.composite import CompositeComponent, CompositeDefinition +from app.infrastructure.persistence.sqlalchemy.base import Base +from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import ( + SqlAlchemyCompositeRepository, +) +from app.main import app +from fastapi.testclient import TestClient +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker + + +@pytest.fixture() +def session(tmp_path): + engine = create_engine(f"sqlite:///{tmp_path / 'c.db'}", future=True) + Base.metadata.create_all(engine) + Session = sessionmaker(bind=engine, expire_on_commit=False) + with Session() as s: + yield s + + +def _comp(name="质量动量", **kw) -> CompositeDefinition: + base = dict( + name=name, + description="测试组合", + components=[ + CompositeComponent(name="momentum_60", weight=0.7, direction="higher_is_better"), + CompositeComponent(name="volatility_60", weight=0.3, direction="lower_is_better"), + ], + ) + base.update(kw) + return CompositeDefinition(**base) + + +class TestCompositeRepository: + def test_save_get_list(self, session) -> None: + repo = SqlAlchemyCompositeRepository(session) + saved = repo.save(_comp().model_copy(update={"id": "CF-TEST-1"})) + session.commit() + assert saved.id == "CF-TEST-1" + got = repo.get("CF-TEST-1") + assert got is not None and got.name == "质量动量" + assert len(got.components) == 2 + assert repo.list()[0].components[0].direction == "higher_is_better" + + def test_duplicate_name_rejected(self, session) -> None: + repo = SqlAlchemyCompositeRepository(session) + repo.save(_comp().model_copy(update={"id": "CF-A"})) + session.commit() + with pytest.raises(ValueError): + repo.save(_comp().model_copy(update={"id": "CF-B"})) # 同名不同 id + + def test_delete(self, session) -> None: + repo = SqlAlchemyCompositeRepository(session) + repo.save(_comp().model_copy(update={"id": "CF-X"})) + session.commit() + assert repo.delete("CF-X") is True + session.commit() + assert repo.get("CF-X") is None + assert repo.delete("CF-X") is False + + +@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 TestCompositesApi: + def test_crud(self, client) -> None: + body = { + "name": "质量动量组合", + "description": "动量+低波", + "components": [ + {"name": "momentum_60", "weight": 0.7}, + {"name": "volatility_60", "weight": 0.3}, + ], + } + created = client.post("/api/composites", json=body) + assert created.status_code == 200 + cid = created.json()["id"] + assert cid.startswith("CF-") + # 方向由注册表自动填充 + dirs = {c["name"]: c["direction"] for c in created.json()["components"]} + assert dirs == {"momentum_60": "higher_is_better", "volatility_60": "lower_is_better"} + + rows = client.get("/api/composites").json() + assert len(rows) == 1 + detail = client.get(f"/api/composites/{cid}").json() + assert detail["name"] == "质量动量组合" + + assert client.delete(f"/api/composites/{cid}").status_code == 200 + assert client.get(f"/api/composites/{cid}").status_code == 404 + assert client.delete(f"/api/composites/{cid}").status_code == 404 + + def test_unknown_factor_rejected(self, client) -> None: + resp = client.post( + "/api/composites", + json={ + "name": "坏组合", + "components": [{"name": "no_such_factor", "weight": 1}], + }, + ) + assert resp.status_code == 400 + assert "no_such_factor" in resp.json()["detail"] + + def test_duplicate_name_400(self, client) -> None: + body = {"name": "同名", "components": [{"name": "momentum_60", "weight": 1}]} + assert client.post("/api/composites", json=body).status_code == 200 + assert client.post("/api/composites", json=body).status_code == 400