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:
+27
-2
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user