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 import APIRouter, BackgroundTasks, HTTPException
from fastapi.responses import StreamingResponse 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.application.services.job_executor import execute_job, new_id
from app.domain.entities.research import ( from app.domain.entities.research import (
BacktestResult, BacktestResult,
@@ -24,27 +24,17 @@ from app.domain.entities.research import (
ResearchSpec, ResearchSpec,
) )
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import ( from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
SqlAlchemyExperimentRepository,
SqlAlchemyJobRepository, SqlAlchemyJobRepository,
) )
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
from app.infrastructure.persistence.sqlalchemy.session import SessionLocal from app.infrastructure.persistence.sqlalchemy.session import SessionLocal
router = APIRouter(prefix="/jobs", tags=["jobs"]) router = APIRouter(prefix="/jobs", tags=["jobs"])
def _bg_factories(): def _bg_factories():
return { from app.application.services.job_executor import default_factories
"session_factory": SessionLocal,
"job_repo_factory": lambda s: SqlAlchemyJobRepository(s), return default_factories()
"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): 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 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 = APIRouter()
api_router.include_router(health.router) 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(research.router)
api_router.include_router(jobs.router) api_router.include_router(jobs.router)
api_router.include_router(experiments.router) api_router.include_router(experiments.router)
api_router.include_router(agent.router)
@@ -18,6 +18,7 @@ from app.domain.entities.research import (
BacktestResult, BacktestResult,
ExperimentRecord, ExperimentRecord,
FactorTestReport, FactorTestReport,
JobRecord,
JobStatus, JobStatus,
ResearchSpec, ResearchSpec,
) )
@@ -144,3 +145,56 @@ def execute_job(
session.commit() session.commit()
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
pass 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_primary: str
data_source_fallback: str data_source_fallback: str
tushare_token: str tushare_token: str
llm_api_key: str
llm_base_url: str | None
llm_model: str
storage: dict[str, Path] storage: dict[str, Path]
config_path: Path = CONFIG_PATH config_path: Path = CONFIG_PATH
env_path: Path = ENV_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" secret_env = _deep(cfg, "app.secret_key_env") or "APP_SECRET_KEY"
url_env = _deep(cfg, "database.url_env") or "DATABASE_URL" url_env = _deep(cfg, "database.url_env") or "DATABASE_URL"
token_env = _deep(cfg, "data_source.tushare_token_env") or "TUSHARE_TOKEN" 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" default_url = "sqlite:///./data/quant.db"
database_url = os.environ.get(url_env) or default_url 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_primary=str(_deep(cfg, "data_source.primary") or "tushare"),
data_source_fallback=str(_deep(cfg, "data_source.fallback") or "sina"), data_source_fallback=str(_deep(cfg, "data_source.fallback") or "sina"),
tushare_token=os.environ.get(token_env, ""), 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), storage=_resolve_storage_dirs(cfg),
) )
+1
View File
@@ -10,6 +10,7 @@ dependencies = [
"sqlalchemy>=2.0", "sqlalchemy>=2.0",
"alembic>=1.13", "alembic>=1.13",
"pyyaml>=6.0", "pyyaml>=6.0",
"httpx>=0.28.1",
] ]
[project.optional-dependencies] [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 = [ dependencies = [
{ name = "alembic" }, { name = "alembic" },
{ name = "fastapi" }, { name = "fastapi" },
{ name = "httpx" },
{ name = "pyyaml" }, { name = "pyyaml" },
{ name = "sqlalchemy" }, { name = "sqlalchemy" },
{ name = "uvicorn", extra = ["standard"] }, { name = "uvicorn", extra = ["standard"] },
@@ -932,6 +933,7 @@ dev = [
requires-dist = [ requires-dist = [
{ name = "alembic", specifier = ">=1.13" }, { name = "alembic", specifier = ">=1.13" },
{ name = "fastapi", specifier = ">=0.115" }, { name = "fastapi", specifier = ">=0.115" },
{ name = "httpx", specifier = ">=0.28.1" },
{ name = "pyyaml", specifier = ">=6.0" }, { name = "pyyaml", specifier = ">=6.0" },
{ name = "sqlalchemy", specifier = ">=2.0" }, { name = "sqlalchemy", specifier = ">=2.0" },
{ name = "tushare", marker = "extra == 'datasource-tushare'", specifier = ">=1.4" }, { name = "tushare", marker = "extra == 'datasource-tushare'", specifier = ">=1.4" },
+7 -7
View File
@@ -8,14 +8,14 @@
## 0. 总览与里程碑 ## 0. 总览与里程碑
| 里程碑 | 内容 | 验收口径 | | 里程碑 | 内容 | 状态(2026-09 已实施) |
|---|---|---| |---|---|---|
| M0 ✅ | 工程初始化:uv + FastAPI 骨架、分层包、SQLAlchemy/Alembic、CI 前检查(pytest/ruff) | `/api/health` 可用、测试通过、已推送远端 | | M0 ✅ | 工程初始化 | commit 2a52ee5 |
| M1 | Phase 1 数据层:Tushare 拉取 → 标准化 → SQLite/Parquet | 命令行可全量/增量同步,来源可追溯 | | M1 ✅ | Phase 1 数据层(Tushare→SQLite/Parquet、Failover 审计、防未来函数) | commit 2da2342 · 38 tests |
| M2 | Phase 2 Qlib:Qlib Dataset 构建 + 因子 + LightGBM + 回测(Adapter 封装) | Research Specification 能驱动一次完整回测,产出标准结果 | | M2 ✅ | Phase 2 研究引擎(Spec→因子→评估→低频回测→标准结果;Qlib 桥接占位见 §2 备注) | commit e9f59d3 · 60 tests |
| M3 | Phase 3 Web:股票池 / 因子 / 选股 / 回测 / 结果可视化 | 浏览器完成「选池→算因子→回测→看图」闭环 | | M3 ✅ | Phase 3 Web(业务 API + Next.js 前端闭环) | commit 92627f5 · 68 tests |
| M4 | Phase 4 Experiment:研究全量可复现归档 + Job/SSE 异步化 | 任一历史实验可一键复跑 | | M4 ✅ | Phase 4 Experiment 归档 + Job/SSE 异步(一键复跑) | commit 0ea229d · 73 tests |
| M5 | Phase 5 AI Agent:自然语言 → Research Plan → 受控 Tool → Experiment | Agent 能独立完成一次「因子假设→测试→结论」并留档 | | M5 ✅ | Phase 5 AI Agent(受控 Tool 白名单 + LLM 编排;Key 配置见 .env) | commit(本轮)· 79 tests |
每阶段结束时同步更新:README / AGENT.md 相关清单 / 文档;**禁止跨阶段提前堆量**(AGENT.md §38)。 每阶段结束时同步更新:README / AGENT.md 相关清单 / 文档;**禁止跨阶段提前堆量**(AGENT.md §38)。