feat: Phase 4 — Experiment 自动归档 + 异步 Job(状态机 / SSE / 一键复跑)
- 数据表:job / experiment(spec/result JSON 存档、code_version),Alembic 迁移 53113c80257f
- Job:queued→running→(success|failed) 状态机,BackgroundTasks 本地执行 + 失败兜底标记;结果与 Experiment 关联
- Experiment:每次研究成功自动归档(含 git commit 与收益摘要),支持一键复跑(同 spec 重建 Job)
- API:POST /api/jobs、GET /api/jobs/{id}(内嵌结果)、SSE /api/jobs/{id}/events、/api/experiments 列表/详情/rerun
- 前端:新增「实验」页(列表 / 详情 / 复跑 + Job 轮询);导航更新
- 端到端验证:真实 20 股 job 提交→后台执行→success→EXP 归档(-12.41%);executor 成功/失败路径单测
- 测试 73 passed(新增 5 项 Job/Experiment)/ ruff clean / 前端 tsc + build 通过
This commit is contained in:
@@ -10,6 +10,7 @@ from typing import Annotated
|
||||
from fastapi import Depends
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.domain.repositories.jobs import ExperimentRepository, JobRepository
|
||||
from app.domain.repositories.market import (
|
||||
DailyBarRepository,
|
||||
StockRepository,
|
||||
@@ -49,3 +50,23 @@ StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
|
||||
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
|
||||
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
|
||||
ResearchServiceDep = Annotated[ResearchService, Depends(_service_factory)]
|
||||
|
||||
|
||||
def _job_repo_factory(session: DbSession):
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
|
||||
return SqlAlchemyJobRepository(session)
|
||||
|
||||
|
||||
def _experiment_repo_factory(session: DbSession):
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
)
|
||||
|
||||
return SqlAlchemyExperimentRepository(session)
|
||||
|
||||
|
||||
JobRepoDep = Annotated[JobRepository, Depends(_job_repo_factory)]
|
||||
ExperimentRepoDep = Annotated[ExperimentRepository, Depends(_experiment_repo_factory)]
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
"""Experiment API(Phase 4):列表 / 详情 / 一键复跑。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, HTTPException
|
||||
|
||||
from app.api.deps import DbSession, ExperimentRepoDep, JobRepoDep
|
||||
from app.api.jobs import _bg_factories
|
||||
from app.application.services.job_executor import execute_job, new_id
|
||||
from app.domain.entities.research import (
|
||||
ExperimentRecord,
|
||||
JobRecord,
|
||||
JobStatus,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/experiments", tags=["experiments"])
|
||||
|
||||
|
||||
def _experiment_meta(exp: ExperimentRecord) -> dict:
|
||||
spec = json.loads(exp.spec_json)
|
||||
return {
|
||||
"id": exp.id,
|
||||
"kind": exp.kind,
|
||||
"factors": [f["name"] for f in spec.get("factors", [])],
|
||||
"period": spec.get("period"),
|
||||
"rebalance": spec.get("rebalance"),
|
||||
"top_n": spec.get("selection", {}).get("top_n"),
|
||||
"summary_text": exp.summary_text,
|
||||
"code_version": exp.code_version,
|
||||
"created_at": exp.created_at,
|
||||
}
|
||||
|
||||
|
||||
def _experiment_full(exp: ExperimentRecord) -> dict:
|
||||
from app.api.jobs import _decode_result
|
||||
|
||||
return {
|
||||
**_experiment_meta(exp),
|
||||
"spec": json.loads(exp.spec_json),
|
||||
"result": _decode_result(exp.kind, exp.result_json),
|
||||
}
|
||||
|
||||
|
||||
@router.get("", summary="Experiment 列表")
|
||||
def list_experiments(experiment_repo: ExperimentRepoDep) -> list[dict]:
|
||||
return [_experiment_meta(e) for e in experiment_repo.list_recent(limit=50)]
|
||||
|
||||
|
||||
@router.get("/{experiment_id}", summary="Experiment 详情(含完整结果)")
|
||||
def get_experiment(experiment_id: str, experiment_repo: ExperimentRepoDep) -> dict:
|
||||
exp = experiment_repo.get(experiment_id)
|
||||
if exp is None:
|
||||
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
|
||||
return _experiment_full(exp)
|
||||
|
||||
|
||||
@router.post("/{experiment_id}/rerun", summary="一键复跑(AGENT §21:历史实验可重放)")
|
||||
def rerun_experiment(
|
||||
experiment_id: str,
|
||||
background: BackgroundTasks,
|
||||
session: DbSession,
|
||||
experiment_repo: ExperimentRepoDep,
|
||||
job_repo: JobRepoDep,
|
||||
) -> dict:
|
||||
exp = experiment_repo.get(experiment_id)
|
||||
if exp is None:
|
||||
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"),
|
||||
kind=exp.kind,
|
||||
spec_json=exp.spec_json,
|
||||
status=JobStatus.QUEUED,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
job_repo.create(job)
|
||||
session.commit()
|
||||
background.add_task(execute_job, job.id, **_bg_factories())
|
||||
return {"job_id": job.id, "status": job.status, "origin_experiment": exp.id}
|
||||
@@ -0,0 +1,110 @@
|
||||
"""异步研究 Job API(Phase 4):提交 / 查询 / SSE 进度。
|
||||
|
||||
POST /api/jobs 创建 Job(BackgroundTasks 本地执行),立即返回 job_id
|
||||
GET /api/jobs/{id} 状态 + 结果(成功时内嵌 result)
|
||||
GET /api/jobs/{id}/events SSE 进度(queued→running→success|failed)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, HTTPException
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from app.api.deps import DbSession, JobRepoDep, _engine_factory
|
||||
from app.application.services.job_executor import execute_job, new_id
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
FactorTestReport,
|
||||
JobRecord,
|
||||
JobStatus,
|
||||
ResearchSpec,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
|
||||
|
||||
router = APIRouter(prefix="/jobs", tags=["jobs"])
|
||||
|
||||
|
||||
def _bg_factories():
|
||||
return {
|
||||
"session_factory": SessionLocal,
|
||||
"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": _engine_factory(),
|
||||
}
|
||||
|
||||
|
||||
def _decode_result(kind: str, result_json: str | None):
|
||||
if result_json is None:
|
||||
return None
|
||||
model = BacktestResult if kind == "backtest" else FactorTestReport
|
||||
return model.model_validate_json(result_json)
|
||||
|
||||
|
||||
def _job_view(job: JobRecord) -> dict:
|
||||
view = job.model_dump()
|
||||
view["result"] = _decode_result(job.kind, job.result_json)
|
||||
view.pop("result_json", None)
|
||||
view["spec"] = json.loads(job.spec_json)
|
||||
return view
|
||||
|
||||
|
||||
@router.post("", summary="创建异步研究 Job")
|
||||
def create_job(
|
||||
spec: ResearchSpec,
|
||||
background: BackgroundTasks,
|
||||
session: DbSession,
|
||||
job_repo: JobRepoDep,
|
||||
) -> dict:
|
||||
job = JobRecord(
|
||||
id=new_id("JOB"),
|
||||
kind=spec.type,
|
||||
spec_json=spec.model_dump_json(),
|
||||
status=JobStatus.QUEUED,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
job_repo.create(job)
|
||||
session.commit()
|
||||
background.add_task(execute_job, job.id, **_bg_factories())
|
||||
return {"job_id": job.id, "status": job.status}
|
||||
|
||||
|
||||
@router.get("/{job_id}", summary="查询 Job 状态与结果")
|
||||
def get_job(job_id: str, job_repo: JobRepoDep) -> dict:
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
raise HTTPException(status_code=404, detail=f"Job {job_id} 不存在")
|
||||
return _job_view(job)
|
||||
|
||||
|
||||
@router.get("/{job_id}/events", summary="Job 进度 SSE")
|
||||
async def job_events(job_id: str) -> StreamingResponse:
|
||||
async def gen():
|
||||
while True:
|
||||
with SessionLocal() as session:
|
||||
job = SqlAlchemyJobRepository(session).get(job_id)
|
||||
if job is None:
|
||||
yield "event: error\ndata: job not found\n\n"
|
||||
return
|
||||
payload = json.dumps(
|
||||
{"job_id": job.id, "status": job.status, "stage": job.stage}, ensure_ascii=False
|
||||
)
|
||||
yield f"data: {payload}\n\n"
|
||||
if job.status in (JobStatus.SUCCESS, JobStatus.FAILED, JobStatus.CANCELLED):
|
||||
return
|
||||
await asyncio.sleep(0.4)
|
||||
|
||||
return StreamingResponse(gen(), media_type="text/event-stream")
|
||||
@@ -8,10 +8,12 @@ from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api import factors, health, research, stocks
|
||||
from app.api import experiments, factors, health, jobs, research, stocks
|
||||
|
||||
api_router = APIRouter()
|
||||
api_router.include_router(health.router)
|
||||
api_router.include_router(stocks.router)
|
||||
api_router.include_router(factors.router)
|
||||
api_router.include_router(research.router)
|
||||
api_router.include_router(jobs.router)
|
||||
api_router.include_router(experiments.router)
|
||||
|
||||
Reference in New Issue
Block a user