"""Phase 4 测试:Job 状态机 / 执行与 Experiment 自动归档 / API 提交与查询。""" from __future__ import annotations import sys from datetime import date, datetime import pytest from app.application.services import job_executor as je from app.application.services.job_executor import execute_job from app.domain.entities.research import ( JobRecord, JobStatus, ResearchSpec, ) from app.infrastructure.persistence.sqlalchemy import session as sess_mod from app.infrastructure.persistence.sqlalchemy.base import Base from app.infrastructure.persistence.sqlalchemy.models.market import ( StockDailyModel, StockModel, ) from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import ( SqlAlchemyExperimentRepository, SqlAlchemyJobRepository, ) 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 Session 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}, factors=[{"name": "momentum_20", "weight": 1.0}], selection={"top_n": 1}, rebalance="monthly", period=["2024-03-01", "2024-10-31"], ) base.update(kw) return ResearchSpec.model_validate(base) class _FakeStockRepo: def __init__(self, stocks): self._stocks = stocks def list(self): return self._stocks class _FakeDailyRepo: def __init__(self, bars): self._bars = bars def get_range_many(self, symbols, start, end, adjust="none"): return [b for b in self._bars if b.symbol in symbols and start <= b.trade_date <= end] def get_range(self, symbol, start, end): return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end] def _fakes(): from app.domain.entities.market import Stock stocks = [ Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS) ] drifts = {s: 0.003 - 0.0015 * i for i, s in enumerate(_SYMS)} bars = bars_dataframe_to_daily_bars(synthetic_daily(drifts, n=300)) return stocks, bars @pytest.fixture() def sf(tmp_path): engine = create_engine(f"sqlite:///{tmp_path / 'jobs.db'}", future=True) Base.metadata.create_all(engine) return lambda: Session(engine) class TestJobExecutor: def test_success_archives_experiment(self, sf) -> None: stocks, bars = _fakes() job = JobRecord( id="JOB-TEST-1", kind="backtest", spec_json=_spec().model_dump_json(), status=JobStatus.QUEUED, created_at=datetime.now(), ) with sf() as session: SqlAlchemyJobRepository(session).create(job) session.commit() execute_job( "JOB-TEST-1", session_factory=sf, job_repo_factory=lambda s: SqlAlchemyJobRepository(s), experiment_repo_factory=lambda s: SqlAlchemyExperimentRepository(s), stock_repo_factory=lambda s: _FakeStockRepo(stocks), daily_repo_factory=lambda s: _FakeDailyRepo(bars), engine=LocalEngine(), ) with sf() as session: done = SqlAlchemyJobRepository(session).get("JOB-TEST-1") assert done is not None assert done.status == JobStatus.SUCCESS assert done.result_json is not None assert done.experiment_id is not None exp = SqlAlchemyExperimentRepository(session).get(done.experiment_id) assert exp is not None assert exp.kind == "backtest" assert exp.summary_text is not None assert "总收益" in (exp.summary_text or "") def test_failure_marks_failed(self, sf) -> None: stocks, _bars = _fakes() job = JobRecord( id="JOB-TEST-2", kind="backtest", spec_json=_spec(factors=[{"name": "no_such", "weight": 1.0}]).model_dump_json(), status=JobStatus.QUEUED, created_at=datetime.now(), ) with sf() as session: SqlAlchemyJobRepository(session).create(job) session.commit() execute_job( "JOB-TEST-2", session_factory=sf, job_repo_factory=lambda s: SqlAlchemyJobRepository(s), experiment_repo_factory=lambda s: SqlAlchemyExperimentRepository(s), stock_repo_factory=lambda s: _FakeStockRepo(stocks), daily_repo_factory=lambda s: _FakeDailyRepo([]), engine=LocalEngine(), ) with sf() as session: done = SqlAlchemyJobRepository(session).get("JOB-TEST-2") assert done is not None assert done.status == JobStatus.FAILED assert done.error class TestJobResilience: """内存隔离专项:启动残留清理 + 子进程异常退出兜底。""" def test_mark_stale_jobs_failed(self, sf) -> None: repo_f = lambda s: SqlAlchemyJobRepository(s) # noqa: E731 rows = [ ("JOB-STALE-1", JobStatus.RUNNING), ("JOB-STALE-2", JobStatus.QUEUED), ("JOB-OK-3", JobStatus.SUCCESS), ] for jid, status in rows: job = JobRecord( id=jid, kind="backtest", spec_json=_spec().model_dump_json(), status=status, created_at=datetime.now(), ) with sf() as session: repo_f(session).create(job) session.commit() n = je.mark_stale_jobs_failed(session_factory=sf, job_repo_factory=repo_f) assert n == 2 with sf() as session: stale = repo_f(session).get("JOB-STALE-1") queued = repo_f(session).get("JOB-STALE-2") ok = repo_f(session).get("JOB-OK-3") assert stale is not None and stale.status == JobStatus.FAILED assert stale.finished_at is not None and "重启" in (stale.error or "") assert queued is not None and queued.status == JobStatus.FAILED assert ok is not None and ok.status == JobStatus.SUCCESS def test_subprocess_abnormal_exit_marks_failed(self, sf, monkeypatch) -> None: repo_f = lambda s: SqlAlchemyJobRepository(s) # noqa: E731 job = JobRecord( id="JOB-SUB-1", kind="backtest", spec_json=_spec().model_dump_json(), status=JobStatus.QUEUED, created_at=datetime.now(), ) with sf() as session: repo_f(session).create(job) session.commit() monkeypatch.setattr(je, "_job_mode", lambda: "subprocess") monkeypatch.setattr( je, "_job_worker_cmd", lambda job_id: [sys.executable, "-c", "import sys; sys.exit(7)"] ) je.run_job_background("JOB-SUB-1", session_factory=sf, job_repo_factory=repo_f) with sf() as session: done = repo_f(session).get("JOB-SUB-1") assert done is not None assert done.status == JobStatus.FAILED assert done.error is not None and "code=7" in done.error @pytest.fixture() def seeded_api_db(tmp_path, monkeypatch) -> None: """把全局 SessionLocal 指向 tmp 种子库(5 只股票 × 300 交易日), Job 后台执行跑小数据,避免依赖真实全市场库(779 万行)导致过慢。""" from sqlalchemy.orm import sessionmaker engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True) Base.metadata.create_all(engine) sf = sessionmaker(bind=engine, expire_on_commit=False) drifts = {s: 0.003 - 0.0015 * i for i, s in enumerate(_SYMS)} daily = synthetic_daily(drifts, n=320) with sf() as session: for i, sym in enumerate(_SYMS): session.add(StockModel(symbol=sym, name=f"测试股份{i}", list_date=date(1999, 1, 1))) rows = [ { "symbol": r.symbol, "trade_date": r.trade_date, "source": "tushare", "adjust": "none", "open": float(r.open), "high": float(r.high), "low": float(r.low), "close": float(r.close), "volume": float(r.volume), "amount": float(r.amount), } for r in daily.itertuples() ] session.execute(StockDailyModel.__table__.insert(), rows) session.commit() monkeypatch.setattr(sess_mod, "SessionLocal", sf) class TestJobsApi: def test_submit_then_query(self, seeded_api_db) -> None: with TestClient(app) as client: resp = client.post("/api/jobs", json=_spec().model_dump(mode="json")) assert resp.status_code == 200 body = resp.json() assert body["status"] == "queued" job_id = body["job_id"] # 后台任务随请求执行完成(TestClient 同步运行 BackgroundTasks) for _ in range(20): state = client.get(f"/api/jobs/{job_id}").json() if state["status"] in (JobStatus.SUCCESS, JobStatus.FAILED): break assert state["status"] == JobStatus.SUCCESS assert state["result"] is not None assert state["spec"]["factors"][0]["name"] == "momentum_20" assert state["experiment_id"] def test_unknown_job_404(self, seeded_api_db) -> None: with TestClient(app) as client: assert client.get("/api/jobs/JOB-NOPE").status_code == 404 def test_experiments_list_and_detail(self, seeded_api_db) -> None: with TestClient(app) as client: # 先跑一个 job 生成归档,再验证列表与详情 resp = client.post("/api/jobs", json=_spec().model_dump(mode="json")) assert resp.status_code == 200 job_id = resp.json()["job_id"] for _ in range(30): st = client.get(f"/api/jobs/{job_id}").json()["status"] if st in (JobStatus.SUCCESS, JobStatus.FAILED): break assert st == JobStatus.SUCCESS exps = client.get("/api/experiments").json() assert exps, "job 成功后应有 Experiment 归档" first = exps[0] detail = client.get(f"/api/experiments/{first['id']}") assert detail.status_code == 200 assert "result" in detail.json()