"""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