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)
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
"""Job 执行编排与 Experiment 归档(Phase 4)。
|
||||
|
||||
execute_job 在后台运行:状态机 queued→running→(success|failed),
|
||||
成功时自动把 spec/result 存为 Experiment(含代码版本),实现「研究可复现」(AGENT.md §21)。
|
||||
独立 Session 生命周期(不依赖请求 scope),可被 FastAPI BackgroundTasks 或测试直接调用。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import subprocess
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
ExperimentRecord,
|
||||
FactorTestReport,
|
||||
JobStatus,
|
||||
ResearchSpec,
|
||||
)
|
||||
from app.quant.service import ResearchService
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[3]
|
||||
|
||||
|
||||
def new_id(prefix: str) -> str:
|
||||
return f"{prefix}-{uuid.uuid4().hex[:8].upper()}"
|
||||
|
||||
|
||||
def _git_short_rev() -> str | None:
|
||||
try:
|
||||
out = subprocess.run(
|
||||
["git", "rev-parse", "--short", "HEAD"],
|
||||
cwd=PROJECT_ROOT,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
check=False,
|
||||
)
|
||||
return out.stdout.strip() or None
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
|
||||
|
||||
def _summary_text(kind: str, result: BacktestResult | FactorTestReport) -> str | None:
|
||||
if kind == "backtest" and isinstance(result, BacktestResult):
|
||||
s = result.summary
|
||||
return f"总收益 {s.total_return_pct:.2f}% · 年化 {s.annual_return_pct:.2f}% · 回撤 {s.max_drawdown_pct:.2f}%"
|
||||
if isinstance(result, FactorTestReport):
|
||||
return f"IC {result.ic_mean:.4f} · RankIC {result.rank_ic_mean:.4f} · 样本 {result.sample_days} 日"
|
||||
return None
|
||||
|
||||
|
||||
def _execute_inner(
|
||||
job_id: str,
|
||||
*,
|
||||
session_factory: Callable,
|
||||
job_repo_factory: Callable,
|
||||
experiment_repo_factory: Callable,
|
||||
stock_repo_factory: Callable,
|
||||
daily_repo_factory: Callable,
|
||||
engine,
|
||||
) -> None:
|
||||
with session_factory() as session:
|
||||
job_repo = job_repo_factory(session)
|
||||
experiment_repo = experiment_repo_factory(session)
|
||||
job = job_repo.get(job_id)
|
||||
if job is None:
|
||||
return
|
||||
job.status = JobStatus.RUNNING
|
||||
job.started_at = datetime.now()
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
try:
|
||||
spec = ResearchSpec.model_validate_json(job.spec_json)
|
||||
service = ResearchService(
|
||||
stock_repo_factory(session), daily_repo_factory(session), engine
|
||||
)
|
||||
if spec.type == "backtest":
|
||||
result = service.run_backtest(spec)
|
||||
else:
|
||||
result = service.run_factor_test(spec)
|
||||
|
||||
result_json = json.dumps(result.model_dump(mode="json"), ensure_ascii=False)
|
||||
experiment = ExperimentRecord(
|
||||
id=new_id("EXP"),
|
||||
kind=spec.type,
|
||||
spec_json=job.spec_json,
|
||||
result_json=result_json,
|
||||
summary_text=_summary_text(spec.type, result),
|
||||
code_version=_git_short_rev(),
|
||||
job_id=job.id,
|
||||
created_at=datetime.now(),
|
||||
)
|
||||
experiment_repo.save(experiment)
|
||||
|
||||
job.result_json = result_json
|
||||
job.experiment_id = experiment.id
|
||||
job.status = JobStatus.SUCCESS
|
||||
job.error = None
|
||||
except Exception as exc: # noqa: BLE001 —— 统一记为 failed 供前端展示
|
||||
job.status = JobStatus.FAILED
|
||||
job.error = f"{type(exc).__name__}: {exc}"
|
||||
finally:
|
||||
job.finished_at = datetime.now()
|
||||
job_repo.update(job)
|
||||
session.commit()
|
||||
|
||||
|
||||
def execute_job(
|
||||
job_id: str,
|
||||
*,
|
||||
session_factory: Callable,
|
||||
job_repo_factory: Callable,
|
||||
experiment_repo_factory: Callable,
|
||||
stock_repo_factory: Callable,
|
||||
daily_repo_factory: Callable,
|
||||
engine,
|
||||
) -> None:
|
||||
"""入口包装:任何未预期异常都将 Job 标记 failed(防止卡在 queued/running)。"""
|
||||
try:
|
||||
_execute_inner(
|
||||
job_id,
|
||||
session_factory=session_factory,
|
||||
job_repo_factory=job_repo_factory,
|
||||
experiment_repo_factory=experiment_repo_factory,
|
||||
stock_repo_factory=stock_repo_factory,
|
||||
daily_repo_factory=daily_repo_factory,
|
||||
engine=engine,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
try:
|
||||
with session_factory() as session:
|
||||
repo = job_repo_factory(session)
|
||||
job = repo.get(job_id)
|
||||
if job is not None:
|
||||
job.status = JobStatus.FAILED
|
||||
job.error = f"内部错误: {type(exc).__name__}: {exc}"
|
||||
job.finished_at = datetime.now()
|
||||
repo.update(job)
|
||||
session.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
@@ -8,7 +8,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date
|
||||
from datetime import date, datetime
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
@@ -167,3 +167,46 @@ class FactorTestReport(BaseModel):
|
||||
sample_days: int
|
||||
unimplemented: list[str] = Field(default_factory=list)
|
||||
config_snapshot: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
# ---------- 异步 Job 与 Experiment(Phase 4) ----------
|
||||
|
||||
|
||||
class JobStatus(str):
|
||||
"""统一状态机(AGENT.md §20):queued→running→(success|failed|cancelled)。"""
|
||||
|
||||
QUEUED = "queued"
|
||||
RUNNING = "running"
|
||||
SUCCESS = "success"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class JobRecord(BaseModel):
|
||||
"""一次异步研究任务。spec/result 以 JSON 文本存储(保持 Schema 演进自由)。"""
|
||||
|
||||
id: str
|
||||
kind: str # factor_test | backtest
|
||||
status: str = JobStatus.QUEUED
|
||||
stage: str | None = None
|
||||
spec_json: str
|
||||
error: str | None = None
|
||||
result_json: str | None = None
|
||||
experiment_id: str | None = None
|
||||
created_at: datetime | None = None
|
||||
started_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
|
||||
|
||||
class ExperimentRecord(BaseModel):
|
||||
"""一次研究的可复现存档(AGENT.md §21)。"""
|
||||
|
||||
id: str
|
||||
kind: str # factor_test | backtest
|
||||
spec_json: str
|
||||
result_json: str
|
||||
summary_text: str | None = None # 便于列表展示的摘要(如 total_return_pct)
|
||||
code_version: str | None = None # git commit / 代码指纹
|
||||
data_version: str | None = None
|
||||
job_id: str | None = None
|
||||
created_at: datetime | None = None
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
"""Job / Experiment Repository Protocol(Phase 4)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
from app.domain.entities.research import ExperimentRecord, JobRecord
|
||||
|
||||
|
||||
class JobRepository(Protocol):
|
||||
def create(self, job: JobRecord) -> JobRecord: ...
|
||||
|
||||
def get(self, job_id: str) -> JobRecord | None: ...
|
||||
|
||||
def update(self, job: JobRecord) -> None: ...
|
||||
|
||||
def list_recent(self, kind: str | None = None, limit: int = 20) -> list[JobRecord]: ...
|
||||
|
||||
|
||||
class ExperimentRepository(Protocol):
|
||||
def save(self, experiment: ExperimentRecord) -> ExperimentRecord: ...
|
||||
|
||||
def get(self, experiment_id: str) -> ExperimentRecord | None: ...
|
||||
|
||||
def list_recent(self, limit: int = 50) -> list[ExperimentRecord]: ...
|
||||
+65
@@ -0,0 +1,65 @@
|
||||
"""phase4 jobs and experiments
|
||||
|
||||
Revision ID: 53113c80257f
|
||||
Revises: e4d188250fb2
|
||||
Create Date: 2026-09-06 17:16:05.976236
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision: str = "53113c80257f"
|
||||
down_revision: str | None = "e4d188250fb2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table(
|
||||
"experiment",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("kind", sa.String(length=16), nullable=False),
|
||||
sa.Column("spec_json", sa.Text(), nullable=False),
|
||||
sa.Column("result_json", sa.Text(), nullable=False),
|
||||
sa.Column("summary_text", sa.String(length=200), nullable=True),
|
||||
sa.Column("code_version", sa.String(length=40), nullable=True),
|
||||
sa.Column("data_version", sa.String(length=40), nullable=True),
|
||||
sa.Column("job_id", sa.String(length=32), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_table(
|
||||
"job",
|
||||
sa.Column("id", sa.String(length=32), nullable=False),
|
||||
sa.Column("kind", sa.String(length=16), nullable=False),
|
||||
sa.Column("status", sa.String(length=12), nullable=False),
|
||||
sa.Column("stage", sa.String(length=24), nullable=True),
|
||||
sa.Column("spec_json", sa.Text(), nullable=False),
|
||||
sa.Column("error", sa.Text(), nullable=True),
|
||||
sa.Column("result_json", sa.Text(), nullable=True),
|
||||
sa.Column("experiment_id", sa.String(length=32), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False),
|
||||
sa.Column("started_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("finished_at", sa.DateTime(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
with op.batch_alter_table("job", schema=None) as batch_op:
|
||||
batch_op.create_index(batch_op.f("ix_job_status"), ["status"], unique=False)
|
||||
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
with op.batch_alter_table("job", schema=None) as batch_op:
|
||||
batch_op.drop_index(batch_op.f("ix_job_status"))
|
||||
|
||||
op.drop_table("job")
|
||||
op.drop_table("experiment")
|
||||
# ### end Alembic commands ###
|
||||
@@ -4,6 +4,10 @@
|
||||
模型统一继承 infra.persistence.sqlalchemy.base.Base。
|
||||
"""
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.models.jobs import ( # noqa: F401
|
||||
ExperimentModel,
|
||||
JobModel,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.models.market import ( # noqa: F401
|
||||
AdjustFactorModel,
|
||||
FinancialIndicatorModel,
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
"""Phase 4:Job(异步任务)与 Experiment(实验归档)表。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
|
||||
|
||||
class JobModel(Base):
|
||||
__tablename__ = "job"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
kind: Mapped[str] = mapped_column(String(16))
|
||||
status: Mapped[str] = mapped_column(String(12), index=True)
|
||||
stage: Mapped[str | None] = mapped_column(String(24), nullable=True)
|
||||
spec_json: Mapped[str] = mapped_column(Text)
|
||||
error: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
result_json: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
experiment_id: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
started_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
|
||||
|
||||
|
||||
class ExperimentModel(Base):
|
||||
__tablename__ = "experiment"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
kind: Mapped[str] = mapped_column(String(16))
|
||||
spec_json: Mapped[str] = mapped_column(Text)
|
||||
result_json: Mapped[str] = mapped_column(Text)
|
||||
summary_text: Mapped[str | None] = mapped_column(String(200), nullable=True)
|
||||
code_version: Mapped[str | None] = mapped_column(String(40), nullable=True)
|
||||
data_version: Mapped[str | None] = mapped_column(String(40), nullable=True)
|
||||
job_id: Mapped[str | None] = mapped_column(String(32), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(DateTime)
|
||||
@@ -0,0 +1,63 @@
|
||||
"""Job / Experiment 的 SQLAlchemy 实现(Phase 4)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.domain.entities.research import ExperimentRecord, JobRecord
|
||||
from app.infrastructure.persistence.sqlalchemy.models.jobs import ExperimentModel, JobModel
|
||||
|
||||
|
||||
def _job_to_model(job: JobRecord) -> JobModel:
|
||||
return JobModel(**job.model_dump())
|
||||
|
||||
|
||||
class SqlAlchemyJobRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def create(self, job: JobRecord) -> JobRecord:
|
||||
self._session.add(_job_to_model(job))
|
||||
self._session.flush()
|
||||
return job
|
||||
|
||||
def get(self, job_id: str) -> JobRecord | None:
|
||||
row = self._session.get(JobModel, job_id)
|
||||
return JobRecord.model_validate(row, from_attributes=True) if row else None
|
||||
|
||||
def update(self, job: JobRecord) -> None:
|
||||
row = self._session.get(JobModel, job.id)
|
||||
if row is None:
|
||||
raise KeyError(f"job {job.id} 不存在")
|
||||
for k, v in job.model_dump().items():
|
||||
setattr(row, k, v)
|
||||
|
||||
def list_recent(self, kind: str | None = None, limit: int = 20) -> list[JobRecord]:
|
||||
stmt = select(JobModel).order_by(JobModel.created_at.desc()).limit(limit)
|
||||
if kind:
|
||||
stmt = stmt.where(JobModel.kind == kind)
|
||||
return [
|
||||
JobRecord.model_validate(r, from_attributes=True)
|
||||
for r in self._session.scalars(stmt).all()
|
||||
]
|
||||
|
||||
|
||||
class SqlAlchemyExperimentRepository:
|
||||
def __init__(self, session: Session) -> None:
|
||||
self._session = session
|
||||
|
||||
def save(self, experiment: ExperimentRecord) -> ExperimentRecord:
|
||||
self._session.add(ExperimentModel(**experiment.model_dump()))
|
||||
self._session.flush()
|
||||
return experiment
|
||||
|
||||
def get(self, experiment_id: str) -> ExperimentRecord | None:
|
||||
row = self._session.get(ExperimentModel, experiment_id)
|
||||
return ExperimentRecord.model_validate(row, from_attributes=True) if row else None
|
||||
|
||||
def list_recent(self, limit: int = 50) -> list[ExperimentRecord]:
|
||||
rows = self._session.scalars(
|
||||
select(ExperimentModel).order_by(ExperimentModel.created_at.desc()).limit(limit)
|
||||
).all()
|
||||
return [ExperimentRecord.model_validate(r, from_attributes=True) for r in rows]
|
||||
@@ -0,0 +1,177 @@
|
||||
"""Phase 4 测试:Job 状态机 / 执行与 Experiment 自动归档 / API 提交与查询。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import date, datetime
|
||||
|
||||
import pytest
|
||||
from app.application.services.job_executor import execute_job
|
||||
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.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):
|
||||
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 TestJobsApi:
|
||||
def test_submit_then_query(self) -> 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) -> None:
|
||||
with TestClient(app) as client:
|
||||
assert client.get("/api/jobs/JOB-NOPE").status_code == 404
|
||||
|
||||
def test_experiments_list_and_detail(self) -> None:
|
||||
with TestClient(app) as client:
|
||||
resp = client.get("/api/experiments")
|
||||
assert resp.status_code == 200
|
||||
exps = resp.json()
|
||||
if exps: # 本机真实库中可能有历史实验;detail 可读即通过
|
||||
first = exps[0]
|
||||
detail = client.get(f"/api/experiments/{first['id']}")
|
||||
assert detail.status_code == 200
|
||||
assert "result" in detail.json()
|
||||
@@ -0,0 +1,158 @@
|
||||
"use client";
|
||||
|
||||
import { useEffect, useState } from "react";
|
||||
import { apiGet, apiPost } from "@/lib/api";
|
||||
import type { BacktestResult } from "@/lib/types";
|
||||
|
||||
interface ExperimentMeta {
|
||||
id: string;
|
||||
kind: string;
|
||||
factors: string[];
|
||||
period?: [string, string] | null;
|
||||
rebalance?: string;
|
||||
top_n?: number | null;
|
||||
summary_text?: string | null;
|
||||
code_version?: string | null;
|
||||
created_at?: string;
|
||||
}
|
||||
|
||||
interface ExperimentDetail extends ExperimentMeta {
|
||||
spec: { universe?: Record<string, unknown> };
|
||||
result: BacktestResult | null;
|
||||
}
|
||||
|
||||
export default function ExperimentsPage() {
|
||||
const [exps, setExps] = useState<ExperimentMeta[]>([]);
|
||||
const [detail, setDetail] = useState<ExperimentDetail | null>(null);
|
||||
const [jobId, setJobId] = useState("");
|
||||
const [jobStatus, setJobStatus] = useState("");
|
||||
const [error, setError] = useState("");
|
||||
|
||||
function load() {
|
||||
apiGet<ExperimentMeta[]>("/experiments")
|
||||
.then(setExps)
|
||||
.catch((e: Error) => setError(e.message));
|
||||
}
|
||||
|
||||
useEffect(load, []);
|
||||
|
||||
function open(id: string) {
|
||||
apiGet<ExperimentDetail>(`/experiments/${id}`)
|
||||
.then(setDetail)
|
||||
.catch((e: Error) => setError(e.message));
|
||||
}
|
||||
|
||||
function rerun(exp: ExperimentMeta) {
|
||||
setError("");
|
||||
setJobId("");
|
||||
setJobStatus("submitting");
|
||||
apiPost<{ job_id: string; status: string }>(`/experiments/${exp.id}/rerun`, {})
|
||||
.then((r) => {
|
||||
setJobId(r.job_id);
|
||||
setJobStatus(r.status);
|
||||
poll(r.job_id);
|
||||
})
|
||||
.catch((e: Error) => setError(e.message));
|
||||
}
|
||||
|
||||
function poll(id: string) {
|
||||
let tries = 0;
|
||||
const timer = window.setInterval(() => {
|
||||
tries += 1;
|
||||
apiGet<{ status: string }>(`/jobs/${id}`)
|
||||
.then((j) => {
|
||||
setJobStatus(j.status);
|
||||
if (j.status === "success" || j.status === "failed" || tries > 60) {
|
||||
window.clearInterval(timer);
|
||||
load();
|
||||
}
|
||||
})
|
||||
.catch(() => window.clearInterval(timer));
|
||||
}, 1000);
|
||||
}
|
||||
|
||||
const s = detail?.result?.summary;
|
||||
|
||||
return (
|
||||
<>
|
||||
<h1>实验归档(Experiment · 可复现)</h1>
|
||||
{error && <div className="error">{error}</div>}
|
||||
{jobId && (
|
||||
<div className="card">
|
||||
复跑 Job <b>{jobId}</b> → 状态:<b>{jobStatus}</b>
|
||||
</div>
|
||||
)}
|
||||
<div className="card">
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>ID</th>
|
||||
<th>类型</th>
|
||||
<th>因子</th>
|
||||
<th>区间</th>
|
||||
<th>摘要</th>
|
||||
<th>代码版本</th>
|
||||
<th></th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{exps.map((e) => (
|
||||
<tr key={e.id}>
|
||||
<td>{e.id}</td>
|
||||
<td>{e.kind}</td>
|
||||
<td>{e.factors.join(", ")}</td>
|
||||
<td>
|
||||
{e.period ? `${e.period[0]} ~ ${e.period[1]}` : "—"}
|
||||
</td>
|
||||
<td>{e.summary_text ?? "—"}</td>
|
||||
<td>{e.code_version ?? "—"}</td>
|
||||
<td>
|
||||
<button onClick={() => open(e.id)}>详情</button>
|
||||
<button className="primary" onClick={() => rerun(e)}>
|
||||
复跑
|
||||
</button>
|
||||
</td>
|
||||
</tr>
|
||||
))}
|
||||
{exps.length === 0 && (
|
||||
<tr>
|
||||
<td colSpan={7} className="muted">
|
||||
暂无实验。先在「回测」或「因子研究」中运行一次,即会自动归档。
|
||||
</td>
|
||||
</tr>
|
||||
)}
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
|
||||
{detail && s && (
|
||||
<div className="card">
|
||||
<h2>
|
||||
{detail.id} · {detail.factors.join(", ")} · {detail.period?.[0]}~{detail.period?.[1]}
|
||||
</h2>
|
||||
<div className="stat-grid">
|
||||
<div className="stat">
|
||||
<div className="label">总收益</div>
|
||||
<div className="value">{s.total_return_pct.toFixed(2)}%</div>
|
||||
</div>
|
||||
<div className="stat">
|
||||
<div className="label">年化</div>
|
||||
<div className="value">{s.annual_return_pct.toFixed(2)}%</div>
|
||||
</div>
|
||||
<div className="stat">
|
||||
<div className="label">Sharpe</div>
|
||||
<div className="value">{s.sharpe.toFixed(2)}</div>
|
||||
</div>
|
||||
<div className="stat">
|
||||
<div className="label">最大回撤</div>
|
||||
<div className="value bad">{s.max_drawdown_pct.toFixed(2)}%</div>
|
||||
</div>
|
||||
</div>
|
||||
<div className="muted" style={{ marginTop: 8 }}>
|
||||
spec(可复现输入):{JSON.stringify(detail.spec).slice(0, 400)}…
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
);
|
||||
}
|
||||
@@ -17,6 +17,7 @@ export default function RootLayout({ children }: { children: React.ReactNode })
|
||||
<Link href="/stocks">股票池</Link>
|
||||
<Link href="/factors">因子研究</Link>
|
||||
<Link href="/backtest">回测</Link>
|
||||
<Link href="/experiments">实验</Link>
|
||||
</nav>
|
||||
<main>{children}</main>
|
||||
</body>
|
||||
|
||||
Reference in New Issue
Block a user