feat(selection): M6.3 选股结果落库 + /api/selections(快照 + 逐候选行)

- domain/repositories/selection.py + SelectionMeta:选股 Repository Protocol
- models/selection.py + migration a6c91d4e7f20:selection_snapshot(查询/统计快照)+
  selection_result(逐候选:rank/score/factor_values/filter_status/reason JSON)
- repositories/selection_impl.py:save/get/list_recent(读回重建 SelectionResult)
- api/selections.py:POST /api/selections(同步执行+落库,返回 selection_id+result)、
  GET /{id}、GET 列表(as_of/method 过滤);deps 装配 SelectionService/Repo
- config._build_mysql_url:仅编码破坏 URL 结构的字符(修 Alembic configparser 遇 %21 崩)
- 迁移已在 MySQL 应用(alembic head a6c91d4e7f20);tests/test_selections_api.py 4 例
  提交→读回一致/404/列表过滤/condition;全量 pytest 通过
This commit is contained in:
Simon
2026-09-09 00:20:42 +08:00
parent 75c5472c31
commit c60dc78c88
11 changed files with 514 additions and 4 deletions
+26
View File
@@ -10,15 +10,22 @@ from typing import Annotated
from fastapi import Depends from fastapi import Depends
from sqlalchemy.orm import Session 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.jobs import ExperimentRepository, JobRepository
from app.domain.repositories.market import ( from app.domain.repositories.market import (
DailyBarRepository, DailyBarRepository,
FinancialRepository,
StockRepository, StockRepository,
) )
from app.domain.repositories.selection import SelectionRepository
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository, SqlAlchemyDailyBarRepository,
SqlAlchemyFinancialRepository,
SqlAlchemyStockRepository, SqlAlchemyStockRepository,
) )
from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import (
SqlAlchemySelectionRepository,
)
from app.infrastructure.persistence.sqlalchemy.session import get_session from app.infrastructure.persistence.sqlalchemy.session import get_session
from app.quant.engine import LocalEngine, QuantEngine from app.quant.engine import LocalEngine, QuantEngine
from app.quant.service import ResearchService from app.quant.service import ResearchService
@@ -34,6 +41,10 @@ def _daily_repo_factory(session: DbSession) -> DailyBarRepository:
return SqlAlchemyDailyBarRepository(session) return SqlAlchemyDailyBarRepository(session)
def _financial_repo_factory(session: DbSession) -> FinancialRepository:
return SqlAlchemyFinancialRepository(session)
def _engine_factory() -> QuantEngine: def _engine_factory() -> QuantEngine:
return LocalEngine() return LocalEngine()
@@ -46,10 +57,25 @@ def _service_factory(
return ResearchService(stock_repo, daily_repo, engine) 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)] StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)] DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)] EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
ResearchServiceDep = Annotated[ResearchService, Depends(_service_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): def _job_repo_factory(session: DbSession):
+2 -1
View File
@@ -8,13 +8,14 @@ from __future__ import annotations
from fastapi import APIRouter 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 = APIRouter()
api_router.include_router(health.router) api_router.include_router(health.router)
api_router.include_router(stocks.router) api_router.include_router(stocks.router)
api_router.include_router(factors.router) api_router.include_router(factors.router)
api_router.include_router(research.router) api_router.include_router(research.router)
api_router.include_router(selections.router)
api_router.include_router(jobs.router) api_router.include_router(jobs.router)
api_router.include_router(experiments.router) api_router.include_router(experiments.router)
api_router.include_router(agent.router) api_router.include_router(agent.router)
+71
View File
@@ -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)
+25 -2
View File
@@ -10,7 +10,6 @@ import os
from dataclasses import dataclass from dataclasses import dataclass
from functools import lru_cache from functools import lru_cache
from pathlib import Path from pathlib import Path
from urllib.parse import quote_plus
import yaml import yaml
@@ -100,6 +99,26 @@ def _normalize_sqlite_url(url: str) -> str:
return f"{_SQLITE_PREFIX}{(PROJECT_ROOT / rest).resolve()}" 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: def _build_mysql_url(mysql: dict | None) -> str | None:
"""由 config.yaml database.mysql 段组装 mysql+pymysql URL。 """由 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 port = mysql.get("port") or 3306
charset = mysql.get("charset") or "utf8mb4" charset = mysql.get("charset") or "utf8mb4"
password = os.environ.get(mysql.get("password_env") or "MYSQL_PASSWORD", "") 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}" return f"mysql+pymysql://{auth}@{host}:{port}/{db}?charset={charset}"
+11 -1
View File
@@ -14,7 +14,7 @@
from __future__ import annotations from __future__ import annotations
from datetime import date from datetime import date, datetime
from pydantic import BaseModel, Field, field_validator, model_validator from pydantic import BaseModel, Field, field_validator, model_validator
@@ -124,3 +124,13 @@ class SelectionResult(BaseModel):
SelectionQuery.model_rebuild() 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
@@ -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 过滤)。"""
@@ -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")
@@ -16,3 +16,7 @@ from app.infrastructure.persistence.sqlalchemy.models.market import ( # noqa: F
SyncLogModel, SyncLogModel,
TradingCalendarModel, TradingCalendarModel,
) )
from app.infrastructure.persistence.sqlalchemy.models.selection import ( # noqa: F401
SelectionResultModel,
SelectionSnapshotModel,
)
@@ -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)
@@ -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
+130
View File
@@ -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"])