feat(job): D1 Job 阶段上报(v3 §23)+ 列表/取消端点

- ResearchService.run_backtest/run_factor_test 支持 on_stage 回调(data_loading →
  factor_calculation/backtesting → analysis);job_executor 经独立短会话把 stage 写库
  (子进程模式同样走 DB),收尾回读最后阶段避免覆盖
- /api/jobs:GET 列表(kind/limit)、POST /{id}/cancel(queued/running → CANCELLED +
  终止子进程 terminate_active;父进程兜底已跳过 CANCELLED)
- SSE /jobs/{id}/events 现会携带 stage
- tests/test_job_stages.py:成功 job 终态 stage=analysis;executor 取消 queued/幂等/
  已完成不可取消;API 列表+终态不可取消+404;全量 pytest 通过
This commit is contained in:
Simon
2026-09-09 07:37:25 +08:00
parent 67d3aa1349
commit 7c268e43df
4 changed files with 271 additions and 8 deletions
+27 -2
View File
@@ -13,12 +13,13 @@ from __future__ import annotations
import asyncio
import json
from datetime import datetime
from typing import Annotated
from fastapi import APIRouter, BackgroundTasks, HTTPException
from fastapi import APIRouter, BackgroundTasks, HTTPException, Query
from fastapi.responses import StreamingResponse
from app.api.deps import DbSession, JobRepoDep
from app.application.services.job_executor import new_id, run_job_background
from app.application.services.job_executor import new_id, run_job_background, terminate_active
from app.domain.entities.research import (
BacktestResult,
FactorTestReport,
@@ -69,6 +70,30 @@ def create_job(
return {"job_id": job.id, "status": job.status}
@router.get("", summary="Job 列表")
def list_jobs(
job_repo: JobRepoDep,
kind: Annotated[str | None, Query(description="按类型过滤(backtest/factor_test)")] = None,
limit: Annotated[int, Query(ge=1, le=200)] = 20,
) -> list[dict]:
return [_job_view(j) for j in job_repo.list_recent(kind=kind, limit=limit)]
@router.post("/{job_id}/cancel", summary="取消 Job(queued/running)")
def cancel_job(job_id: str, session: DbSession, job_repo: JobRepoDep) -> dict:
job = job_repo.get(job_id)
if job is None:
raise HTTPException(status_code=404, detail=f"Job {job_id} 不存在")
if job.status not in (JobStatus.QUEUED, JobStatus.RUNNING):
return {"job_id": job_id, "status": job.status, "cancelled": False}
job.status = JobStatus.CANCELLED
job.stage = None
job_repo.update(job)
session.commit()
terminate_active(job_id) # 终止研究子进程(若有);父进程兜底已跳过 CANCELLED
return {"job_id": job_id, "status": JobStatus.CANCELLED, "cancelled": True}
@router.get("/{job_id}", summary="查询 Job 状态与结果")
def get_job(job_id: str, job_repo: JobRepoDep) -> dict:
job = job_repo.get(job_id)
@@ -95,10 +95,23 @@ def _execute_inner(
service = ResearchService(
stock_repo_factory(session), daily_repo_factory(session), engine
)
def _set_stage(name: str) -> None:
"""阶段上报(v3 §23):独立短会话写 job.stage 并 commit(子进程同样走 DB)。"""
try:
with session_factory() as st_sess:
st = job_repo_factory(st_sess).get(job_id)
if st is not None and st.status == JobStatus.RUNNING:
st.stage = name
job_repo_factory(st_sess).update(st)
st_sess.commit()
except Exception: # noqa: BLE001 —— 阶段上报失败不阻断执行
pass
if spec.type == "backtest":
result = service.run_backtest(spec)
result = service.run_backtest(spec, on_stage=_set_stage)
else:
result = service.run_factor_test(spec)
result = service.run_factor_test(spec, on_stage=_set_stage)
result_json = json.dumps(result.model_dump(mode="json"), ensure_ascii=False)
experiment = ExperimentRecord(
@@ -122,6 +135,14 @@ def _execute_inner(
job.error = f"{type(exc).__name__}: {exc}"
finally:
job.finished_at = datetime.now()
# 回读最后一次阶段上报(_set_stage 经独立会话写库),避免被本会话覆盖
try:
with session_factory() as last_sess:
last = job_repo_factory(last_sess).get(job_id)
if last is not None:
job.stage = last.stage
except Exception: # noqa: BLE001
pass
job_repo.update(job)
session.commit()
@@ -353,6 +374,42 @@ def _run_in_subprocess(
)
def terminate_active(job_id: str) -> bool:
"""终止该 Job 的活动子进程(如有)。返回是否找到并终止。"""
with _active_lock:
proc = _active_jobs.get(job_id)
if proc is None:
return False
import contextlib
with contextlib.suppress(Exception):
proc.terminate()
return True
def cancel_job(
job_id: str,
*,
session_factory: Callable | None = None,
job_repo_factory: Callable | None = None,
) -> bool:
"""取消 queued/running Job(置 CANCELLED 并终止子进程);不可取消返回 False。"""
facts = default_factories()
sf = session_factory or facts["session_factory"]
jr = job_repo_factory or facts["job_repo_factory"]
with sf() as session:
repo = jr(session)
job = repo.get(job_id)
if job is None or job.status not in (JobStatus.QUEUED, JobStatus.RUNNING):
return False
job.status = JobStatus.CANCELLED
job.stage = None
repo.update(job)
session.commit()
terminate_active(job_id)
return True
def mark_stale_jobs_failed(
*,
session_factory: Callable | None = None,
+21 -4
View File
@@ -120,17 +120,29 @@ class ResearchService:
self._engine = engine
self._index_repo = index_repo
def run_factor_test(self, spec: ResearchSpec, horizon_days: int = 21) -> FactorTestReport:
def run_factor_test(
self, spec: ResearchSpec, horizon_days: int = 21, on_stage=None
) -> FactorTestReport:
"""on_stage(str):执行阶段回调(data_loading / factor_calculation / analysis),
供 Job 状态机上报 stage(v3 §23)。"""
if spec.type != "factor_test":
raise ValueError("factor_test 用例需要 spec.type=factor_test")
_stage(on_stage, "data_loading")
daily = self._load_daily(spec)
return self._engine.run_factor_test(daily, spec, horizon_days=horizon_days)
_stage(on_stage, "factor_calculation")
report = self._engine.run_factor_test(daily, spec, horizon_days=horizon_days)
_stage(on_stage, "analysis")
return report
def run_backtest(self, spec: ResearchSpec) -> BacktestResult:
def run_backtest(self, spec: ResearchSpec, on_stage=None) -> BacktestResult:
if spec.type != "backtest":
raise ValueError("backtest 用例需要 spec.type=backtest")
_stage(on_stage, "data_loading")
daily = self._load_daily(spec)
return self._engine.run_backtest(daily, spec)
_stage(on_stage, "backtesting")
result = self._engine.run_backtest(daily, spec)
_stage(on_stage, "analysis")
return result
def run_factor_correlation(self, spec: ResearchSpec) -> FactorCorrelationReport:
"""多因子两两相关(v3 §12):同 universe/period 装配 → 横截面相关矩阵。"""
@@ -158,3 +170,8 @@ class ResearchService:
sorted(required),
adjust=spec.price_adjustment,
)
def _stage(cb, name: str) -> None:
if cb is not None:
cb(name)
+164
View File
@@ -0,0 +1,164 @@
"""D1 Job 阶段上报 / 取消 / 列表测试。"""
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.research import JobRecord, JobStatus, ResearchSpec
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 app.quant.engine import LocalEngine
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)]
def _spec(**kw) -> ResearchSpec:
base = dict(
type="backtest",
universe={"exclude_st": False, "min_listing_days": 0, "symbols": _SYMS},
factors=[{"name": "momentum_60", "weight": 1.0}],
selection={"top_n": 2},
rebalance="monthly",
period=[date(2024, 5, 1), date(2024, 8, 31)],
)
base.update(kw)
return ResearchSpec(**base)
@pytest.fixture()
def sf(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'job.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(
[
__import__("app.domain.entities.market", fromlist=["Stock"]).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
def _create(sf, job_id, spec=None):
with sf() as session:
SqlAlchemyJobRepository(session).create(
JobRecord(id=job_id, kind="backtest", spec_json=(spec or _spec()).model_dump_json(),
status=JobStatus.QUEUED, created_at=datetime.now())
)
session.commit()
class TestJobStages:
def test_stage_reported_and_success(self, sf) -> None:
job_id = "JOB-STAGE-1"
_create(sf, job_id)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository as D,
)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyStockRepository as S,
)
je.execute_job(
job_id,
session_factory=sf,
job_repo_factory=lambda s: SqlAlchemyJobRepository(s),
experiment_repo_factory=lambda s: SqlAlchemyExperimentRepository(s),
stock_repo_factory=lambda s: S(s),
daily_repo_factory=lambda s: D(s),
engine=LocalEngine(),
)
with sf() as session:
job = SqlAlchemyJobRepository(session).get(job_id)
assert job is not None and job.status == JobStatus.SUCCESS
# 阶段在成功时保留最后 analysis(v3 §23 阶段名)
assert job.stage == "analysis"
def test_cancel_queued_and_idempotent(self, sf) -> None:
job_id = "JOB-CANCEL-1"
_create(sf, job_id)
assert je.cancel_job(job_id, session_factory=sf,
job_repo_factory=lambda s: SqlAlchemyJobRepository(s)) is True
with sf() as session:
job = SqlAlchemyJobRepository(session).get(job_id)
assert job is not None and job.status == JobStatus.CANCELLED
# 已取消不可再次取消
assert je.cancel_job(job_id, session_factory=sf,
job_repo_factory=lambda s: SqlAlchemyJobRepository(s)) is False
def test_cancel_completed_false(self, sf) -> None:
job_id = "JOB-DONE"
_create(sf, job_id)
with sf() as session:
repo = SqlAlchemyJobRepository(session)
job = repo.get(job_id)
job.status = JobStatus.SUCCESS
repo.update(job)
session.commit()
assert je.cancel_job(job_id, session_factory=sf,
job_repo_factory=lambda s: SqlAlchemyJobRepository(s)) is False
@pytest.fixture()
def client(tmp_path, monkeypatch):
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
# 后台执行(JOB_MODE=local)走全局 SessionLocal → 指向同一测试库
from app.infrastructure.persistence.sqlalchemy import session as sess_mod
monkeypatch.setattr(sess_mod, "SessionLocal", Session)
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()
class TestJobsApi:
def test_list_and_cancel_finished(self, client) -> None:
# 提交一个会失败(未知因子)的 job:后台同步执行完 → failed
payload = _spec(factors=[{"name": "no_such", "weight": 1.0}]).model_dump(mode="json")
resp = client.post("/api/jobs", json=payload)
assert resp.status_code == 200
job_id = resp.json()["job_id"]
# 等待后台任务终态(TestClient 同步执行后台任务于响应后)
state = None
for _ in range(30):
state = client.get(f"/api/jobs/{job_id}").json()
if state["status"] in ("success", "failed", "cancelled"):
break
assert state["status"] == "failed" # 未知因子 → 失败
# 列表含该 job
rows = client.get("/api/jobs").json()
assert any(r["id"] == job_id for r in rows)
# 已完成(失败)不可取消
cancel = client.post(f"/api/jobs/{job_id}/cancel").json()
assert cancel["cancelled"] is False
assert client.post("/api/jobs/JOB-NOPE/cancel").status_code == 404