From 0ea229d766715a7de4efa226db2136db633cc581 Mon Sep 17 00:00:00 2001 From: Simon Date: Sun, 6 Sep 2026 17:18:46 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20Phase=204=20=E2=80=94=20Experiment=20?= =?UTF-8?q?=E8=87=AA=E5=8A=A8=E5=BD=92=E6=A1=A3=20+=20=E5=BC=82=E6=AD=A5?= =?UTF-8?q?=20Job=EF=BC=88=E7=8A=B6=E6=80=81=E6=9C=BA=20/=20SSE=20/=20?= =?UTF-8?q?=E4=B8=80=E9=94=AE=E5=A4=8D=E8=B7=91=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 数据表: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 通过 --- backend/app/api/deps.py | 21 +++ backend/app/api/experiments.py | 81 ++++++++ backend/app/api/jobs.py | 110 +++++++++++ backend/app/api/router.py | 4 +- .../app/application/services/job_executor.py | 146 +++++++++++++++ backend/app/domain/entities/research.py | 45 ++++- backend/app/domain/repositories/jobs.py | 25 +++ ...3113c80257f_phase4_jobs_and_experiments.py | 65 +++++++ .../persistence/sqlalchemy/models/__init__.py | 4 + .../persistence/sqlalchemy/models/jobs.py | 40 ++++ .../sqlalchemy/repositories/jobs_impl.py | 63 +++++++ backend/tests/test_jobs_experiments.py | 177 ++++++++++++++++++ frontend/web/app/experiments/page.tsx | 158 ++++++++++++++++ frontend/web/app/layout.tsx | 1 + 14 files changed, 938 insertions(+), 2 deletions(-) create mode 100644 backend/app/api/experiments.py create mode 100644 backend/app/api/jobs.py create mode 100644 backend/app/application/services/job_executor.py create mode 100644 backend/app/domain/repositories/jobs.py create mode 100644 backend/app/infrastructure/persistence/migrations/versions/53113c80257f_phase4_jobs_and_experiments.py create mode 100644 backend/app/infrastructure/persistence/sqlalchemy/models/jobs.py create mode 100644 backend/app/infrastructure/persistence/sqlalchemy/repositories/jobs_impl.py create mode 100644 backend/tests/test_jobs_experiments.py create mode 100644 frontend/web/app/experiments/page.tsx diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index fb00bc5..d455ee3 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -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)] diff --git a/backend/app/api/experiments.py b/backend/app/api/experiments.py new file mode 100644 index 0000000..e2d6bb5 --- /dev/null +++ b/backend/app/api/experiments.py @@ -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} diff --git a/backend/app/api/jobs.py b/backend/app/api/jobs.py new file mode 100644 index 0000000..ee898af --- /dev/null +++ b/backend/app/api/jobs.py @@ -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") diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 933db7c..390ecad 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -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) diff --git a/backend/app/application/services/job_executor.py b/backend/app/application/services/job_executor.py new file mode 100644 index 0000000..c060446 --- /dev/null +++ b/backend/app/application/services/job_executor.py @@ -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 diff --git a/backend/app/domain/entities/research.py b/backend/app/domain/entities/research.py index df69c37..bf531e4 100644 --- a/backend/app/domain/entities/research.py +++ b/backend/app/domain/entities/research.py @@ -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 diff --git a/backend/app/domain/repositories/jobs.py b/backend/app/domain/repositories/jobs.py new file mode 100644 index 0000000..5fe036c --- /dev/null +++ b/backend/app/domain/repositories/jobs.py @@ -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]: ... diff --git a/backend/app/infrastructure/persistence/migrations/versions/53113c80257f_phase4_jobs_and_experiments.py b/backend/app/infrastructure/persistence/migrations/versions/53113c80257f_phase4_jobs_and_experiments.py new file mode 100644 index 0000000..1a55359 --- /dev/null +++ b/backend/app/infrastructure/persistence/migrations/versions/53113c80257f_phase4_jobs_and_experiments.py @@ -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 ### diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py index a7cd56f..43e3410 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py @@ -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, diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/jobs.py b/backend/app/infrastructure/persistence/sqlalchemy/models/jobs.py new file mode 100644 index 0000000..35314b6 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/jobs.py @@ -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) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/jobs_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/jobs_impl.py new file mode 100644 index 0000000..ee3fd56 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/jobs_impl.py @@ -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] diff --git a/backend/tests/test_jobs_experiments.py b/backend/tests/test_jobs_experiments.py new file mode 100644 index 0000000..fa15837 --- /dev/null +++ b/backend/tests/test_jobs_experiments.py @@ -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() diff --git a/frontend/web/app/experiments/page.tsx b/frontend/web/app/experiments/page.tsx new file mode 100644 index 0000000..ebf0f41 --- /dev/null +++ b/frontend/web/app/experiments/page.tsx @@ -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 }; + result: BacktestResult | null; +} + +export default function ExperimentsPage() { + const [exps, setExps] = useState([]); + const [detail, setDetail] = useState(null); + const [jobId, setJobId] = useState(""); + const [jobStatus, setJobStatus] = useState(""); + const [error, setError] = useState(""); + + function load() { + apiGet("/experiments") + .then(setExps) + .catch((e: Error) => setError(e.message)); + } + + useEffect(load, []); + + function open(id: string) { + apiGet(`/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 ( + <> +

实验归档(Experiment · 可复现)

+ {error &&
{error}
} + {jobId && ( +
+ 复跑 Job {jobId} → 状态:{jobStatus} +
+ )} +
+ + + + + + + + + + + + + + {exps.map((e) => ( + + + + + + + + + + ))} + {exps.length === 0 && ( + + + + )} + +
ID类型因子区间摘要代码版本
{e.id}{e.kind}{e.factors.join(", ")} + {e.period ? `${e.period[0]} ~ ${e.period[1]}` : "—"} + {e.summary_text ?? "—"}{e.code_version ?? "—"} + + +
+ 暂无实验。先在「回测」或「因子研究」中运行一次,即会自动归档。 +
+
+ + {detail && s && ( +
+

+ {detail.id} · {detail.factors.join(", ")} · {detail.period?.[0]}~{detail.period?.[1]} +

+
+
+
总收益
+
{s.total_return_pct.toFixed(2)}%
+
+
+
年化
+
{s.annual_return_pct.toFixed(2)}%
+
+
+
Sharpe
+
{s.sharpe.toFixed(2)}
+
+
+
最大回撤
+
{s.max_drawdown_pct.toFixed(2)}%
+
+
+
+ spec(可复现输入):{JSON.stringify(detail.spec).slice(0, 400)}… +
+
+ )} + + ); +} diff --git a/frontend/web/app/layout.tsx b/frontend/web/app/layout.tsx index 83c370c..6587c7f 100644 --- a/frontend/web/app/layout.tsx +++ b/frontend/web/app/layout.tsx @@ -17,6 +17,7 @@ export default function RootLayout({ children }: { children: React.ReactNode }) 股票池 因子研究 回测 + 实验
{children}