diff --git a/backend/app/api/jobs.py b/backend/app/api/jobs.py index 96c5580..fb02b20 100644 --- a/backend/app/api/jobs.py +++ b/backend/app/api/jobs.py @@ -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) diff --git a/backend/app/api/selections.py b/backend/app/api/selections.py index dd4ec9d..91a597a 100644 --- a/backend/app/api/selections.py +++ b/backend/app/api/selections.py @@ -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, diff --git a/backend/app/application/services/job_executor.py b/backend/app/application/services/job_executor.py index 844984f..21db646 100644 --- a/backend/app/application/services/job_executor.py +++ b/backend/app/application/services/job_executor.py @@ -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(), diff --git a/backend/tests/test_replays.py b/backend/tests/test_replays.py index 223b365..96bb27d 100644 --- a/backend/tests/test_replays.py +++ b/backend/tests/test_replays.py @@ -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 # 后段持续暴涨 diff --git a/backend/tests/test_selection_job.py b/backend/tests/test_selection_job.py new file mode 100644 index 0000000..0f128a8 --- /dev/null +++ b/backend/tests/test_selection_job.py @@ -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) diff --git a/frontend/web/app/selection/page.tsx b/frontend/web/app/selection/page.tsx index 6eacb98..bfd5190 100644 --- a/frontend/web/app/selection/page.tsx +++ b/frontend/web/app/selection/page.tsx @@ -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(null); const [selectionId, setSelectionId] = useState(""); const [running, setRunning] = useState(false); + const [asyncBusy, setAsyncBusy] = useState(false); const [error, setError] = useState(""); const [history, setHistory] = useState(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(job_id); + if (out.status === "success" && out.result) { + setSelectionId(job_id); + setResult(out.result); + apiGet("/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() { )}
- + {running ? "筛选中…" : "执行选股"} + + {asyncBusy ? "后台执行中…" : "异步(全市场)"} +
{error ? (