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:
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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}"
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 过滤)。"""
|
||||
+62
@@ -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,
|
||||
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
|
||||
@@ -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"])
|
||||
Reference in New Issue
Block a user