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