200 lines
7.4 KiB
Python
200 lines
7.4 KiB
Python
#!/usr/bin/env python3
|
|
"""
|
|
LangChain + LangGraph 脚手架基础测试
|
|
"""
|
|
|
|
import unittest
|
|
import sys
|
|
import os
|
|
from unittest.mock import patch
|
|
from typing import cast
|
|
from langchain_core.messages import AIMessage
|
|
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
|
|
|
config_path = os.path.join(os.path.dirname(__file__), '..', 'config', 'config.ini')
|
|
if not os.path.exists(config_path):
|
|
with open(config_path, 'w') as f:
|
|
f.write("""
|
|
[General]
|
|
DEFAULT_MODEL_SECTION = gpt-4o
|
|
MAX_RETRIES = 1
|
|
TIMEOUT = 10
|
|
|
|
[gpt-4o]
|
|
MODEL_NAME = gpt-4o
|
|
OPENAI_API_KEY = your_openai_api_key_here
|
|
|
|
[gpt-3.5-turbo]
|
|
MODEL_NAME = gpt-3.5-turbo
|
|
OPENAI_API_KEY = your_openai_api_key_here
|
|
""")
|
|
|
|
from workflows.workflow_manager import WorkflowManager, WorkflowType
|
|
from services.common.app_errors import AppError, ErrorCode
|
|
|
|
|
|
class FakeModel:
|
|
def invoke(self, messages):
|
|
if len(messages) == 2 and getattr(messages[0], 'content', '').startswith('You are a translation and normalization assistant'):
|
|
return AIMessage(content=messages[1].content)
|
|
return AIMessage(content='fallback')
|
|
|
|
|
|
class EmptyTemplateMatcher:
|
|
def match(self, normalized_text: str):
|
|
return {
|
|
'table_name': None,
|
|
'candidates': [],
|
|
'raw': {'query': normalized_text},
|
|
}
|
|
|
|
|
|
class TestWorkflowManager(unittest.TestCase):
|
|
"""测试 WorkflowManager 功能"""
|
|
|
|
def setUp(self):
|
|
"""设置测试夹具"""
|
|
self.model_patcher = patch("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
|
|
self.matcher_patcher = patch("agent.core.nodes.get_template_matcher", lambda: EmptyTemplateMatcher())
|
|
self.model_patcher.start()
|
|
self.matcher_patcher.start()
|
|
self.manager = WorkflowManager(enable_multi_turn=False)
|
|
|
|
def tearDown(self):
|
|
self.matcher_patcher.stop()
|
|
self.model_patcher.stop()
|
|
|
|
def test_get_available_workflows(self):
|
|
"""测试可用工作流返回"""
|
|
workflows = self.manager.get_available_workflows()
|
|
self.assertIsInstance(workflows, list)
|
|
self.assertGreater(len(workflows), 0)
|
|
self.assertIn("conversation", workflows)
|
|
self.assertIn("tool_using", workflows)
|
|
|
|
def test_get_workflow(self):
|
|
"""测试获取工作流实例"""
|
|
conversation_workflow = self.manager.get_workflow(WorkflowType.CONVERSATION)
|
|
self.assertIsNotNone(conversation_workflow)
|
|
|
|
tool_workflow = self.manager.get_workflow(WorkflowType.TOOL_USING)
|
|
self.assertIsNotNone(tool_workflow)
|
|
|
|
def test_session_management(self):
|
|
"""测试会话创建与获取"""
|
|
# 执行工作流以创建会话
|
|
result = self.manager.execute_workflow(
|
|
WorkflowType.CONVERSATION,
|
|
"Hello, test session"
|
|
)
|
|
|
|
session_id = result["session_id"]
|
|
self.assertIsNotNone(session_id)
|
|
|
|
# 测试会话信息获取
|
|
session_info = self.manager.get_session_info(session_id)
|
|
self.assertIsNotNone(session_info)
|
|
self.assertEqual(session_info["workflow_type"], WorkflowType.CONVERSATION)
|
|
|
|
def test_conversation_session_is_temporarily_stateless(self):
|
|
"""测试关闭多轮后,同一 session_id 也不会自动累积对话历史"""
|
|
first = self.manager.execute_workflow(
|
|
WorkflowType.CONVERSATION,
|
|
"Hello, first turn",
|
|
session_id="session-1",
|
|
)
|
|
second = self.manager.execute_workflow(
|
|
WorkflowType.CONVERSATION,
|
|
"Hello, second turn",
|
|
session_id="session-1",
|
|
)
|
|
other = self.manager.execute_workflow(
|
|
WorkflowType.CONVERSATION,
|
|
"Hello, other session",
|
|
session_id="session-2",
|
|
)
|
|
|
|
session_one = self.manager.get_session_info("session-1")
|
|
session_two = self.manager.get_session_info("session-2")
|
|
|
|
self.assertEqual(first["session_id"], "session-1")
|
|
self.assertEqual(second["session_id"], "session-1")
|
|
self.assertEqual(other["session_id"], "session-2")
|
|
self.assertNotIn("conversation_history", session_one)
|
|
self.assertNotIn("last_context", session_one)
|
|
self.assertNotIn("conversation_history", session_two)
|
|
self.assertNotIn("last_context", session_two)
|
|
self.assertEqual(len(first["result"].get("conversation_history") or []), 2)
|
|
self.assertEqual(len(second["result"].get("conversation_history") or []), 2)
|
|
self.assertEqual(len(other["result"].get("conversation_history") or []), 2)
|
|
|
|
def test_conversation_session_memory_can_be_enabled(self):
|
|
"""测试开启多轮后,同一 session_id 会保存并复用会话级历史"""
|
|
manager = WorkflowManager(enable_multi_turn=True)
|
|
|
|
first = manager.execute_workflow(
|
|
WorkflowType.CONVERSATION,
|
|
"Hello, first turn",
|
|
session_id="session-enabled",
|
|
)
|
|
second = manager.execute_workflow(
|
|
WorkflowType.CONVERSATION,
|
|
"Hello, second turn",
|
|
session_id="session-enabled",
|
|
)
|
|
|
|
session_info = manager.get_session_info("session-enabled")
|
|
|
|
self.assertIn("conversation_history", session_info)
|
|
self.assertIn("last_context", session_info)
|
|
self.assertGreaterEqual(len(session_info["conversation_history"]), 4)
|
|
self.assertEqual(len(first["result"].get("conversation_history") or []), 2)
|
|
self.assertGreaterEqual(len(second["result"].get("conversation_history") or []), 4)
|
|
|
|
def test_reusing_session_id_with_different_workflow_raises(self):
|
|
"""测试同一个 session_id 不能绑定到不同工作流"""
|
|
self.manager.execute_workflow(
|
|
WorkflowType.CONVERSATION,
|
|
"Hello, test session",
|
|
session_id="shared-session",
|
|
)
|
|
|
|
with self.assertRaises(ValueError):
|
|
self.manager.execute_workflow(
|
|
WorkflowType.TOOL_USING,
|
|
"2 + 2",
|
|
session_id="shared-session",
|
|
)
|
|
|
|
def test_execute_workflow_rejects_none_user_input(self):
|
|
"""测试空 query 会在进入 Agent 之前被拒绝"""
|
|
with self.assertRaises(AppError) as ctx:
|
|
self.manager.execute_workflow(WorkflowType.CONVERSATION, cast(str, None))
|
|
|
|
self.assertEqual(ctx.exception.code, ErrorCode.INVALID_REQUEST)
|
|
self.assertEqual(ctx.exception.status_code, 400)
|
|
self.assertEqual(ctx.exception.detail, {"field": "user_input", "reason": "missing_or_blank"})
|
|
|
|
def test_execute_workflow_rejects_blank_user_input(self):
|
|
"""测试全空白 query 会在进入 Agent 之前被拒绝"""
|
|
with self.assertRaises(AppError) as ctx:
|
|
self.manager.execute_workflow(WorkflowType.CONVERSATION, " ")
|
|
|
|
self.assertEqual(ctx.exception.code, ErrorCode.INVALID_REQUEST)
|
|
self.assertEqual(ctx.exception.status_code, 400)
|
|
self.assertEqual(ctx.exception.detail, {"field": "user_input", "reason": "missing_or_blank"})
|
|
|
|
|
|
class TestConfiguration(unittest.TestCase):
|
|
"""测试配置校验"""
|
|
|
|
def test_config_loading(self):
|
|
"""测试能从 config.ini 加载配置"""
|
|
from config import Config
|
|
model_config = Config.get_model_config()
|
|
self.assertIn('model', model_config)
|
|
self.assertIn('api_key', model_config)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main() |