301 lines
10 KiB
Python
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
|