- 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
136 lines
4.9 KiB
Python
136 lines
4.9 KiB
Python
"""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
|