- 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
101 lines
3.1 KiB
Python
101 lines
3.1 KiB
Python
"""异步研究 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
|
||
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 (
|
||
SqlAlchemyJobRepository,
|
||
)
|
||
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
|
||
|
||
router = APIRouter(prefix="/jobs", tags=["jobs"])
|
||
|
||
|
||
def _bg_factories():
|
||
from app.application.services.job_executor import default_factories
|
||
|
||
return default_factories()
|
||
|
||
|
||
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")
|