Files
qlib/backend/app/application/services/job_executor.py
T
Simon d9be75a98f feat: Phase 5 — AI Research Agent(受控工具白名单 + LLM 编排 + API)
- agent/tools.py:Tool 元数据(JSON Schema)+ 白名单调用(异常转可读反馈,不中断对话)
- agent/tools_impl.py:6 个受控工具 search_stocks / get_market_data / test_factor / run_backtest / get_experiment / compare_experiments —— 全部只读经 Job/Experiment 链路,研究自动归档;无 shell/任意执行/写删数据能力
- agent/llm.py:LLMClient 抽象 + OpenAI 兼容客户端(LLM_API_KEY/LLM_BASE_URL/LLM_MODEL 走 .env,未配置给出引导提示)+ 研究纪律 system prompt(反过拟合/样本外/成本)
- agent/service.py:编排循环(tool/final JSON 决策 → 执行 → 回喂 → 结论),轮次上限兜底,未知工具拒绝
- /api/agent/chat;httpx 移至主依赖;Job 默认工厂抽取(api/agent/executor 复用)
- 测试 6 项(白名单无 shell、完整研究循环产出、未知工具拒绝、轮次兜底),全量 79 passed / ruff clean
2026-09-06 17:22:05 +08:00

201 lines
6.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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,
JobRecord,
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
def default_factories() -> dict:
"""后台执行所需的独立 Session / Repository / 引擎装配(跨请求生命周期)。"""
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
from app.quant.engine import LocalEngine
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": LocalEngine(),
}
def submit_and_run(spec: ResearchSpec, *, factories: dict | None = None) -> JobRecord:
"""同步创建并执行一个 Job(复用 Job 状态机与 Experiment 归档),返回终态 Job。"""
facts = factories or default_factories()
session_factory = facts["session_factory"]
job_repo = facts["job_repo_factory"]
job = JobRecord(
id=new_id("JOB"),
kind=spec.type,
spec_json=json.dumps(spec.model_dump(mode="json"), ensure_ascii=False),
status=JobStatus.QUEUED,
created_at=datetime.now(),
)
with session_factory() as session:
job_repo(session).create(job)
session.commit()
execute_job(
job.id,
session_factory=session_factory,
job_repo_factory=job_repo,
experiment_repo_factory=facts["experiment_repo_factory"],
stock_repo_factory=facts["stock_repo_factory"],
daily_repo_factory=facts["daily_repo_factory"],
engine=facts["engine"],
)
with session_factory() as session:
done = job_repo(session).get(job.id)
assert done is not None
return done