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:
Simon
2026-09-06 17:22:05 +08:00
parent 0ea229d766
commit d9be75a98f
14 changed files with 722 additions and 23 deletions
+77
View File
@@ -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()))
+98
View File
@@ -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}
+45
View File
@@ -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
View File
@@ -1 +0,0 @@
"""Agent 受控工具集:每个 Tool 是后端服务的只读 / 沙箱化入口。"""
+248
View File
@@ -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 "-"
+40
View File
@@ -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
View File
@@ -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):
+2 -1
View File
@@ -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
+9
View File
@@ -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),
)
+1
View File
@@ -10,6 +10,7 @@ dependencies = [
"sqlalchemy>=2.0",
"alembic>=1.13",
"pyyaml>=6.0",
"httpx>=0.28.1",
]
[project.optional-dependencies]
+135
View File
@@ -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
+2
View File
@@ -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
View File
@@ -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)。