feat(job): D2 全市场选股 Job 化(kind=selection 异步)
- Job kind=selection:executor 分支用 SelectionService(spec=SelectionQuery),结果 落 result_json、自动归档 Experiment(summary:as_of + 选出 N);/api/jobs 与 /api/experiments 的 result 解码支持 SelectionResult - POST /api/selections/jobs:提交 SelectionQuery 为异步 Job(BackgroundTasks/subprocess) - Web 选股页新增「异步(全市场)」按钮:提交 Job → waitJob 轮询结果(解决同步 60s+) - tests/test_selection_job.py:executor 执行归档(SelectionResult/Experiment kind)、 API 提交→轮询→结果与实验列表;全量 pytest + tsc 通过
This commit is contained in:
@@ -36,8 +36,12 @@ router = APIRouter(prefix="/jobs", tags=["jobs"])
|
||||
|
||||
|
||||
def _decode_result(kind: str, result_json: str | None):
|
||||
from app.domain.entities.selection import SelectionResult
|
||||
|
||||
if result_json is None:
|
||||
return None
|
||||
if kind == "selection":
|
||||
return SelectionResult.model_validate_json(result_json)
|
||||
model = BacktestResult if kind == "backtest" else FactorTestReport
|
||||
return model.model_validate_json(result_json)
|
||||
|
||||
|
||||
@@ -9,18 +9,20 @@ GET /api/selections 历史选股元数据(可过滤 as_of/method)
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from datetime import date, datetime
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Query
|
||||
from fastapi import APIRouter, BackgroundTasks, HTTPException, Query
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.api.deps import (
|
||||
DbSession,
|
||||
JobRepoDep,
|
||||
SelectionRepoDep,
|
||||
SelectionServiceDep,
|
||||
)
|
||||
from app.application.services.job_executor import new_id
|
||||
from app.application.services.job_executor import new_id, run_job_background
|
||||
from app.domain.entities.research import JobRecord, JobStatus
|
||||
from app.domain.entities.selection import SelectionMeta, SelectionQuery, SelectionResult
|
||||
|
||||
router = APIRouter(prefix="/selections", tags=["selections"])
|
||||
@@ -31,6 +33,26 @@ class SelectionRun(BaseModel):
|
||||
result: SelectionResult
|
||||
|
||||
|
||||
@router.post("/jobs", summary="提交全市场/长任务选股为异步 Job")
|
||||
def submit_selection_job(
|
||||
query: SelectionQuery,
|
||||
background: BackgroundTasks,
|
||||
session: DbSession,
|
||||
job_repo: JobRepoDep,
|
||||
) -> dict:
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"),
|
||||
kind="selection",
|
||||
spec_json=query.model_dump_json(),
|
||||
status=JobStatus.QUEUED,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
job_repo.create(job)
|
||||
session.commit()
|
||||
background.add_task(run_job_background, job.id)
|
||||
return {"job_id": job.id, "status": job.status}
|
||||
|
||||
|
||||
@router.post("", response_model=SelectionRun, summary="执行一次选股(同步)并落库")
|
||||
def run_selection(
|
||||
query: SelectionQuery,
|
||||
|
||||
@@ -61,12 +61,16 @@ def _git_short_rev() -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _summary_text(kind: str, result: BacktestResult | FactorTestReport) -> str | None:
|
||||
def _summary_text(kind: str, result) -> str | None:
|
||||
from app.domain.entities.selection import SelectionResult
|
||||
|
||||
if kind == "backtest" and isinstance(result, BacktestResult):
|
||||
s = result.summary
|
||||
return f"总收益 {s.total_return_pct:.2f}% · 年化 {s.annual_return_pct:.2f}% · 回撤 {s.max_drawdown_pct:.2f}%"
|
||||
if isinstance(result, FactorTestReport):
|
||||
return f"IC {result.ic_mean:.4f} · RankIC {result.rank_ic_mean:.4f} · 样本 {result.sample_days} 日"
|
||||
if kind == "selection" and isinstance(result, SelectionResult):
|
||||
return f"as_of {result.as_of_date} · 选出 {result.statistics.selected} / 评估 {result.statistics.evaluated}"
|
||||
return None
|
||||
|
||||
|
||||
@@ -91,10 +95,20 @@ def _execute_inner(
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
try:
|
||||
spec = ResearchSpec.model_validate_json(job.spec_json)
|
||||
service = ResearchService(
|
||||
stock_repo_factory(session), daily_repo_factory(session), engine
|
||||
)
|
||||
is_selection = job.kind == "selection"
|
||||
if is_selection:
|
||||
from app.application.services.selection_service import SelectionService
|
||||
from app.domain.entities.selection import SelectionQuery
|
||||
|
||||
spec = SelectionQuery.model_validate_json(job.spec_json)
|
||||
service = SelectionService(
|
||||
stock_repo_factory(session), daily_repo_factory(session)
|
||||
)
|
||||
else:
|
||||
spec = ResearchSpec.model_validate_json(job.spec_json)
|
||||
service = ResearchService(
|
||||
stock_repo_factory(session), daily_repo_factory(session), engine
|
||||
)
|
||||
|
||||
def _set_stage(name: str) -> None:
|
||||
"""阶段上报(v3 §23):独立短会话写 job.stage 并 commit(子进程同样走 DB)。"""
|
||||
@@ -108,7 +122,10 @@ def _execute_inner(
|
||||
except Exception: # noqa: BLE001 —— 阶段上报失败不阻断执行
|
||||
pass
|
||||
|
||||
if spec.type == "backtest":
|
||||
if is_selection:
|
||||
_set_stage("selection")
|
||||
result = service.select(spec)
|
||||
elif spec.type == "backtest":
|
||||
result = service.run_backtest(spec, on_stage=_set_stage)
|
||||
else:
|
||||
result = service.run_factor_test(spec, on_stage=_set_stage)
|
||||
@@ -116,10 +133,10 @@ def _execute_inner(
|
||||
result_json = json.dumps(result.model_dump(mode="json"), ensure_ascii=False)
|
||||
experiment = ExperimentRecord(
|
||||
id=new_id("EXP"),
|
||||
kind=spec.type,
|
||||
kind=job.kind,
|
||||
spec_json=job.spec_json,
|
||||
result_json=result_json,
|
||||
summary_text=_summary_text(spec.type, result),
|
||||
summary_text=_summary_text(job.kind, result),
|
||||
code_version=_git_short_rev(),
|
||||
job_id=job.id,
|
||||
created_at=datetime.now(),
|
||||
|
||||
@@ -106,7 +106,7 @@ class TestReplayConsistency:
|
||||
dates = sorted(df["trade_date"].unique())
|
||||
mid = dates[len(dates) // 2]
|
||||
b = _SYMS[1]
|
||||
for i, d in enumerate(dates):
|
||||
for d in dates:
|
||||
if d > mid:
|
||||
mask = (df["symbol"] == b) & (df["trade_date"] == d)
|
||||
df.loc[mask, "close"] = df.loc[mask, "close"] * 1.03 # 后段持续暴涨
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
"""D2 选股 Job 化测试:kind=selection 异步执行(executor + API 提交/轮询)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
import pytest
|
||||
from app.api import deps
|
||||
from app.application.services import job_executor as je
|
||||
from app.domain.entities.market import Stock
|
||||
from app.domain.entities.research import JobRecord, JobStatus
|
||||
from app.domain.entities.selection import SelectionQuery, SelectionResult
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
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 = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
|
||||
|
||||
|
||||
def _query(symbols=_SYMS) -> SelectionQuery:
|
||||
return SelectionQuery(
|
||||
universe={"exclude_st": False, "min_listing_days": 0, "symbols": list(symbols)},
|
||||
factors=[{"name": "momentum_60", "weight": 1.0}],
|
||||
top_n=3,
|
||||
as_of=date(2024, 12, 31),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def engine_session(tmp_path):
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'db.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
||||
df = synthetic_daily({s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}, n=320)
|
||||
with Session() as session:
|
||||
SqlAlchemyStockRepository(session).upsert_many(
|
||||
[Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS)]
|
||||
)
|
||||
SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df))
|
||||
session.commit()
|
||||
return Session
|
||||
|
||||
|
||||
class TestSelectionJobExecutor:
|
||||
def test_execute_and_archive(self, engine_session) -> None:
|
||||
Session = engine_session
|
||||
with Session() as session:
|
||||
SqlAlchemyJobRepository(session).create(
|
||||
JobRecord(id="JOB-SEL-1", kind="selection",
|
||||
spec_json=_query().model_dump_json(),
|
||||
status=JobStatus.QUEUED, created_at=datetime.now())
|
||||
)
|
||||
session.commit()
|
||||
je.execute_job(
|
||||
"JOB-SEL-1",
|
||||
session_factory=Session,
|
||||
job_repo_factory=lambda s: SqlAlchemyJobRepository(s),
|
||||
experiment_repo_factory=lambda s: SqlAlchemyExperimentRepository(s),
|
||||
stock_repo_factory=lambda s: SqlAlchemyStockRepository(s),
|
||||
daily_repo_factory=lambda s: SqlAlchemyDailyBarRepository(s),
|
||||
engine=None,
|
||||
)
|
||||
with Session() as session:
|
||||
job = SqlAlchemyJobRepository(session).get("JOB-SEL-1")
|
||||
exp = SqlAlchemyExperimentRepository(session).get(job.experiment_id or "")
|
||||
assert job.status == JobStatus.SUCCESS
|
||||
result = SelectionResult.model_validate_json(job.result_json or "{}")
|
||||
assert len(result.candidates) == 3
|
||||
assert exp is not None and exp.kind == "selection"
|
||||
assert "as_of" in (exp.summary_text or "") and "选出 3" in (exp.summary_text or "")
|
||||
|
||||
|
||||
class TestSelectionJobApi:
|
||||
@pytest.fixture()
|
||||
def client(self, tmp_path, monkeypatch):
|
||||
from app.infrastructure.persistence.sqlalchemy import session as sess_mod
|
||||
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine, expire_on_commit=False)
|
||||
monkeypatch.setattr(sess_mod, "SessionLocal", Session)
|
||||
df = synthetic_daily({s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}, n=320)
|
||||
with Session() as session:
|
||||
SqlAlchemyStockRepository(session).upsert_many(
|
||||
[Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1))
|
||||
for i, s in enumerate(_SYMS)]
|
||||
)
|
||||
SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df))
|
||||
session.commit()
|
||||
|
||||
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()
|
||||
|
||||
def test_submit_poll(self, client) -> None:
|
||||
body = _query().model_dump(mode="json")
|
||||
resp = client.post("/api/selections/jobs", json=body)
|
||||
assert resp.status_code == 200
|
||||
job_id = resp.json()["job_id"]
|
||||
state = None
|
||||
for _ in range(40):
|
||||
state = client.get(f"/api/jobs/{job_id}").json()
|
||||
if state["status"] in ("success", "failed", "cancelled"):
|
||||
break
|
||||
assert state["status"] == "success"
|
||||
result = state["result"]
|
||||
assert result is not None and len(result["candidates"]) == 3
|
||||
assert result["statistics"]["selected"] == 3
|
||||
# 实验归档存在(selection 类型)
|
||||
exps = client.get("/api/experiments").json()
|
||||
assert any(e["kind"] == "selection" for e in exps)
|
||||
@@ -6,6 +6,7 @@
|
||||
*/
|
||||
import { useEffect, useState } from "react";
|
||||
import { apiGet, apiPost } from "@/lib/api";
|
||||
import { waitJob } from "@/lib/jobs";
|
||||
import type {
|
||||
FactorMeta,
|
||||
SelectionCondition,
|
||||
@@ -53,6 +54,7 @@ export default function SelectionPage() {
|
||||
const [result, setResult] = useState<SelectionResult | null>(null);
|
||||
const [selectionId, setSelectionId] = useState("");
|
||||
const [running, setRunning] = useState(false);
|
||||
const [asyncBusy, setAsyncBusy] = useState(false);
|
||||
const [error, setError] = useState("");
|
||||
const [history, setHistory] = useState<SelectionMeta[] | null>(null);
|
||||
|
||||
@@ -111,6 +113,28 @@ export default function SelectionPage() {
|
||||
}
|
||||
}
|
||||
|
||||
/** 异步(全市场等长任务):提交 Job 后台执行并轮询(解决同步 60s+)。 */
|
||||
async function runAsync() {
|
||||
setAsyncBusy(true);
|
||||
setError("");
|
||||
setResult(null);
|
||||
try {
|
||||
const { job_id } = await apiPost<{ job_id: string }>("/selections/jobs", buildQuery());
|
||||
const out = await waitJob<SelectionResult>(job_id);
|
||||
if (out.status === "success" && out.result) {
|
||||
setSelectionId(job_id);
|
||||
setResult(out.result);
|
||||
apiGet<SelectionMeta[]>("/selections?limit=8").then(setHistory).catch(() => undefined);
|
||||
} else {
|
||||
setError(`任务${out.status}${out.error ? `:${out.error}` : ""}`);
|
||||
}
|
||||
} catch (e) {
|
||||
setError((e as Error).message);
|
||||
} finally {
|
||||
setAsyncBusy(false);
|
||||
}
|
||||
}
|
||||
|
||||
async function openHistory(id: string) {
|
||||
setError("");
|
||||
try {
|
||||
@@ -279,9 +303,12 @@ export default function SelectionPage() {
|
||||
)}
|
||||
|
||||
<div style={{ marginTop: 14 }}>
|
||||
<Btn variant="primary" icon="play" loading={running} disabled={running} onClick={run}>
|
||||
<Btn variant="primary" icon="play" loading={running} disabled={running || asyncBusy} onClick={run}>
|
||||
{running ? "筛选中…" : "执行选股"}
|
||||
</Btn>
|
||||
<Btn variant="ghost" icon="clock" loading={asyncBusy} disabled={running || asyncBusy} onClick={runAsync}>
|
||||
{asyncBusy ? "后台执行中…" : "异步(全市场)"}
|
||||
</Btn>
|
||||
</div>
|
||||
{error ? (
|
||||
<div style={{ marginTop: 12 }}>
|
||||
|
||||
Reference in New Issue
Block a user