diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index d455ee3..8b76268 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -10,15 +10,22 @@ from typing import Annotated from fastapi import Depends from sqlalchemy.orm import Session +from app.application.services.selection_service import SelectionService from app.domain.repositories.jobs import ExperimentRepository, JobRepository from app.domain.repositories.market import ( DailyBarRepository, + FinancialRepository, StockRepository, ) +from app.domain.repositories.selection import SelectionRepository from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( SqlAlchemyDailyBarRepository, + SqlAlchemyFinancialRepository, SqlAlchemyStockRepository, ) +from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import ( + SqlAlchemySelectionRepository, +) from app.infrastructure.persistence.sqlalchemy.session import get_session from app.quant.engine import LocalEngine, QuantEngine from app.quant.service import ResearchService @@ -34,6 +41,10 @@ def _daily_repo_factory(session: DbSession) -> DailyBarRepository: return SqlAlchemyDailyBarRepository(session) +def _financial_repo_factory(session: DbSession) -> FinancialRepository: + return SqlAlchemyFinancialRepository(session) + + def _engine_factory() -> QuantEngine: return LocalEngine() @@ -46,10 +57,25 @@ def _service_factory( return ResearchService(stock_repo, daily_repo, engine) +def _selection_service_factory( + stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)], + daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)], + financial_repo: Annotated[FinancialRepository, Depends(_financial_repo_factory)], +) -> SelectionService: + + return SelectionService(stock_repo, daily_repo, financial_repo) + + +def _selection_repo_factory(session: DbSession) -> SelectionRepository: + return SqlAlchemySelectionRepository(session) + + StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)] DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)] EngineDep = Annotated[QuantEngine, Depends(_engine_factory)] ResearchServiceDep = Annotated[ResearchService, Depends(_service_factory)] +SelectionServiceDep = Annotated[SelectionService, Depends(_selection_service_factory)] +SelectionRepoDep = Annotated[SelectionRepository, Depends(_selection_repo_factory)] def _job_repo_factory(session: DbSession): diff --git a/backend/app/api/router.py b/backend/app/api/router.py index ef7c007..94be253 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -8,13 +8,14 @@ from __future__ import annotations from fastapi import APIRouter -from app.api import agent, experiments, factors, health, jobs, research, stocks +from app.api import agent, 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(research.router) +api_router.include_router(selections.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/selections.py b/backend/app/api/selections.py new file mode 100644 index 0000000..dd4ec9d --- /dev/null +++ b/backend/app/api/selections.py @@ -0,0 +1,71 @@ +"""选股 API(M6.3):提交/查询选股,结果落库可复现。 + +POST /api/selections 同步执行一次选股并落库 → {selection_id, result} +GET /api/selections/{id} 读回某次选股完整结果 +GET /api/selections 历史选股元数据(可过滤 as_of/method) + +同步执行:单日全市场因子评分/条件计算量轻(秒级);未来若超时再迁 Job。 +""" + +from __future__ import annotations + +from datetime import date +from typing import Annotated + +from fastapi import APIRouter, HTTPException, Query +from pydantic import BaseModel + +from app.api.deps import ( + DbSession, + SelectionRepoDep, + SelectionServiceDep, +) +from app.application.services.job_executor import new_id +from app.domain.entities.selection import SelectionMeta, SelectionQuery, SelectionResult + +router = APIRouter(prefix="/selections", tags=["selections"]) + + +class SelectionRun(BaseModel): + selection_id: str + result: SelectionResult + + +@router.post("", response_model=SelectionRun, summary="执行一次选股(同步)并落库") +def run_selection( + query: SelectionQuery, + service: SelectionServiceDep, + selection_repo: SelectionRepoDep, + session: DbSession, +) -> SelectionRun: + result = service.select(query) + selection_id = new_id("SEL") + selection_repo.save(selection_id, result) + session.commit() + return SelectionRun(selection_id=selection_id, result=result) + + +@router.get("/{selection_id}", response_model=SelectionResult, summary="读回一次选股结果") +def get_selection( + selection_id: str, + selection_repo: SelectionRepoDep, +) -> SelectionResult: + result = selection_repo.get(selection_id) + if result is None: + raise HTTPException(status_code=404, detail=f"选股记录 {selection_id} 不存在") + return result + + +_AsOfQuery = Annotated[date | None, Query(description="按选股时点过滤")] +_MethodQuery = Annotated[str | None, Query(pattern="^(score|condition)$")] +_LimitQuery = Annotated[int, Query(ge=1, le=200)] + + +@router.get("", response_model=list[SelectionMeta], summary="历史选股元数据列表") +def list_selections( + selection_repo: SelectionRepoDep, + as_of: _AsOfQuery = None, + method: _MethodQuery = None, + limit: _LimitQuery = 20, +) -> list[SelectionMeta]: + return selection_repo.list_recent(as_of=as_of, method=method, limit=limit) diff --git a/backend/app/core/config.py b/backend/app/core/config.py index 1921907..5b9dea1 100644 --- a/backend/app/core/config.py +++ b/backend/app/core/config.py @@ -10,7 +10,6 @@ import os from dataclasses import dataclass from functools import lru_cache from pathlib import Path -from urllib.parse import quote_plus import yaml @@ -100,6 +99,26 @@ def _normalize_sqlite_url(url: str) -> str: return f"{_SQLITE_PREFIX}{(PROJECT_ROOT / rest).resolve()}" +_URL_RESERVED = set("@:/?#%") + + +def _quote_userinfo(value: str) -> str: + """仅编码会破坏 SQLAlchemy URL 解析的字符(@ : / ? # % 与空白)。 + + 其余字符(含 ! - _ . 等非保留符)原样保留:避免 URL 出现 %XX 干扰 + Alembic configparser 的 interpolation。 + """ + out: list[str] = [] + for ch in value: + if ch in _URL_RESERVED or ch.isspace(): + out.append(f"%{ord(ch):02X}") + elif ord(ch) > 127: # 非 ASCII:按 UTF-8 逐字节 percent 编码 + out.append("".join(f"%{b:02X}" for b in ch.encode("utf-8"))) + else: + out.append(ch) + return "".join(out) + + def _build_mysql_url(mysql: dict | None) -> str | None: """由 config.yaml database.mysql 段组装 mysql+pymysql URL。 @@ -116,7 +135,11 @@ def _build_mysql_url(mysql: dict | None) -> str | None: port = mysql.get("port") or 3306 charset = mysql.get("charset") or "utf8mb4" password = os.environ.get(mysql.get("password_env") or "MYSQL_PASSWORD", "") - auth = f"{quote_plus(user)}:{quote_plus(password)}" if password else quote_plus(user) + auth = ( + f"{_quote_userinfo(user)}:{_quote_userinfo(password)}" + if password + else _quote_userinfo(user) + ) return f"mysql+pymysql://{auth}@{host}:{port}/{db}?charset={charset}" diff --git a/backend/app/domain/entities/selection.py b/backend/app/domain/entities/selection.py index d5c575b..2c3ffc3 100644 --- a/backend/app/domain/entities/selection.py +++ b/backend/app/domain/entities/selection.py @@ -14,7 +14,7 @@ from __future__ import annotations -from datetime import date +from datetime import date, datetime from pydantic import BaseModel, Field, field_validator, model_validator @@ -124,3 +124,13 @@ class SelectionResult(BaseModel): SelectionQuery.model_rebuild() + +class SelectionMeta(BaseModel): + """选股运行元数据(列表/历史查询用,不含候选明细)。""" + + id: str + as_of: date + method: str + universe_size: int = 0 + selected: int = 0 + created_at: datetime | None = None diff --git a/backend/app/domain/repositories/selection.py b/backend/app/domain/repositories/selection.py new file mode 100644 index 0000000..6d3a201 --- /dev/null +++ b/backend/app/domain/repositories/selection.py @@ -0,0 +1,29 @@ +"""选股 Repository Protocol(M6.3 落库)。 + +业务层只依赖本 Protocol;实现位于 infrastructure/persistence。 +结果按 v2 §8:selection_snapshot(一次运行的查询与统计)+ selection_result(逐候选行), +用于回答「2023-08-15 为什么选这只股票」「某日选了什么」。 +""" + +from __future__ import annotations + +from datetime import date +from typing import Protocol + +from app.domain.entities.selection import SelectionMeta, SelectionResult + + +class SelectionRepository(Protocol): + def save(self, selection_id: str, result: SelectionResult) -> None: + """落库一次选股:snapshot 一行 + 候选逐行(同一事务,由调用方 commit)。""" + + def get(self, selection_id: str) -> SelectionResult | None: + """按 id 读回完整结果(重建 SelectionResult)。""" + + def list_recent( + self, + as_of: date | None = None, + method: str | None = None, + limit: int = 20, + ) -> list[SelectionMeta]: + """历史选股元数据列表(按 created_at 倒序;可选 as_of/method 过滤)。""" diff --git a/backend/app/infrastructure/persistence/migrations/versions/a6c91d4e7f20_selection_tables.py b/backend/app/infrastructure/persistence/migrations/versions/a6c91d4e7f20_selection_tables.py new file mode 100644 index 0000000..96d4e25 --- /dev/null +++ b/backend/app/infrastructure/persistence/migrations/versions/a6c91d4e7f20_selection_tables.py @@ -0,0 +1,62 @@ +"""selection_snapshot / selection_result 表(M6.3 选股落库) + +Revision ID: a6c91d4e7f20 +Revises: d3f6c9a21b04 +Create Date: 2026-09-08 + +""" + +from __future__ import annotations + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +revision: str = "a6c91d4e7f20" +down_revision: str | None = "d3f6c9a21b04" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "selection_snapshot", + sa.Column("id", sa.String(length=32), nullable=False), + sa.Column("as_of", sa.Date(), nullable=False), + sa.Column("method", sa.String(length=16), nullable=False), + sa.Column("query_json", sa.Text(), nullable=False), + sa.Column("statistics_json", sa.Text(), nullable=False), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("ix_selection_snapshot_as_of", "selection_snapshot", ["as_of"]) + op.create_index("ix_selection_snapshot_created_at", "selection_snapshot", ["created_at"]) + + op.create_table( + "selection_result", + sa.Column( + "id", + sa.BigInteger().with_variant(sa.Integer(), "sqlite"), + autoincrement=True, + nullable=False, + ), + sa.Column("selection_id", sa.String(length=32), nullable=False), + sa.Column("symbol", sa.String(length=12), nullable=False), + sa.Column("rank", sa.Integer(), nullable=False), + sa.Column("score", sa.Numeric(precision=14, scale=6), nullable=False), + sa.Column("factor_values_json", sa.Text(), nullable=True), + sa.Column("filter_status_json", sa.Text(), nullable=True), + sa.Column("reason_json", sa.Text(), nullable=True), + sa.Column("created_at", sa.DateTime(), nullable=False), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("ix_selection_result_selection_id", "selection_result", ["selection_id"]) + + +def downgrade() -> None: + op.drop_index("ix_selection_result_selection_id", table_name="selection_result") + op.drop_table("selection_result") + op.drop_index("ix_selection_snapshot_created_at", table_name="selection_snapshot") + op.drop_index("ix_selection_snapshot_as_of", table_name="selection_snapshot") + op.drop_table("selection_snapshot") diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py index 43e3410..9504790 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py @@ -16,3 +16,7 @@ from app.infrastructure.persistence.sqlalchemy.models.market import ( # noqa: F SyncLogModel, TradingCalendarModel, ) +from app.infrastructure.persistence.sqlalchemy.models.selection import ( # noqa: F401 + SelectionResultModel, + SelectionSnapshotModel, +) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/selection.py b/backend/app/infrastructure/persistence/sqlalchemy/models/selection.py new file mode 100644 index 0000000..2eeb38d --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/selection.py @@ -0,0 +1,42 @@ +"""选股持久化表(M6.3)。 + +selection_snapshot:一次选股运行的查询与统计快照(复现/历史查询用) +selection_result:逐候选行(symbol/rank/score + 可解释字段 JSON), +回答 v2 §8「某日为什么选这只股票」。 +""" + +from __future__ import annotations + +from datetime import date, datetime + +from sqlalchemy import BigInteger, Date, DateTime, Integer, Numeric, String, Text +from sqlalchemy.orm import Mapped, mapped_column + +from app.infrastructure.persistence.sqlalchemy.base import Base + +PK_INT = BigInteger().with_variant(Integer, "sqlite") + + +class SelectionSnapshotModel(Base): + __tablename__ = "selection_snapshot" + + id: Mapped[str] = mapped_column(String(32), primary_key=True) + as_of: Mapped[date] = mapped_column(Date, index=True) + method: Mapped[str] = mapped_column(String(16)) + query_json: Mapped[str] = mapped_column(Text) + statistics_json: Mapped[str] = mapped_column(Text) + created_at: Mapped[datetime] = mapped_column(DateTime, index=True) + + +class SelectionResultModel(Base): + __tablename__ = "selection_result" + + id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True) + selection_id: Mapped[str] = mapped_column(String(32), index=True) + symbol: Mapped[str] = mapped_column(String(12)) + rank: Mapped[int] = mapped_column(Integer) + score: Mapped[float] = mapped_column(Numeric(14, 6)) + factor_values_json: Mapped[str | None] = mapped_column(Text, nullable=True) + filter_status_json: Mapped[str | None] = mapped_column(Text, nullable=True) + reason_json: Mapped[str | None] = mapped_column(Text, nullable=True) + created_at: Mapped[datetime] = mapped_column(DateTime) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/selection_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/selection_impl.py new file mode 100644 index 0000000..3f23265 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/selection_impl.py @@ -0,0 +1,112 @@ +"""选股 Repository 的 SQLAlchemy 实现(M6.3)。 + +save:snapshot + 候选逐行(同 session,由调用方 commit); +get:读回并重建 SelectionResult;list_recent:历史元数据。 +""" + +from __future__ import annotations + +import json +from datetime import date, datetime + +from sqlalchemy import select +from sqlalchemy.orm import Session + +from app.domain.entities.selection import ( + SelectionCandidate, + SelectionMeta, + SelectionResult, + SelectionStatistics, +) +from app.infrastructure.persistence.sqlalchemy.models.selection import ( + SelectionResultModel, + SelectionSnapshotModel, +) + + +class SqlAlchemySelectionRepository: + def __init__(self, session: Session) -> None: + self._session = session + + def save(self, selection_id: str, result: SelectionResult) -> None: + self._session.add( + SelectionSnapshotModel( + id=selection_id, + as_of=result.as_of_date, + method=result.method, + query_json=json.dumps(result.config_snapshot, ensure_ascii=False), + statistics_json=json.dumps(result.statistics.model_dump(mode="json")), + created_at=datetime.now(), + ) + ) + now = datetime.now() + for c in result.candidates: + self._session.add( + SelectionResultModel( + selection_id=selection_id, + symbol=c.symbol, + rank=c.rank, + score=c.score, + factor_values_json=json.dumps(c.factor_values, ensure_ascii=False), + filter_status_json=json.dumps(c.filter_status, ensure_ascii=False), + reason_json=json.dumps(c.selection_reason, ensure_ascii=False), + created_at=now, + ) + ) + self._session.flush() + + def get(self, selection_id: str) -> SelectionResult | None: + snap = self._session.get(SelectionSnapshotModel, selection_id) + if snap is None: + return None + rows = self._session.scalars( + select(SelectionResultModel) + .where(SelectionResultModel.selection_id == selection_id) + .order_by(SelectionResultModel.rank) + ).all() + stats = SelectionStatistics.model_validate_json(snap.statistics_json) + candidates = [ + SelectionCandidate( + symbol=r.symbol, + rank=r.rank, + score=float(r.score), + factor_values=json.loads(r.factor_values_json or "{}"), + filter_status=json.loads(r.filter_status_json or "[]"), + selection_reason=json.loads(r.reason_json or "[]"), + ) + for r in rows + ] + return SelectionResult( + as_of_date=snap.as_of, + method=snap.method, + statistics=stats, + candidates=candidates, + config_snapshot=json.loads(snap.query_json), + ) + + def list_recent( + self, + as_of: date | None = None, + method: str | None = None, + limit: int = 20, + ) -> list[SelectionMeta]: + stmt = select(SelectionSnapshotModel).order_by(SelectionSnapshotModel.created_at.desc()) + if as_of is not None: + stmt = stmt.where(SelectionSnapshotModel.as_of == as_of) + if method is not None: + stmt = stmt.where(SelectionSnapshotModel.method == method) + stmt = stmt.limit(limit) + metas: list[SelectionMeta] = [] + for snap in self._session.scalars(stmt).all(): + stats = SelectionStatistics.model_validate_json(snap.statistics_json) + metas.append( + SelectionMeta( + id=snap.id, + as_of=snap.as_of, + method=snap.method, + universe_size=stats.universe_size, + selected=stats.selected, + created_at=snap.created_at, + ) + ) + return metas diff --git a/backend/tests/test_selections_api.py b/backend/tests/test_selections_api.py new file mode 100644 index 0000000..5d6c350 --- /dev/null +++ b/backend/tests/test_selections_api.py @@ -0,0 +1,130 @@ +"""M6.3 选股 API 集成测试:POST /api/selections 落库 → GET 读回 → 历史列表。 + +使用 tmp SQLite + 真实 SQLAlchemy Repository(override get_session), +验证「提交→落库→读回一致」闭环与 v2 §8 历史选股查询。 +""" + +from __future__ import annotations + +from datetime import date + +import pytest +from app.api import deps +from app.domain.entities.market import Stock +from app.infrastructure.persistence.sqlalchemy.base import Base +from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( + SqlAlchemyDailyBarRepository, + SqlAlchemyStockRepository, +) +from app.main import app +from fastapi.testclient import TestClient +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker + +from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily + +_SYMS = ["60000" + str(i) + ".SH" for i in range(5)] # 600000~600004 +_AS_OF = date(2024, 12, 31) + + +@pytest.fixture() +def client(tmp_path) -> TestClient: + engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True) + Base.metadata.create_all(engine) + SessionFactory = sessionmaker(bind=engine, expire_on_commit=False) + + drifts = {s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)} + daily_df = synthetic_daily(drifts, n=320) + + with SessionFactory() as session: + stocks = [ + Stock(symbol=s, name=f"测试股份{i}", industry="白酒", + list_date=date(1999, 1, 1)) + for i, s in enumerate(_SYMS) + ] + SqlAlchemyStockRepository(session).upsert_many(stocks) + bars = bars_dataframe_to_daily_bars(daily_df) + SqlAlchemyDailyBarRepository(session).upsert_many(bars) + session.commit() + + def _session_override(): + with SessionFactory() as session: + yield session + + app.dependency_overrides[deps.get_session] = _session_override + with TestClient(app) as c: + yield c + app.dependency_overrides.clear() + + +_SCORE_BODY = { + "universe": {"min_listing_days": 0}, + "method": "score", + "factors": [{"name": "momentum_60", "weight": 1.0}], + "top_n": 3, + "as_of": _AS_OF.isoformat(), +} + + +class TestSelectionsApi: + def test_submit_then_read_back(self, client: TestClient) -> None: + resp = client.post("/api/selections", json=_SCORE_BODY) + assert resp.status_code == 200 + body = resp.json() + sel_id = body["selection_id"] + assert sel_id.startswith("SEL-") + result = body["result"] + assert len(result["candidates"]) == 3 + assert [c["rank"] for c in result["candidates"]] == [1, 2, 3] + assert result["candidates"][0]["factor_values"] # 有因子值 + assert result["candidates"][0]["selection_reason"] # 可解释 + assert result["as_of_date"] <= _AS_OF.isoformat() + + # 读回一致 + got = client.get(f"/api/selections/{sel_id}") + assert got.status_code == 200 + g = got.json() + assert g["as_of_date"] == result["as_of_date"] + assert [c["symbol"] for c in g["candidates"]] == [ + c["symbol"] for c in result["candidates"] + ] + assert g["statistics"]["selected"] == 3 + + def test_missing_id_404(self, client: TestClient) -> None: + assert client.get("/api/selections/SEL-NOPE").status_code == 404 + + def test_list_and_filter(self, client: TestClient) -> None: + client.post("/api/selections", json=_SCORE_BODY) + client.post( + "/api/selections", + json={**_SCORE_BODY, "top_n": 2, "method": "score", + "factors": [{"name": "momentum_20", "weight": 1.0}]}, + ) + rows = client.get("/api/selections").json() + assert len(rows) >= 2 + assert all(r["id"].startswith("SEL-") for r in rows) + assert all(r["method"] == "score" for r in rows) + assert all(r["selected"] > 0 for r in rows) + # as_of 过滤 + by_date = client.get(f"/api/selections?as_of={_AS_OF.isoformat()}").json() + assert len(by_date) == len(rows) + empty = client.get("/api/selections?as_of=2020-01-01").json() + assert empty == [] + + def test_condition_submit(self, client: TestClient) -> None: + resp = client.post( + "/api/selections", + json={ + "universe": {"min_listing_days": 0}, + "method": "condition", + "conditions": [ + {"field": "static.industry", "op": "eq", "value": "白酒"}, + {"field": "momentum_60", "op": "gt", "value": 0}, + ], + "as_of": _AS_OF.isoformat(), + }, + ) + assert resp.status_code == 200 + body = resp.json() + assert len(body["result"]["candidates"]) == 5 # 全部上涨 → 5 只全过 + assert all(c["filter_status"] for c in body["result"]["candidates"])