Files
group_xinghuo_jinrong/tests/test_step11_agent_graph.py
T

301 lines
10 KiB
Python

"""Step 11 测试:LangGraph 编排层 — 顾问 Agent StateGraph"""
import pytest
from unittest.mock import Mock
from langchain_core.messages import HumanMessage, AIMessage, ToolMessage
from langchain_core.tools import tool as _tool_decorator
from app.service.agent_graph import (
AgentState,
AdvisorAgent,
build_advisor_graph,
)
# ---------------------------------------------------------------------------
# 辅助:创建真实的 LangGraph 工具(用于测试)
# ---------------------------------------------------------------------------
def _make_test_tool(name="test_tool", return_value=None, side_effect=None):
"""创建一个真实的 @tool 装饰函数用于测试。"""
if side_effect is not None:
@_tool_decorator
def test_tool(input_text: str = "") -> str:
"""测试工具"""
raise side_effect
else:
result = return_value or {"result": "ok"}
@_tool_decorator
def test_tool(input_text: str = "") -> dict:
"""测试工具"""
return result
# 修改工具名称
test_tool.name = name
return test_tool
# ---------------------------------------------------------------------------
# AgentState 测试
# ---------------------------------------------------------------------------
class TestAgentState:
"""测试 Agent 状态定义"""
def test_state_keys(self):
"""状态应包含必要字段"""
state: AgentState = {
"messages": [],
"current_customer_id": None,
"conversation_history": [],
}
assert "messages" in state
assert "current_customer_id" in state
assert "conversation_history" in state
def test_state_type_annotations(self):
"""AgentState 应有正确的类型注解"""
hints = AgentState.__annotations__
assert "messages" in hints
assert "current_customer_id" in hints
assert "conversation_history" in hints
# ---------------------------------------------------------------------------
# AdvisorAgent 测试
# ---------------------------------------------------------------------------
class TestAdvisorAgent:
"""测试顾问 Agent 类"""
def test_init_with_tools_and_llm(self):
"""应能用工具和 LLM 客户端初始化"""
test_tool = _make_test_tool()
mock_llm = Mock()
agent = AdvisorAgent(tools=[test_tool], llm=mock_llm)
assert len(agent.tools) == 1
assert agent.llm == mock_llm
def test_build_graph_returns_compiled_graph(self):
"""build_graph 应返回编译后的图"""
test_tool = _make_test_tool()
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
agent = AdvisorAgent(tools=[test_tool], llm=mock_llm)
graph = agent.build_graph()
assert graph is not None
assert hasattr(graph, "invoke")
def test_build_graph_without_tools(self):
"""无工具时也能构建图"""
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
agent = AdvisorAgent(tools=[], llm=mock_llm)
graph = agent.build_graph()
assert graph is not None
# ---------------------------------------------------------------------------
# build_advisor_graph 工厂函数测试
# ---------------------------------------------------------------------------
class TestBuildAdvisorGraph:
"""测试构建顾问 Agent 图的工厂函数"""
def test_build_advisor_graph_with_dependencies(self):
"""应能用依赖项构建完整的图"""
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
graph = build_advisor_graph(
core_repo=Mock(),
kyc_service=Mock(),
template_service=Mock(),
compliance_service=Mock(),
llm=mock_llm,
)
assert graph is not None
assert hasattr(graph, "invoke")
# ---------------------------------------------------------------------------
# 图执行测试
# ---------------------------------------------------------------------------
class TestGraphExecution:
"""测试图的执行流程"""
def test_simple_conversation_without_tools(self):
"""简单对话应不调用工具直接返回"""
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
mock_llm.invoke = Mock(return_value=AIMessage(content="您好,我是您的投资顾问。"))
agent = AdvisorAgent(tools=[], llm=mock_llm)
graph = agent.build_graph()
result = graph.invoke({
"messages": [HumanMessage(content="你好")],
"current_customer_id": None,
"conversation_history": [],
})
# 用户输入 + AI 回复
assert len(result["messages"]) == 2
assert result["messages"][-1].content == "您好,我是您的投资顾问。"
def test_tool_calling_flow(self):
"""工具调用流程应正确执行"""
holdings_result = {"customer_id": "CUST-001", "holdings": []}
test_tool = _make_test_tool(name="query_holdings", return_value=holdings_result)
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
# 第一次 LLM 调用:决定调用工具
tool_call_message = AIMessage(
content="",
tool_calls=[{
"name": "query_holdings",
"args": {"input_text": "CUST-001"},
"id": "call_1",
}]
)
# 第二次 LLM 调用:基于工具结果生成回复
final_message = AIMessage(content="客户 CUST-001 目前没有持仓。")
mock_llm.invoke = Mock(side_effect=[tool_call_message, final_message])
agent = AdvisorAgent(tools=[test_tool], llm=mock_llm)
graph = agent.build_graph()
result = graph.invoke({
"messages": [HumanMessage(content="查一下客户 CUST-001 的持仓")],
"current_customer_id": None,
"conversation_history": [],
})
# 用户输入 + AI 工具调用 + 工具结果 + AI 最终回复
assert len(result["messages"]) == 4
assert isinstance(result["messages"][2], ToolMessage)
def test_error_handling_in_tool_execution(self):
"""工具执行错误应被正确抛出"""
test_tool = _make_test_tool(
name="failing_tool",
side_effect=Exception("工具执行失败"),
)
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
tool_call_message = AIMessage(
content="",
tool_calls=[{
"name": "failing_tool",
"args": {"input_text": ""},
"id": "call_1",
}]
)
mock_llm.invoke = Mock(return_value=tool_call_message)
agent = AdvisorAgent(tools=[test_tool], llm=mock_llm)
graph = agent.build_graph()
# LangGraph ToolNode 默认不捕获异常,验证异常被正确抛出
with pytest.raises(Exception, match="工具执行失败"):
graph.invoke({
"messages": [HumanMessage(content="测试错误处理")],
"current_customer_id": None,
"conversation_history": [],
})
def test_conversation_history_is_preserved(self):
"""对话历史应被保留"""
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
mock_llm.invoke = Mock(return_value=AIMessage(content="好的,我记住了。"))
history = [
{"role": "user", "content": "第一轮对话"},
{"role": "assistant", "content": "第一轮回复"},
]
agent = AdvisorAgent(tools=[], llm=mock_llm)
graph = agent.build_graph()
result = graph.invoke({
"messages": [HumanMessage(content="继续对话")],
"current_customer_id": None,
"conversation_history": history,
})
assert "conversation_history" in result
assert len(result["conversation_history"]) == 2
def test_current_customer_id_is_tracked(self):
"""当前客户 ID 应被追踪"""
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
mock_llm.invoke = Mock(return_value=AIMessage(content="好的。"))
agent = AdvisorAgent(tools=[], llm=mock_llm)
graph = agent.build_graph()
result = graph.invoke({
"messages": [HumanMessage(content="查看客户 CUST-001")],
"current_customer_id": "CUST-001",
"conversation_history": [],
})
assert result["current_customer_id"] == "CUST-001"
# ---------------------------------------------------------------------------
# 图配置测试
# ---------------------------------------------------------------------------
class TestGraphConfiguration:
"""测试图的配置选项"""
def test_max_iterations_prevents_infinite_loops(self):
"""max_iterations 应防止无限循环"""
test_tool = _make_test_tool(name="loop_tool", return_value={"result": "ok"})
mock_llm = Mock()
mock_llm.bind_tools = Mock(return_value=mock_llm)
# 总是返回工具调用
tool_call_message = AIMessage(
content="",
tool_calls=[{
"name": "loop_tool",
"args": {"input_text": ""},
"id": "call_1",
}]
)
mock_llm.invoke = Mock(return_value=tool_call_message)
agent = AdvisorAgent(tools=[test_tool], llm=mock_llm, max_iterations=3)
graph = agent.build_graph()
result = graph.invoke({
"messages": [HumanMessage(content="触发循环")],
"current_customer_id": None,
"conversation_history": [],
})
# 应该在 max_iterations 后停止
# 每次迭代产生 1 条 AI 消息 + 1 条工具消息 = 2 条
# 加上初始 1 条用户消息,总共不超过 1 + 3*2 = 7 条
assert len(result["messages"]) <= 1 + 3 * 2
def test_default_max_iterations(self):
"""默认 max_iterations 应为 10"""
mock_llm = Mock()
agent = AdvisorAgent(tools=[], llm=mock_llm)
assert agent.max_iterations == 10