diff --git a/backend/app/api/jobs.py b/backend/app/api/jobs.py index 12c86da..96c5580 100644 --- a/backend/app/api/jobs.py +++ b/backend/app/api/jobs.py @@ -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) diff --git a/backend/app/application/services/job_executor.py b/backend/app/application/services/job_executor.py index 0fa2965..844984f 100644 --- a/backend/app/application/services/job_executor.py +++ b/backend/app/application/services/job_executor.py @@ -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, diff --git a/backend/app/quant/service.py b/backend/app/quant/service.py index 3007dbf..c9fcea4 100644 --- a/backend/app/quant/service.py +++ b/backend/app/quant/service.py @@ -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) diff --git a/backend/tests/test_job_stages.py b/backend/tests/test_job_stages.py new file mode 100644 index 0000000..d566e93 --- /dev/null +++ b/backend/tests/test_job_stages.py @@ -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