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
This commit is contained in:
@@ -0,0 +1,77 @@
|
||||
"""LLM 客户端抽象(Phase 5)。
|
||||
|
||||
实现约定:真实 Key 来自 .env(LLM_API_KEY / LLM_BASE_URL / LLM_MODEL,见 config.yaml 引用)。
|
||||
测试注入 FakeLLMClient 走完整编排链路,不触网。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
from app.agent.tools import tools_schema
|
||||
from app.agent.tools_impl import build_tools
|
||||
from app.core.config import get_settings
|
||||
|
||||
SYSTEM_TEMPLATE = """你是个人 A 股量化研究助手(Research Assistant),不是系统管理员。
|
||||
|
||||
你可以调用以下工具(每次只能输出一个动作):
|
||||
{tools}
|
||||
|
||||
输出规则:只输出一行 JSON,两种形态之一:
|
||||
1. 需要调用工具:{{"tool": "<工具名>", "args": {{...}}}}
|
||||
2. 给出结论:{{"final": "结论文本"}}
|
||||
|
||||
研究纪律(必须遵守):
|
||||
- 先提出假设 → 用工具做因子测试或回测 → 基于实验事实分析,再给结论
|
||||
- 不得仅凭单次回测高收益就宣布策略有效;要主动说明样本外、过拟合、
|
||||
look-ahead bias、交易成本、参数敏感性等风险(未做验证的项要明说「未验证」)
|
||||
- 全程只读:不得要求删除/修改数据或执行任意命令(你也没有这类工具)
|
||||
- 回答使用简体中文
|
||||
"""
|
||||
|
||||
|
||||
class LLMClient(Protocol):
|
||||
def chat(self, messages: list[dict]) -> str: ...
|
||||
|
||||
|
||||
class OpenAICompatibleClient:
|
||||
"""OpenAI 兼容 Chat Completions(qwen/dashscope、deepseek、openai 等均适用)。"""
|
||||
|
||||
def __init__(self, *, base_url: str, api_key: str, model: str, timeout: float = 60.0) -> None:
|
||||
self._base_url = base_url.rstrip("/")
|
||||
self._api_key = api_key
|
||||
self._model = model
|
||||
self._timeout = timeout
|
||||
|
||||
def chat(self, messages: list[dict]) -> str:
|
||||
with httpx.Client(timeout=self._timeout) as client:
|
||||
resp = client.post(
|
||||
f"{self._base_url}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {self._api_key}"},
|
||||
json={
|
||||
"model": self._model,
|
||||
"messages": messages,
|
||||
"temperature": 0.2,
|
||||
},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data["choices"][0]["message"]["content"]
|
||||
|
||||
|
||||
def build_llm_from_settings():
|
||||
"""从配置构造 LLM;未配置 Key 时返回 None(调用方给出引导提示)。"""
|
||||
settings = get_settings()
|
||||
if not settings.llm_api_key:
|
||||
return None
|
||||
return OpenAICompatibleClient(
|
||||
base_url=settings.llm_base_url or "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
api_key=settings.llm_api_key,
|
||||
model=settings.llm_model,
|
||||
)
|
||||
|
||||
|
||||
def system_prompt() -> str:
|
||||
return SYSTEM_TEMPLATE.format(tools=tools_schema(build_tools()))
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Agent 编排:自然语言 → 受控工具调用循环 → 结论(AGENT.md §29 假设-实验-分析循环)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from app.agent.llm import SYSTEM_TEMPLATE, LLMClient
|
||||
from app.agent.tools import Tool, tools_schema
|
||||
from app.agent.tools_impl import build_tools
|
||||
|
||||
MAX_TOOL_ROUNDS = 5
|
||||
_DECISION_PATTERN = re.compile(r"\{.*\}", re.DOTALL)
|
||||
|
||||
|
||||
def _parse_decision(text: str) -> dict[str, Any]:
|
||||
"""容忍 LLM 输出中的代码块/前后缀,提取首个 JSON 对象。"""
|
||||
match = _DECISION_PATTERN.search(text)
|
||||
if not match:
|
||||
raise ValueError(f"无法从模型输出中解析动作 JSON:{text[:200]}")
|
||||
try:
|
||||
return json.loads(match.group(0))
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"模型输出的 JSON 不合法:{text[:200]}") from exc
|
||||
|
||||
|
||||
class AgentService:
|
||||
"""受控研究 Agent:每轮让 LLM 决策 tool/final,执行工具并把结果回喂,直至 final。"""
|
||||
|
||||
def __init__(self, llm: LLMClient, tools: list[Tool] | None = None) -> None:
|
||||
self._llm = llm
|
||||
self._tools = {t.name: t for t in (tools or build_tools())}
|
||||
|
||||
def _by_name(self, name: str) -> Tool:
|
||||
tool = self._tools.get(name)
|
||||
if tool is None:
|
||||
raise ValueError(f"工具不存在:{name}(可用 {sorted(self._tools)})")
|
||||
return tool
|
||||
|
||||
def chat(self, message: str, max_rounds: int = MAX_TOOL_ROUNDS) -> dict:
|
||||
messages: list[dict] = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": SYSTEM_TEMPLATE.format(tools=tools_schema(list(self._tools.values()))),
|
||||
},
|
||||
{"role": "user", "content": message},
|
||||
]
|
||||
actions: list[dict] = []
|
||||
|
||||
for _round in range(max_rounds):
|
||||
try:
|
||||
decision = _parse_decision(self._llm.chat(messages))
|
||||
except ValueError as exc:
|
||||
# LLM 回复格式异常:回传错误要求重试一次结构化输出
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"输出格式错误,请只输出一行 JSON(tool 或 final):{exc}",
|
||||
}
|
||||
)
|
||||
continue
|
||||
if "final" in decision:
|
||||
return {"reply": str(decision["final"]), "actions": actions}
|
||||
tool_name = str(decision.get("tool", ""))
|
||||
args = decision.get("args") or {}
|
||||
if not isinstance(args, dict):
|
||||
args = {}
|
||||
tool = self._by_name(tool_name) # 白名单之外的调用直接报错
|
||||
output = tool.invoke(args)
|
||||
actions.append({"tool": tool_name, "args": args, "output": output[:1000]})
|
||||
messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": f'{{"tool": "{tool_name}", "args": {json.dumps(args, ensure_ascii=False)}}}',
|
||||
}
|
||||
)
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"工具 {tool_name} 返回:\n{output}\n请继续(如需再调用输出 tool,否则输出 final)。",
|
||||
}
|
||||
)
|
||||
|
||||
# 轮次耗尽:强制收尾
|
||||
messages.append(
|
||||
{
|
||||
"role": "user",
|
||||
"content": "已达工具调用轮次上限,请直接给出基于已有事实的最终结论(只输出 final JSON)。",
|
||||
}
|
||||
)
|
||||
try:
|
||||
decision = _parse_decision(self._llm.chat(messages))
|
||||
except ValueError:
|
||||
decision = {
|
||||
"final": "研究轮次耗尽且模型未给出结构化结论,请人工查看 actions 中的实验输出。"
|
||||
}
|
||||
return {"reply": str(decision.get("final", decision)), "actions": actions}
|
||||
@@ -0,0 +1,45 @@
|
||||
"""AI Research Agent(Phase 5)。
|
||||
|
||||
Agent 是 Research Assistant(AGENT.md §28/§29):
|
||||
- 只能调用本目录 tools 提供的**白名单受控工具**(只读研究能力)
|
||||
- 禁止:执行任意 shell / 修改删除数据 / 修改配置与凭证
|
||||
- 研究行为:提出假设 → 建立实验 → 运行测试 → 分析 → 下一步,禁止以单次高收益宣告策略有效
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
TOOL_CALL_TAG = "__tool__"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Tool:
|
||||
"""受控工具元数据(LLM 可见的 JSON Schema + 调用实现)。"""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
parameters: dict[str, Any]
|
||||
handler: Callable[[dict[str, Any]], str] = field(repr=False)
|
||||
|
||||
def schema(self) -> dict[str, Any]:
|
||||
return {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"parameters": self.parameters,
|
||||
}
|
||||
|
||||
def invoke(self, args: dict[str, Any]) -> str:
|
||||
"""执行工具;任何异常都转为可读错误串(不向上抛,避免中断整轮对话)。"""
|
||||
try:
|
||||
result = self.handler(args)
|
||||
return result if isinstance(result, str) else json.dumps(result, ensure_ascii=False)
|
||||
except Exception as exc: # noqa: BLE001 —— 工具异常反馈给 LLM 而非崩溃
|
||||
return json.dumps({"error": f"{type(exc).__name__}: {exc}"}, ensure_ascii=False)
|
||||
|
||||
|
||||
def tools_schema(tools: list[Tool]) -> str:
|
||||
return json.dumps([t.schema() for t in tools], ensure_ascii=False, indent=1)
|
||||
@@ -1 +0,0 @@
|
||||
"""Agent 受控工具集:每个 Tool 是后端服务的只读 / 沙箱化入口。"""
|
||||
@@ -0,0 +1,248 @@
|
||||
"""受控工具集实现(AGENT.md §28):Agent 只能调用这里的白名单工具。
|
||||
|
||||
全部工具经 Job/Experiment 链路或只读查询执行:
|
||||
- 不提供 shell / 任意代码执行 / 修改配置与凭证 / 删除数据
|
||||
- 任何研究都会产出 Experiment 归档(可复现)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import date
|
||||
|
||||
from app.agent.tools import Tool
|
||||
from app.application.services.job_executor import default_factories, submit_and_run
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
FactorTestReport,
|
||||
ResearchSpec,
|
||||
)
|
||||
|
||||
|
||||
def _day(text: str) -> date:
|
||||
return date.fromisoformat(text)
|
||||
|
||||
|
||||
def _pick(mapping: dict, key: str, default=None):
|
||||
val = mapping.get(key, default)
|
||||
if isinstance(val, str):
|
||||
val = val.strip()
|
||||
if val == "":
|
||||
return default
|
||||
return val
|
||||
|
||||
|
||||
def build_tools(factories: dict | None = None) -> list[Tool]:
|
||||
facts = factories or default_factories()
|
||||
session_factory = facts["session_factory"]
|
||||
stock_repo_f = facts["stock_repo_factory"]
|
||||
daily_repo_f = facts["daily_repo_factory"]
|
||||
exp_repo_f = facts["experiment_repo_factory"]
|
||||
|
||||
def search_stocks(args: dict) -> str:
|
||||
q = str(_pick(args, "q", "") or "").upper()
|
||||
with session_factory() as session:
|
||||
stocks = stock_repo_f(session).list()
|
||||
rows = [
|
||||
s for s in stocks if (not q) or q in s.symbol.upper() or q in (s.name or "").upper()
|
||||
][:15]
|
||||
if not rows:
|
||||
return "未找到匹配股票"
|
||||
return "\n".join(
|
||||
f"{s.symbol} {s.name} 行业={s.industry or '-'} 上市={s.list_date}" for s in rows
|
||||
)
|
||||
|
||||
def get_market_data(args: dict) -> str:
|
||||
symbol = str(_pick(args, "symbol", "")).upper()
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", date.today().isoformat())))
|
||||
with session_factory() as session:
|
||||
bars = daily_repo_f(session).get_range(symbol, start, end)
|
||||
if not bars:
|
||||
return f"{symbol} 在 {start}~{end} 无日线数据(可能未同步)"
|
||||
head, tail = bars[0], bars[-1]
|
||||
last = "\n".join(f"{b.trade_date} close={b.close}" for b in bars[-8:])
|
||||
change = float(tail.close) / float(head.close) - 1 if head.close and tail.close else None
|
||||
return (
|
||||
f"{symbol} {start}~{end} 共 {len(bars)} 根日线;"
|
||||
f"区间 {head.trade_date}→{tail.trade_date} 收盘 {head.close}→{tail.close}"
|
||||
f"(涨跌 {change * 100:.2f}% 若数据完整);最近 8 根:\n{last}"
|
||||
)
|
||||
|
||||
def _run_spec(spec: ResearchSpec, desc: str) -> str:
|
||||
job = submit_and_run(spec, factories=facts)
|
||||
if job.status != "success":
|
||||
return f"{desc} 执行失败:{job.error}"
|
||||
if spec.type == "backtest":
|
||||
result = BacktestResult.model_validate_json(job.result_json or "{}")
|
||||
s = result.summary
|
||||
return (
|
||||
f"回测完成(Experiment {job.experiment_id},代码版本 {_code_version(job, exp_repo_f)})。"
|
||||
f"总收益 {s.total_return_pct:.2f}%,年化 {s.annual_return_pct:.2f}%,"
|
||||
f"Sharpe {s.sharpe:.2f},最大回撤 {s.max_drawdown_pct:.2f}%,"
|
||||
f"交易 {s.total_trades} 笔,平均换手 {s.avg_turnover_pct:.1f}%。"
|
||||
f"未建模约束 {len(result.unimplemented)} 项(成本/涨跌停近似见实验详情)。"
|
||||
)
|
||||
report = FactorTestReport.model_validate_json(job.result_json or "{}")
|
||||
qs = ", ".join(f"Q{q.quantile + 1}: {q.return_pct:.2f}%" for q in report.quantile_returns)
|
||||
return (
|
||||
f"因子测试完成(Experiment {job.experiment_id})。IC {report.ic_mean:.4f},"
|
||||
f"RankIC {report.rank_ic_mean:.4f},ICIR {report.icir:.2f},正收益占比 "
|
||||
f"{report.positive_ratio_pct:.1f}%,样本 {report.sample_days} 日;分层未来收益 {qs}。"
|
||||
f"注意:单因子测试不代表策略有效,需结合稳健性分析。"
|
||||
)
|
||||
|
||||
def test_factor(args: dict) -> str:
|
||||
name = str(_pick(args, "name", ""))
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", "2024-12-31")))
|
||||
spec = ResearchSpec(
|
||||
type="factor_test",
|
||||
universe={"exclude_st": True, "min_listing_days": 0},
|
||||
factors=[{"name": name, "weight": 1.0}],
|
||||
selection={"top_n": 10},
|
||||
rebalance="monthly",
|
||||
period=(start, end),
|
||||
)
|
||||
return _run_spec(spec, f"因子 {name} 测试")
|
||||
|
||||
def run_backtest(args: dict) -> str:
|
||||
factor_names = [f.strip() for f in str(_pick(args, "factors", "momentum_60")).split(",")]
|
||||
top_n = int(_pick(args, "top_n", 5) or 5)
|
||||
rebalance = str(_pick(args, "rebalance", "monthly"))
|
||||
exclude_st = bool(_pick(args, "exclude_st", True))
|
||||
start = _day(str(_pick(args, "start", "2024-01-01")))
|
||||
end = _day(str(_pick(args, "end", "2024-12-31")))
|
||||
spec = ResearchSpec(
|
||||
type="backtest",
|
||||
universe={"exclude_st": exclude_st, "min_listing_days": 0},
|
||||
factors=[{"name": n, "weight": 1.0} for n in factor_names],
|
||||
selection={"top_n": top_n},
|
||||
rebalance=rebalance,
|
||||
period=(start, end),
|
||||
)
|
||||
return _run_spec(spec, "回测")
|
||||
|
||||
def get_experiment(args: dict) -> str:
|
||||
exp_id = str(_pick(args, "experiment_id", "")).upper()
|
||||
with session_factory() as session:
|
||||
exp = exp_repo_f(session).get(exp_id)
|
||||
if exp is None:
|
||||
return f"Experiment {exp_id} 不存在(可用列表:GET /api/experiments)"
|
||||
spec = json.loads(exp.spec_json)
|
||||
return (
|
||||
f"Experiment {exp.id} [{exp.kind}] 因子={[f['name'] for f in spec.get('factors', [])]} "
|
||||
f"区间={spec.get('period')} 调仓={spec.get('rebalance')};摘要:{exp.summary_text or '-'} "
|
||||
f"代码版本={exp.code_version or '-'} 创建={exp.created_at}"
|
||||
)
|
||||
|
||||
def compare_experiments(args: dict) -> str:
|
||||
ids = [
|
||||
x.strip().upper()
|
||||
for x in str(_pick(args, "experiment_ids", "")).split(",")
|
||||
if x.strip()
|
||||
]
|
||||
if not ids:
|
||||
return "请提供 experiment_ids(逗号分隔)"
|
||||
with session_factory() as session:
|
||||
repo = exp_repo_f(session)
|
||||
rows = [(i, repo.get(i)) for i in ids]
|
||||
out = []
|
||||
for exp_id, exp in rows:
|
||||
if exp is None:
|
||||
out.append(f"{exp_id}: 不存在")
|
||||
else:
|
||||
spec = json.loads(exp.spec_json)
|
||||
out.append(
|
||||
f"{exp.id}: 因子={[f['name'] for f in spec.get('factors', [])]} "
|
||||
f"区间={spec.get('period')} → {exp.summary_text or '-'}"
|
||||
)
|
||||
return "\n".join(out)
|
||||
|
||||
return [
|
||||
Tool(
|
||||
"search_stocks",
|
||||
"按代码或名称搜索股票,返回基础信息(只读)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"q": {"type": "string", "description": "代码或名称关键字"}},
|
||||
},
|
||||
search_stocks,
|
||||
),
|
||||
Tool(
|
||||
"get_market_data",
|
||||
"读取一只股票一段区间的日线行情摘要(只读,不复权)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"symbol": {"type": "string", "description": "如 600519.SH"},
|
||||
"start": {"type": "string", "description": "YYYY-MM-DD"},
|
||||
"end": {"type": "string", "description": "YYYY-MM-DD"},
|
||||
},
|
||||
"required": ["symbol"],
|
||||
},
|
||||
get_market_data,
|
||||
),
|
||||
Tool(
|
||||
"test_factor",
|
||||
"对单个因子做 IC/RankIC/分层测试并归档 Experiment",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "因子名(momentum_60 / volatility_20 等)",
|
||||
},
|
||||
"start": {"type": "string"},
|
||||
"end": {"type": "string"},
|
||||
},
|
||||
"required": ["name"],
|
||||
},
|
||||
test_factor,
|
||||
),
|
||||
Tool(
|
||||
"run_backtest",
|
||||
"运行 TopK 低频回测并归档 Experiment(成本/涨跌停近似建模)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"factors": {"type": "string", "description": "逗号分隔的因子名"},
|
||||
"top_n": {"type": "integer"},
|
||||
"rebalance": {"type": "string", "enum": ["monthly", "weekly"]},
|
||||
"exclude_st": {"type": "boolean"},
|
||||
"start": {"type": "string"},
|
||||
"end": {"type": "string"},
|
||||
},
|
||||
},
|
||||
run_backtest,
|
||||
),
|
||||
Tool(
|
||||
"get_experiment",
|
||||
"读取已归档实验的摘要",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"experiment_id": {"type": "string"}},
|
||||
"required": ["experiment_id"],
|
||||
},
|
||||
get_experiment,
|
||||
),
|
||||
Tool(
|
||||
"compare_experiments",
|
||||
"对比多个实验(因子/区间/收益摘要)",
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {"experiment_ids": {"type": "string"}},
|
||||
"required": ["experiment_ids"],
|
||||
},
|
||||
compare_experiments,
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _code_version(job, exp_repo_f) -> str:
|
||||
try:
|
||||
with default_factories()["session_factory"]() as session:
|
||||
exp = exp_repo_f(session).get(job.experiment_id or "")
|
||||
return exp.code_version or "-" if exp else "-"
|
||||
except Exception: # noqa: BLE001
|
||||
return "-"
|
||||
@@ -0,0 +1,40 @@
|
||||
"""AI Research Agent API(Phase 5)。
|
||||
|
||||
POST /api/agent/chat {message} → {reply, actions:[{tool,args,output}]}
|
||||
未配置 LLM Key 时返回 400 引导(不会崩溃)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.agent.llm import LLMClient, build_llm_from_settings
|
||||
from app.agent.service import AgentService
|
||||
|
||||
router = APIRouter(prefix="/agent", tags=["agent"])
|
||||
|
||||
|
||||
class AgentChatRequest(BaseModel):
|
||||
message: str = Field(min_length=1, max_length=2000)
|
||||
|
||||
|
||||
def _llm_or_raise() -> LLMClient:
|
||||
llm = build_llm_from_settings()
|
||||
if llm is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="未配置 LLM:请在根目录 .env 中设置 LLM_API_KEY(可选 LLM_BASE_URL / LLM_MODEL),"
|
||||
"参考 .env.example 与 AGENT.md §33",
|
||||
)
|
||||
return llm
|
||||
|
||||
|
||||
@router.post("/chat", summary="与 AI 研究助手对话(受控工具)")
|
||||
def agent_chat(
|
||||
body: AgentChatRequest,
|
||||
llm: Annotated[LLMClient, Depends(_llm_or_raise)],
|
||||
) -> dict:
|
||||
return AgentService(llm).chat(body.message)
|
||||
+4
-14
@@ -14,7 +14,7 @@ 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.api.deps import DbSession, JobRepoDep
|
||||
from app.application.services.job_executor import execute_job, new_id
|
||||
from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
@@ -24,27 +24,17 @@ from app.domain.entities.research import (
|
||||
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(),
|
||||
}
|
||||
from app.application.services.job_executor import default_factories
|
||||
|
||||
return default_factories()
|
||||
|
||||
|
||||
def _decode_result(kind: str, result_json: str | None):
|
||||
|
||||
@@ -8,7 +8,7 @@ from __future__ import annotations
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api import experiments, factors, health, jobs, research, stocks
|
||||
from app.api import agent, experiments, factors, health, jobs, research, stocks
|
||||
|
||||
api_router = APIRouter()
|
||||
api_router.include_router(health.router)
|
||||
@@ -17,3 +17,4 @@ api_router.include_router(factors.router)
|
||||
api_router.include_router(research.router)
|
||||
api_router.include_router(jobs.router)
|
||||
api_router.include_router(experiments.router)
|
||||
api_router.include_router(agent.router)
|
||||
|
||||
@@ -18,6 +18,7 @@ from app.domain.entities.research import (
|
||||
BacktestResult,
|
||||
ExperimentRecord,
|
||||
FactorTestReport,
|
||||
JobRecord,
|
||||
JobStatus,
|
||||
ResearchSpec,
|
||||
)
|
||||
@@ -144,3 +145,56 @@ def execute_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
|
||||
|
||||
@@ -68,6 +68,9 @@ class Settings:
|
||||
data_source_primary: str
|
||||
data_source_fallback: str
|
||||
tushare_token: str
|
||||
llm_api_key: str
|
||||
llm_base_url: str | None
|
||||
llm_model: str
|
||||
storage: dict[str, Path]
|
||||
config_path: Path = CONFIG_PATH
|
||||
env_path: Path = ENV_PATH
|
||||
@@ -102,6 +105,9 @@ def get_settings() -> Settings:
|
||||
secret_env = _deep(cfg, "app.secret_key_env") or "APP_SECRET_KEY"
|
||||
url_env = _deep(cfg, "database.url_env") or "DATABASE_URL"
|
||||
token_env = _deep(cfg, "data_source.tushare_token_env") or "TUSHARE_TOKEN"
|
||||
llm_key_env = _deep(cfg, "agent.llm_key_env") or "LLM_API_KEY"
|
||||
llm_url_env = _deep(cfg, "agent.llm_base_url_env") or "LLM_BASE_URL"
|
||||
llm_model_env = _deep(cfg, "agent.llm_model_env") or "LLM_MODEL"
|
||||
|
||||
default_url = "sqlite:///./data/quant.db"
|
||||
database_url = os.environ.get(url_env) or default_url
|
||||
@@ -122,5 +128,8 @@ def get_settings() -> Settings:
|
||||
data_source_primary=str(_deep(cfg, "data_source.primary") or "tushare"),
|
||||
data_source_fallback=str(_deep(cfg, "data_source.fallback") or "sina"),
|
||||
tushare_token=os.environ.get(token_env, ""),
|
||||
llm_api_key=os.environ.get(llm_key_env, ""),
|
||||
llm_base_url=os.environ.get(llm_url_env) or None,
|
||||
llm_model=os.environ.get(llm_model_env) or "qwen-plus",
|
||||
storage=_resolve_storage_dirs(cfg),
|
||||
)
|
||||
|
||||
@@ -10,6 +10,7 @@ dependencies = [
|
||||
"sqlalchemy>=2.0",
|
||||
"alembic>=1.13",
|
||||
"pyyaml>=6.0",
|
||||
"httpx>=0.28.1",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
|
||||
@@ -0,0 +1,135 @@
|
||||
"""Phase 5 Agent 测试:受控工具白名单 / 编排循环 / 异常输入(用 FakeLLM,不触网)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
from app.agent.service import AgentService
|
||||
from app.agent.tools_impl import build_tools
|
||||
from app.infrastructure.persistence.sqlalchemy.base import Base
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
class FakeLLM:
|
||||
"""按顺序返回预设决策的假 LLM。"""
|
||||
|
||||
def __init__(self, decisions: list[dict]) -> None:
|
||||
self._decisions = decisions
|
||||
self.calls = 0
|
||||
|
||||
def chat(self, messages: list[dict]) -> str:
|
||||
self.calls += 1
|
||||
return json.dumps(self._decisions.pop(0), ensure_ascii=False)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def tmp_factories(tmp_path) -> dict:
|
||||
"""注入 tmp SQLite 的默认工厂(隔离本机开发库)。"""
|
||||
engine = create_engine(f"sqlite:///{tmp_path / 'agent.db'}", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = lambda: Session(engine) # noqa: E731
|
||||
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
|
||||
SqlAlchemyExperimentRepository,
|
||||
SqlAlchemyJobRepository,
|
||||
)
|
||||
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
||||
SqlAlchemyDailyBarRepository,
|
||||
SqlAlchemyStockRepository,
|
||||
)
|
||||
from app.quant.engine import LocalEngine
|
||||
|
||||
return {
|
||||
"session_factory": session_factory,
|
||||
"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(),
|
||||
}
|
||||
|
||||
|
||||
class TestToolCatalog:
|
||||
def test_whitelist_and_metadata(self) -> None:
|
||||
tools = build_tools()
|
||||
names = {t.name for t in tools}
|
||||
assert {
|
||||
"search_stocks",
|
||||
"get_market_data",
|
||||
"test_factor",
|
||||
"run_backtest",
|
||||
"get_experiment",
|
||||
"compare_experiments",
|
||||
} <= names
|
||||
for t in tools:
|
||||
assert t.description
|
||||
assert t.parameters["type"] == "object"
|
||||
|
||||
def test_no_arbitrary_shell_tool(self) -> None:
|
||||
names = {t.name for t in build_tools()}
|
||||
assert "shell" not in names
|
||||
assert "execute" not in names
|
||||
|
||||
|
||||
class TestAgentOrchestration:
|
||||
def test_full_research_loop(self, tmp_factories) -> None:
|
||||
llm = FakeLLM(
|
||||
[
|
||||
{
|
||||
"tool": "test_factor",
|
||||
"args": {"name": "momentum_60", "start": "2024-03-01", "end": "2024-10-31"},
|
||||
},
|
||||
{
|
||||
"tool": "run_backtest",
|
||||
"args": {
|
||||
"factors": "momentum_60",
|
||||
"top_n": 3,
|
||||
"start": "2024-03-01",
|
||||
"end": "2024-10-31",
|
||||
},
|
||||
},
|
||||
{"final": "结论:单次样本内回测较好,但需要样本外验证,暂不判定策略有效。"},
|
||||
]
|
||||
)
|
||||
svc = AgentService(llm, build_tools(factories=tmp_factories))
|
||||
out = svc.chat("请测试 momentum_60 因子并跑一次回测,再给结论。")
|
||||
|
||||
assert [a["tool"] for a in out["actions"]] == ["test_factor", "run_backtest"]
|
||||
# 工具输出真实来自引擎/归档链路(空库也返回结构化文本)
|
||||
assert all("output" in a and a["output"] for a in out["actions"])
|
||||
assert "样本外" in out["reply"]
|
||||
|
||||
def test_tool_error_fed_back_without_abort(self, tmp_factories) -> None:
|
||||
llm = FakeLLM(
|
||||
[
|
||||
{"tool": "search_stocks", "args": {"q": "茅台"}},
|
||||
{"final": "完成"},
|
||||
]
|
||||
)
|
||||
svc = AgentService(llm, build_tools(factories=tmp_factories))
|
||||
out = svc.chat("查一下茅台")
|
||||
assert out["actions"][0]["tool"] == "search_stocks"
|
||||
|
||||
def test_unknown_tool_is_rejected(self) -> None:
|
||||
class RigidLLM:
|
||||
def chat(self, messages): # noqa: ARG002
|
||||
return json.dumps({"tool": "delete_database", "args": {}}, ensure_ascii=False)
|
||||
|
||||
svc = AgentService(RigidLLM(), build_tools())
|
||||
with pytest.raises(ValueError, match="工具不存在"):
|
||||
svc.chat("恶意调用")
|
||||
|
||||
def test_round_limit_forces_final(self, tmp_factories) -> None:
|
||||
# LLM 一直要求调工具 → 轮次耗尽强制 final
|
||||
class LoopLLM:
|
||||
def chat(self, messages): # noqa: ARG002
|
||||
return json.dumps(
|
||||
{"tool": "search_stocks", "args": {"q": "600519"}}, ensure_ascii=False
|
||||
)
|
||||
|
||||
svc = AgentService(LoopLLM(), build_tools(factories=tmp_factories))
|
||||
out = svc.chat("一直搜索")
|
||||
assert "actions" in out
|
||||
assert len(out["actions"]) <= 6
|
||||
Generated
+2
@@ -911,6 +911,7 @@ source = { virtual = "." }
|
||||
dependencies = [
|
||||
{ name = "alembic" },
|
||||
{ name = "fastapi" },
|
||||
{ name = "httpx" },
|
||||
{ name = "pyyaml" },
|
||||
{ name = "sqlalchemy" },
|
||||
{ name = "uvicorn", extra = ["standard"] },
|
||||
@@ -932,6 +933,7 @@ dev = [
|
||||
requires-dist = [
|
||||
{ name = "alembic", specifier = ">=1.13" },
|
||||
{ name = "fastapi", specifier = ">=0.115" },
|
||||
{ name = "httpx", specifier = ">=0.28.1" },
|
||||
{ name = "pyyaml", specifier = ">=6.0" },
|
||||
{ name = "sqlalchemy", specifier = ">=2.0" },
|
||||
{ name = "tushare", marker = "extra == 'datasource-tushare'", specifier = ">=1.4" },
|
||||
|
||||
+7
-7
@@ -8,14 +8,14 @@
|
||||
|
||||
## 0. 总览与里程碑
|
||||
|
||||
| 里程碑 | 内容 | 验收口径 |
|
||||
| 里程碑 | 内容 | 状态(2026-09 已实施) |
|
||||
|---|---|---|
|
||||
| M0 ✅ | 工程初始化:uv + FastAPI 骨架、分层包、SQLAlchemy/Alembic、CI 前检查(pytest/ruff) | `/api/health` 可用、测试通过、已推送远端 |
|
||||
| M1 | Phase 1 数据层:Tushare 拉取 → 标准化 → SQLite/Parquet | 命令行可全量/增量同步,来源可追溯 |
|
||||
| M2 | Phase 2 Qlib:Qlib Dataset 构建 + 因子 + LightGBM + 回测(Adapter 封装) | Research Specification 能驱动一次完整回测,产出标准结果 |
|
||||
| M3 | Phase 3 Web:股票池 / 因子 / 选股 / 回测 / 结果可视化 | 浏览器完成「选池→算因子→回测→看图」闭环 |
|
||||
| M4 | Phase 4 Experiment:研究全量可复现归档 + Job/SSE 异步化 | 任一历史实验可一键复跑 |
|
||||
| M5 | Phase 5 AI Agent:自然语言 → Research Plan → 受控 Tool → Experiment | Agent 能独立完成一次「因子假设→测试→结论」并留档 |
|
||||
| M0 ✅ | 工程初始化 | commit 2a52ee5 |
|
||||
| M1 ✅ | Phase 1 数据层(Tushare→SQLite/Parquet、Failover 审计、防未来函数) | commit 2da2342 · 38 tests |
|
||||
| M2 ✅ | Phase 2 研究引擎(Spec→因子→评估→低频回测→标准结果;Qlib 桥接占位见 §2 备注) | commit e9f59d3 · 60 tests |
|
||||
| M3 ✅ | Phase 3 Web(业务 API + Next.js 前端闭环) | commit 92627f5 · 68 tests |
|
||||
| M4 ✅ | Phase 4 Experiment 归档 + Job/SSE 异步(一键复跑) | commit 0ea229d · 73 tests |
|
||||
| M5 ✅ | Phase 5 AI Agent(受控 Tool 白名单 + LLM 编排;Key 配置见 .env) | commit(本轮)· 79 tests |
|
||||
|
||||
每阶段结束时同步更新:README / AGENT.md 相关清单 / 文档;**禁止跨阶段提前堆量**(AGENT.md §38)。
|
||||
|
||||
|
||||
Reference in New Issue
Block a user