init
This commit is contained in:
@@ -0,0 +1,12 @@
|
||||
# Tests 模块
|
||||
|
||||
## 作用
|
||||
|
||||
维护项目自动化测试,覆盖工作流、API 与脚本行为。
|
||||
|
||||
## 文件
|
||||
|
||||
- `test_basic.py`:基础可用性测试
|
||||
- `test_endpoints.py`:接口行为测试
|
||||
- `test_sql_workflow_refactor.py`:SQL 流程关键逻辑测试
|
||||
- `test_console_chat.py` / `test_demo_chat.py`:脚本相关测试
|
||||
@@ -0,0 +1 @@
|
||||
"""测试模块"""
|
||||
@@ -0,0 +1,51 @@
|
||||
"""
|
||||
pytest 配置文件
|
||||
|
||||
提供全局的 fixtures 和配置
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
# 添加项目根目录到 Python 路径
|
||||
project_root = Path(__file__).parent.parent
|
||||
sys.path.insert(0, str(project_root))
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def project_dir() -> Path:
|
||||
"""获取项目根目录"""
|
||||
return project_root
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def config_dir() -> Path:
|
||||
"""获取配置目录"""
|
||||
return project_root / "config"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_user_input() -> str:
|
||||
"""示例用户输入"""
|
||||
return "你好,帮我查询订单信息"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_sql() -> str:
|
||||
"""示例 SQL 语句"""
|
||||
return "SELECT * FROM orders LIMIT 10"
|
||||
|
||||
|
||||
# 自动使用的 fixture(可选)
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup_environment():
|
||||
"""为所有测试设置环境变量"""
|
||||
# 可以在这里设置测试环境变量
|
||||
os.environ.setdefault("TESTING", "true")
|
||||
yield
|
||||
# 清理(如果需要)
|
||||
if "TESTING" in os.environ:
|
||||
del os.environ["TESTING"]
|
||||
+118
-2
@@ -6,6 +6,9 @@ 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')
|
||||
@@ -27,6 +30,23 @@ 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):
|
||||
@@ -34,8 +54,16 @@ class TestWorkflowManager(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
"""设置测试夹具"""
|
||||
self.manager = WorkflowManager()
|
||||
|
||||
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()
|
||||
@@ -68,6 +96,94 @@ class TestWorkflowManager(unittest.TestCase):
|
||||
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):
|
||||
"""测试配置校验"""
|
||||
|
||||
@@ -0,0 +1,99 @@
|
||||
from schemas.agent_input import AgentInput
|
||||
from schemas.chat_message_request import ChatMessageRequestDTO
|
||||
from schemas.stream_input import StreamInputDTO
|
||||
|
||||
|
||||
def test_chat_message_request_accepts_frontend_dto_shape():
|
||||
payload = ChatMessageRequestDTO(
|
||||
query="查询 SO 4020438779 的 eta 信息",
|
||||
inputs={"region": "ANZ", "filters": ["top10"]},
|
||||
response_mode="streaming",
|
||||
user="tester-001",
|
||||
conversation_id="cid-123",
|
||||
files=[{"type": "image", "transfer_method": "remote_url", "url": "https://example.com/a.png"}],
|
||||
)
|
||||
|
||||
assert payload.query == "查询 SO 4020438779 的 eta 信息"
|
||||
assert payload.inputs == {"region": "ANZ", "filters": ["top10"]}
|
||||
assert payload.response_mode == "streaming"
|
||||
assert payload.user == "tester-001"
|
||||
assert payload.conversation_id == "cid-123"
|
||||
assert payload.files[0].model_dump()["type"] == "image"
|
||||
|
||||
|
||||
|
||||
def test_chat_message_request_defaults_inputs_and_files():
|
||||
payload = ChatMessageRequestDTO(
|
||||
query="hello",
|
||||
response_mode="blocking",
|
||||
user="tester-002",
|
||||
)
|
||||
|
||||
assert payload.inputs == {}
|
||||
assert payload.files == []
|
||||
assert payload.conversation_id is None
|
||||
|
||||
|
||||
|
||||
def test_chat_message_request_accepts_non_dict_inputs_object():
|
||||
payload = ChatMessageRequestDTO(
|
||||
query="hello",
|
||||
inputs=[{"name": "foo"}],
|
||||
response_mode="streaming",
|
||||
user="tester-003",
|
||||
)
|
||||
|
||||
assert payload.inputs == [{"name": "foo"}]
|
||||
|
||||
|
||||
|
||||
def test_chat_message_request_allows_nullable_java_dto_fields():
|
||||
payload = ChatMessageRequestDTO()
|
||||
|
||||
assert payload.query is None
|
||||
assert payload.response_mode is None
|
||||
assert payload.user is None
|
||||
assert payload.conversation_id is None
|
||||
assert payload.inputs == {}
|
||||
assert payload.files == []
|
||||
|
||||
|
||||
def test_chat_message_request_accepts_legacy_auto_generate_name_field():
|
||||
payload = ChatMessageRequestDTO(
|
||||
query="hello",
|
||||
response_mode="streaming",
|
||||
user="tester-legacy",
|
||||
auto_generate_name=True,
|
||||
)
|
||||
|
||||
assert payload.auto_generate_name is True
|
||||
assert payload.query == "hello"
|
||||
|
||||
|
||||
def test_agent_input_uses_same_schema_as_chat_message_request():
|
||||
payload = AgentInput(
|
||||
query="hello",
|
||||
response_mode="blocking",
|
||||
user="tester-004",
|
||||
inputs={"k": "v"},
|
||||
files=[{"type": "text"}],
|
||||
)
|
||||
|
||||
assert isinstance(payload, ChatMessageRequestDTO)
|
||||
assert payload.inputs == {"k": "v"}
|
||||
assert payload.files[0].model_dump()["type"] == "text"
|
||||
|
||||
|
||||
def test_stream_input_uses_same_schema_as_chat_message_request():
|
||||
payload = StreamInputDTO(
|
||||
query="hello",
|
||||
response_mode="streaming",
|
||||
user="tester-005",
|
||||
inputs={"region": "ANZ"},
|
||||
)
|
||||
|
||||
assert isinstance(payload, ChatMessageRequestDTO)
|
||||
assert payload.response_mode == "streaming"
|
||||
assert payload.user == "tester-005"
|
||||
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
from schemas.chat_message_response import ChatMessageResponseDTO
|
||||
|
||||
|
||||
EXPECTED_KEYS = [
|
||||
"id",
|
||||
"event",
|
||||
"task_id",
|
||||
"message_id",
|
||||
"conversation_id",
|
||||
"answer",
|
||||
"created_at",
|
||||
]
|
||||
|
||||
|
||||
def test_chat_message_response_matches_java_dto_shape():
|
||||
dto = ChatMessageResponseDTO(
|
||||
id="id-1",
|
||||
task_id="task-1",
|
||||
message_id="msg-1",
|
||||
conversation_id="cid-1",
|
||||
answer="hello",
|
||||
created_at=1705395332,
|
||||
)
|
||||
|
||||
dumped = dto.model_dump()
|
||||
|
||||
assert list(dumped.keys()) == EXPECTED_KEYS
|
||||
assert dumped["event"] == "message"
|
||||
assert isinstance(dumped["created_at"], int)
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
from scripts.console_chat import format_result, main, run_turn
|
||||
|
||||
|
||||
class FakeMessage:
|
||||
def __init__(self, content: str):
|
||||
self.content = content
|
||||
|
||||
|
||||
class FakeAgent:
|
||||
def __init__(self, model_section=None):
|
||||
self.model_section = model_section
|
||||
self.calls = []
|
||||
|
||||
def run(self, query, **kwargs):
|
||||
self.calls.append((query, kwargs))
|
||||
return {
|
||||
"messages": [FakeMessage("mock answer")],
|
||||
"context": {
|
||||
"table_name": "apbo_eta_ful",
|
||||
"query_mode": "detail",
|
||||
"final_sql": "SELECT service_order_id FROM dwd_ai.apbo_eta_ful",
|
||||
"sql_plan": {"selected_table": "apbo_eta_ful", "query_mode": "detail"},
|
||||
},
|
||||
"final_step": "response_generated",
|
||||
}
|
||||
|
||||
|
||||
def test_format_result_includes_optional_blocks():
|
||||
result = {
|
||||
"messages": [FakeMessage("hello")],
|
||||
"context": {
|
||||
"final_sql": "SELECT 1",
|
||||
"sql_plan": {"mode": "detail"},
|
||||
"foo": "bar",
|
||||
},
|
||||
}
|
||||
|
||||
text = format_result(result, show_sql=True, show_context=True, show_plan=True)
|
||||
assert "Answer:" in text
|
||||
assert "SQL:" in text
|
||||
assert "SQL Plan:" in text
|
||||
assert "Context:" in text
|
||||
|
||||
|
||||
def test_format_result_renders_table_from_wrapped_sr_api_result():
|
||||
wrapped = json.dumps(
|
||||
{
|
||||
"status_code": 200,
|
||||
"text": json.dumps(
|
||||
{
|
||||
"data": [
|
||||
{"service_order_id": "4020438779", "ship_to_country": "VN"},
|
||||
{"service_order_id": "4020438780", "ship_to_country": "PH"},
|
||||
]
|
||||
},
|
||||
ensure_ascii=False,
|
||||
),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
result = {
|
||||
"messages": [FakeMessage('{"status_code": 200, "text": "..."}')],
|
||||
"context": {
|
||||
"sr_api_result": wrapped,
|
||||
"final_sql": "SELECT service_order_id, ship_to_country FROM dwd_ai.apbo_eta_ful",
|
||||
},
|
||||
}
|
||||
|
||||
text = format_result(result)
|
||||
assert "Query Result: 2 row(s)" in text
|
||||
assert "status=200" in text
|
||||
assert "service_order_id" in text
|
||||
assert "ship_to_country" in text
|
||||
assert "4020438779" in text
|
||||
assert "VN" in text
|
||||
|
||||
|
||||
def test_format_result_renders_columns_and_rows_payload():
|
||||
result = {
|
||||
"messages": [FakeMessage("ok")],
|
||||
"context": {
|
||||
"sr_api_result": {
|
||||
"status_code": 200,
|
||||
"text": {
|
||||
"columns": ["region", "qty"],
|
||||
"rows": [["ANZ", 12], ["CAP", 8]],
|
||||
},
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
text = format_result(result)
|
||||
assert "Query Result: 2 row(s)" in text
|
||||
assert "status=200" in text
|
||||
assert "region" in text
|
||||
assert "qty" in text
|
||||
assert "ANZ" in text
|
||||
assert "12" in text
|
||||
|
||||
|
||||
def test_format_result_falls_back_for_non_tabular_error():
|
||||
result = {
|
||||
"messages": [FakeMessage("请求失败: timeout")],
|
||||
"context": {"sr_api_result": "请求失败: timeout"},
|
||||
}
|
||||
|
||||
text = format_result(result)
|
||||
assert "Answer:" in text
|
||||
assert "请求失败: timeout" in text
|
||||
assert "Query Result:" not in text
|
||||
|
||||
|
||||
def test_main_one_shot_success(capsys):
|
||||
with patch("scripts.console_chat.Config.validate_config", return_value=None), \
|
||||
patch("scripts.console_chat.ConversationAgent", FakeAgent):
|
||||
exit_code = main([
|
||||
"--query",
|
||||
"查询 SO 4020438779 的 eta 信息",
|
||||
"--skip-sr-api",
|
||||
"--show-sql",
|
||||
])
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert exit_code == 0
|
||||
assert "mock answer" in captured.out
|
||||
assert "SELECT service_order_id FROM dwd_ai.apbo_eta_ful" in captured.out
|
||||
|
||||
|
||||
def test_run_turn_enables_debug_node_trace(capsys):
|
||||
agent = FakeAgent()
|
||||
|
||||
run_turn(
|
||||
agent,
|
||||
"查询 SO 4020438779 的 eta 信息",
|
||||
user="tester",
|
||||
conversation_id="cid-1",
|
||||
skip_sr_api=True,
|
||||
show_sql=False,
|
||||
show_context=False,
|
||||
show_plan=False,
|
||||
)
|
||||
|
||||
_, kwargs = agent.calls[-1]
|
||||
assert kwargs["debug_node_trace"] is True
|
||||
|
||||
|
||||
def test_main_config_error(capsys):
|
||||
with patch("scripts.console_chat.Config.validate_config", side_effect=ValueError("bad config")):
|
||||
exit_code = main(["--query", "hello"])
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert exit_code == 1
|
||||
assert "Configuration error" in captured.err
|
||||
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from api.dependencies import get_workflow_manager
|
||||
from api.endpoints import router
|
||||
from workflows.workflow_manager import WorkflowType
|
||||
|
||||
|
||||
class StubWorkflowManager:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
def execute_workflow(self, workflow_type, user_input, session_id=None, **kwargs):
|
||||
self.calls.append(
|
||||
{
|
||||
"workflow_type": workflow_type,
|
||||
"user_input": user_input,
|
||||
"session_id": session_id,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
)
|
||||
return {
|
||||
"session_id": session_id,
|
||||
"workflow_type": workflow_type.value if isinstance(workflow_type, WorkflowType) else str(workflow_type),
|
||||
"result": {"ok": True, "context": {"final_sql": "SELECT 1"}},
|
||||
}
|
||||
|
||||
|
||||
class StubMessageStorage:
|
||||
def __init__(self):
|
||||
self.enabled = True
|
||||
self.created = []
|
||||
self.updated = []
|
||||
self.saved = []
|
||||
self.existing = {}
|
||||
self.create_should_fail = False
|
||||
self.update_should_fail = False
|
||||
|
||||
def create_conversation(self, conversation_id, user, name, status, introduction, created_at, updated_at):
|
||||
if self.create_should_fail:
|
||||
return False
|
||||
self.created.append(
|
||||
{
|
||||
"conversation_id": conversation_id,
|
||||
"user": user,
|
||||
"name": name,
|
||||
"status": status,
|
||||
"introduction": introduction,
|
||||
"created_at": created_at,
|
||||
"updated_at": updated_at,
|
||||
}
|
||||
)
|
||||
self.existing[conversation_id] = self.created[-1]
|
||||
return True
|
||||
|
||||
def get_conversation_by_id(self, conversation_id):
|
||||
return self.existing.get(conversation_id)
|
||||
|
||||
def update_conversation_updated_at(self, conversation_id, updated_at):
|
||||
self.updated.append({"conversation_id": conversation_id, "updated_at": updated_at})
|
||||
if self.update_should_fail:
|
||||
return False
|
||||
return conversation_id in self.existing
|
||||
|
||||
def save_message(self, **kwargs):
|
||||
self.saved.append(kwargs)
|
||||
return True
|
||||
|
||||
|
||||
def _build_client(workflow_manager, monkeypatch, storage):
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
app.dependency_overrides[get_workflow_manager] = lambda: workflow_manager
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: storage)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_run_workflow_creates_conversation_when_missing_id(monkeypatch):
|
||||
workflow_manager = StubWorkflowManager()
|
||||
storage = StubMessageStorage()
|
||||
client = _build_client(workflow_manager, monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
"query": "查询订单",
|
||||
"inputs": {},
|
||||
"response_mode": "blocking",
|
||||
"user": "tester",
|
||||
"conversation_id": None,
|
||||
"files": [],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert len(storage.created) == 1
|
||||
created = storage.created[0]
|
||||
assert created["status"] == "normal"
|
||||
assert created["name"] == "查询订单"
|
||||
assert created["introduction"] is None
|
||||
assert created["created_at"] == created["updated_at"]
|
||||
assert workflow_manager.calls[0]["session_id"] == created["conversation_id"]
|
||||
assert response.json()["session_id"] == created["conversation_id"]
|
||||
assert storage.saved[0]["created_at"] == created["created_at"]
|
||||
assert storage.saved[0]["updated_at"] == created["updated_at"]
|
||||
|
||||
|
||||
def test_run_workflow_first_turn_name_uses_first_20_chars(monkeypatch):
|
||||
workflow_manager = StubWorkflowManager()
|
||||
storage = StubMessageStorage()
|
||||
client = _build_client(workflow_manager, monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
"query": "12345678901234567890EXTRA_TEXT",
|
||||
"inputs": {},
|
||||
"response_mode": "blocking",
|
||||
"user": "tester",
|
||||
"conversation_id": None,
|
||||
"files": [],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert storage.created[0]["name"] == "12345678901234567890"
|
||||
|
||||
|
||||
def test_run_workflow_updates_existing_conversation(monkeypatch):
|
||||
workflow_manager = StubWorkflowManager()
|
||||
storage = StubMessageStorage()
|
||||
storage.existing["cid-exists"] = {
|
||||
"conversation_id": "cid-exists",
|
||||
"user": "tester",
|
||||
"name": "old",
|
||||
"status": "normal",
|
||||
"created_at": 1,
|
||||
"updated_at": 1,
|
||||
}
|
||||
client = _build_client(workflow_manager, monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
"query": "查询订单",
|
||||
"inputs": {},
|
||||
"response_mode": "blocking",
|
||||
"user": "tester",
|
||||
"conversation_id": "cid-exists",
|
||||
"files": [],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert storage.updated and storage.updated[0]["conversation_id"] == "cid-exists"
|
||||
assert workflow_manager.calls[0]["session_id"] == "cid-exists"
|
||||
|
||||
|
||||
def test_run_workflow_returns_400_when_provided_conversation_id_not_found(monkeypatch):
|
||||
workflow_manager = StubWorkflowManager()
|
||||
storage = StubMessageStorage()
|
||||
client = _build_client(workflow_manager, monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
"query": "查询订单",
|
||||
"inputs": {},
|
||||
"response_mode": "blocking",
|
||||
"user": "tester",
|
||||
"conversation_id": "cid-missing",
|
||||
"files": [],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.json()["detail"]["code"] == "CONVERSATION_NOT_FOUND"
|
||||
|
||||
|
||||
def test_run_workflow_returns_500_when_conversation_create_fails(monkeypatch):
|
||||
workflow_manager = StubWorkflowManager()
|
||||
storage = StubMessageStorage()
|
||||
storage.create_should_fail = True
|
||||
client = _build_client(workflow_manager, monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
"query": "查询订单",
|
||||
"inputs": {},
|
||||
"response_mode": "blocking",
|
||||
"user": "tester",
|
||||
"conversation_id": None,
|
||||
"files": [],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 500
|
||||
assert response.json()["detail"]["code"] == "CONVERSATION_CREATE_FAILED"
|
||||
|
||||
|
||||
def test_run_workflow_returns_500_when_conversation_update_fails(monkeypatch):
|
||||
workflow_manager = StubWorkflowManager()
|
||||
storage = StubMessageStorage()
|
||||
storage.existing["cid-exists"] = {
|
||||
"conversation_id": "cid-exists",
|
||||
"user": "tester",
|
||||
"name": "old",
|
||||
"status": "normal",
|
||||
"created_at": 1,
|
||||
"updated_at": 1,
|
||||
}
|
||||
storage.update_should_fail = True
|
||||
client = _build_client(workflow_manager, monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
"query": "查询订单",
|
||||
"inputs": {},
|
||||
"response_mode": "blocking",
|
||||
"user": "tester",
|
||||
"conversation_id": "cid-exists",
|
||||
"files": [],
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 500
|
||||
assert response.json()["detail"]["code"] == "CONVERSATION_UPDATE_FAILED"
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
from datetime import datetime
|
||||
|
||||
import pytest
|
||||
|
||||
from services.common.datetime_utils import DateTimeGenerator
|
||||
|
||||
|
||||
def test_datetime_generator_supports_10_digit_timestamp():
|
||||
bundle = DateTimeGenerator.bundle("1710912000")
|
||||
assert bundle.epoch_seconds == 1710912000
|
||||
assert bundle.epoch_millis == 1710912000000
|
||||
|
||||
|
||||
def test_datetime_generator_supports_13_digit_timestamp():
|
||||
bundle = DateTimeGenerator.bundle("1710912000123")
|
||||
assert bundle.epoch_millis == 1710912000123
|
||||
|
||||
|
||||
def test_datetime_generator_supports_yyyymmdd():
|
||||
bundle = DateTimeGenerator.bundle("20260320")
|
||||
assert bundle.yyyymmdd == "20260320"
|
||||
assert bundle.date_str == "2026-03-20"
|
||||
|
||||
|
||||
def test_datetime_generator_supports_date_and_datetime_formats():
|
||||
from_date = DateTimeGenerator.bundle("2026-03-20")
|
||||
from_dt = DateTimeGenerator.bundle("2026-03-20 11:27:53")
|
||||
|
||||
assert from_date.date_str == "2026-03-20"
|
||||
assert from_dt.datetime_str == "2026-03-20 11:27:53"
|
||||
|
||||
|
||||
def test_datetime_generator_supports_iso_and_datetime_objects():
|
||||
from_iso = DateTimeGenerator.bundle("2026-03-20T11:27:53")
|
||||
from_obj = DateTimeGenerator.bundle(datetime(2026, 3, 20, 11, 27, 53))
|
||||
|
||||
assert from_iso.date_str == "2026-03-20"
|
||||
assert from_obj.datetime_str == "2026-03-20 11:27:53"
|
||||
|
||||
|
||||
def test_datetime_generator_invalid_value_behaviour():
|
||||
fallback = DateTimeGenerator.bundle("not-a-date")
|
||||
assert fallback.epoch_millis > 0
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
DateTimeGenerator.bundle("not-a-date", default_to_now=False)
|
||||
|
||||
@@ -0,0 +1,166 @@
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
|
||||
from scripts.demo_chat import format_demo_result, main, run_turn
|
||||
|
||||
|
||||
class FakeMessage:
|
||||
def __init__(self, content: str):
|
||||
self.content = content
|
||||
|
||||
|
||||
class FakeAgent:
|
||||
def __init__(self, model_section=None):
|
||||
self.model_section = model_section
|
||||
self.calls = []
|
||||
|
||||
def run(self, query, **kwargs):
|
||||
self.calls.append((query, kwargs))
|
||||
return {
|
||||
"messages": [FakeMessage('{"status_code": 200, "text": "..."}')],
|
||||
"context": {
|
||||
"final_sql": "SELECT service_order_id, ship_to_country FROM dwd_ai.apbo_eta_ful",
|
||||
"sr_api_result": json.dumps(
|
||||
{
|
||||
"status_code": 200,
|
||||
"text": json.dumps(
|
||||
{
|
||||
"data": [
|
||||
{"service_order_id": "4020438779", "ship_to_country": "VN"},
|
||||
{"service_order_id": "4020438780", "ship_to_country": "PH"},
|
||||
]
|
||||
},
|
||||
ensure_ascii=False,
|
||||
),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
),
|
||||
},
|
||||
"final_step": "response_generated",
|
||||
}
|
||||
|
||||
|
||||
def test_format_demo_result_includes_time_sql_row_count_and_table():
|
||||
result = FakeAgent().run("查询")
|
||||
|
||||
text = format_demo_result(result, 1.234)
|
||||
|
||||
assert "耗时: 1.23s" in text
|
||||
assert "SQL:" in text
|
||||
assert "SELECT service_order_id, ship_to_country FROM dwd_ai.apbo_eta_ful" in text
|
||||
assert "数据行数: 2" in text
|
||||
assert "SQL执行结果表:" in text
|
||||
assert "service_order_id" in text
|
||||
assert "4020438779" in text
|
||||
|
||||
|
||||
def test_format_demo_result_falls_back_when_result_is_not_tabular():
|
||||
result = {
|
||||
"messages": [FakeMessage("请求失败: timeout")],
|
||||
"context": {
|
||||
"final_sql": "SELECT 1",
|
||||
"sr_api_result": "请求失败: timeout",
|
||||
},
|
||||
}
|
||||
|
||||
text = format_demo_result(result, 0.4)
|
||||
|
||||
assert "耗时: 0.40s" in text
|
||||
assert "数据行数: 0" in text
|
||||
assert "SQL执行结果:" in text
|
||||
assert "请求失败: timeout" in text
|
||||
|
||||
|
||||
def test_format_demo_result_renders_empty_structured_result_as_table():
|
||||
result = {
|
||||
"messages": [FakeMessage('{"status_code": 200, "text": "..."}')],
|
||||
"context": {
|
||||
"final_sql": "SELECT 1",
|
||||
"sr_api_result": json.dumps(
|
||||
{
|
||||
"status_code": 200,
|
||||
"text": json.dumps(
|
||||
{"code": "0", "data": [], "msg": "操作成功", "total": 0},
|
||||
ensure_ascii=False,
|
||||
),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
text = format_demo_result(result, 0.4)
|
||||
|
||||
assert "数据行数: 0" in text
|
||||
assert "SQL执行结果表:" in text
|
||||
assert "<empty table>" in text
|
||||
assert '"status_code": 200' not in text
|
||||
|
||||
|
||||
def test_format_demo_result_prefers_fallback_answer_for_empty_result():
|
||||
result = {
|
||||
"messages": [FakeMessage("未查询到符合条件的数据,请尝试调整筛选条件。")],
|
||||
"context": {
|
||||
"final_sql": "SELECT 1",
|
||||
"sr_api_result": json.dumps(
|
||||
{
|
||||
"status_code": 200,
|
||||
"text": json.dumps(
|
||||
{"code": "0", "data": [], "msg": "操作成功", "total": 0},
|
||||
ensure_ascii=False,
|
||||
),
|
||||
},
|
||||
ensure_ascii=False,
|
||||
),
|
||||
"is_empty_result": True,
|
||||
"response_source": "model_empty_result_fallback",
|
||||
},
|
||||
}
|
||||
|
||||
text = format_demo_result(result, 0.4)
|
||||
|
||||
assert "结果说明:" in text
|
||||
assert "未查询到符合条件的数据" in text
|
||||
assert "SQL执行结果表:" not in text
|
||||
|
||||
|
||||
def test_main_one_shot_success(capsys):
|
||||
with patch("scripts.demo_chat.Config.validate_config", return_value=None), \
|
||||
patch("scripts.demo_chat.ConversationAgent", FakeAgent):
|
||||
exit_code = main([
|
||||
"--query",
|
||||
"查询 SO 4020438779 的 eta 信息",
|
||||
])
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert exit_code == 0
|
||||
assert "耗时:" in captured.out
|
||||
assert "数据行数: 2" in captured.out
|
||||
assert "SELECT service_order_id, ship_to_country FROM dwd_ai.apbo_eta_ful" in captured.out
|
||||
|
||||
|
||||
def test_run_turn_disables_debug_node_trace(capsys):
|
||||
agent = FakeAgent()
|
||||
|
||||
run_turn(
|
||||
agent,
|
||||
"查询 SO 4020438779 的 eta 信息",
|
||||
user="tester",
|
||||
conversation_id="cid-1",
|
||||
)
|
||||
|
||||
_, kwargs = agent.calls[-1]
|
||||
assert kwargs["debug_node_trace"] is False
|
||||
|
||||
|
||||
def test_main_config_error(capsys):
|
||||
with patch("scripts.demo_chat.Config.validate_config", side_effect=ValueError("bad config")):
|
||||
exit_code = main(["--query", "hello"])
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert exit_code == 1
|
||||
assert "Configuration error" in captured.err
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,201 @@
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from api.dependencies import get_workflow_manager
|
||||
from api.endpoints import router
|
||||
from workflows.workflow_manager import WorkflowType
|
||||
|
||||
|
||||
class GuardWorkflowManager:
|
||||
def execute_workflow(self, *args, **kwargs):
|
||||
raise AssertionError("execute_workflow should not be called for invalid query payloads")
|
||||
|
||||
|
||||
class StubWorkflowManager:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
def execute_workflow(self, workflow_type, user_input, session_id=None, **kwargs):
|
||||
self.calls.append(
|
||||
{
|
||||
"workflow_type": workflow_type,
|
||||
"user_input": user_input,
|
||||
"session_id": session_id,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
)
|
||||
return {
|
||||
"session_id": session_id or "cid-1",
|
||||
"workflow_type": workflow_type.value if isinstance(workflow_type, WorkflowType) else str(workflow_type),
|
||||
"result": {"context": {"final_sql": "SELECT 1"}, "ok": True},
|
||||
}
|
||||
|
||||
|
||||
class DisabledMessageStorage:
|
||||
enabled = False
|
||||
|
||||
def save_message(self, **kwargs):
|
||||
return False
|
||||
|
||||
|
||||
class CaptureMessageStorage:
|
||||
enabled = True
|
||||
|
||||
def __init__(self):
|
||||
self.saved = []
|
||||
self.created = []
|
||||
self.existing = {}
|
||||
|
||||
def create_conversation(self, conversation_id, user, name, status, introduction, created_at, updated_at):
|
||||
record = {
|
||||
"conversation_id": conversation_id,
|
||||
"user": user,
|
||||
"name": name,
|
||||
"status": status,
|
||||
"introduction": introduction,
|
||||
"created_at": created_at,
|
||||
"updated_at": updated_at,
|
||||
}
|
||||
self.created.append(record)
|
||||
self.existing[conversation_id] = record
|
||||
return True
|
||||
|
||||
def get_conversation_by_id(self, conversation_id):
|
||||
return self.existing.get(conversation_id)
|
||||
|
||||
def update_conversation_updated_at(self, conversation_id, updated_at):
|
||||
if conversation_id not in self.existing:
|
||||
return False
|
||||
self.existing[conversation_id]["updated_at"] = updated_at
|
||||
return True
|
||||
|
||||
def save_message(self, **kwargs):
|
||||
self.saved.append(kwargs)
|
||||
return True
|
||||
|
||||
|
||||
def _build_client(workflow_manager) -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
app.dependency_overrides[get_workflow_manager] = lambda: workflow_manager
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _payload(query, conversation_id="cid-123"):
|
||||
return {
|
||||
"query": query,
|
||||
"inputs": {},
|
||||
"response_mode": "streaming",
|
||||
"user": "tester",
|
||||
"conversation_id": conversation_id,
|
||||
"files": [],
|
||||
}
|
||||
|
||||
|
||||
def test_invalid_query_returns_400_for_stream_endpoint(monkeypatch):
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage())
|
||||
client = _build_client(GuardWorkflowManager())
|
||||
|
||||
for invalid_query in (None, "", " "):
|
||||
response = client.post("/api/workflows/stream", json=_payload(invalid_query))
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.json() == {
|
||||
"detail": {
|
||||
"code": "INVALID_REQUEST",
|
||||
"message": "query 不能为空",
|
||||
"detail": {"field": "query", "reason": "missing_or_blank"},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_invalid_query_returns_400_for_blocking_endpoints(monkeypatch):
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage())
|
||||
client = _build_client(GuardWorkflowManager())
|
||||
|
||||
for path in ("/api/workflows", "/api/sql/generate"):
|
||||
response = client.post(path, json=_payload(" "))
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.json()["detail"]["code"] == "INVALID_REQUEST"
|
||||
assert response.json()["detail"]["detail"] == {"field": "query", "reason": "missing_or_blank"}
|
||||
|
||||
|
||||
def test_valid_query_still_reaches_workflow_manager(monkeypatch):
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage())
|
||||
workflow_manager = StubWorkflowManager()
|
||||
client = _build_client(workflow_manager)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
**_payload("查询 SO 4020438779 的 eta 信息"),
|
||||
"response_mode": "blocking",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert workflow_manager.calls[0]["user_input"] == "查询 SO 4020438779 的 eta 信息"
|
||||
assert response.json()["session_id"] == "cid-123"
|
||||
assert response.json()["workflow_type"] == "conversation"
|
||||
|
||||
|
||||
def test_valid_query_persists_message_record(monkeypatch):
|
||||
storage = CaptureMessageStorage()
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: storage)
|
||||
workflow_manager = StubWorkflowManager()
|
||||
client = _build_client(workflow_manager)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
**_payload("查询 SO 4020438779 的 eta 信息", conversation_id=None),
|
||||
"response_mode": "blocking",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert len(storage.saved) == 1
|
||||
saved = storage.saved[0]
|
||||
assert saved["conversation_id"] == response.json()["session_id"]
|
||||
assert saved["query"] == "查询 SO 4020438779 的 eta 信息"
|
||||
assert saved["workflow_type"] == "conversation"
|
||||
assert saved["created_at"] == saved["updated_at"]
|
||||
assert saved["logs"][0].startswith("run_workflow.start")
|
||||
assert any(item.startswith("run_workflow.success") for item in saved["logs"])
|
||||
assert storage.created[0]["name"] == "查询 SO 4020438779 的 e"
|
||||
|
||||
|
||||
def test_missing_conversation_id_returns_400_for_blocking_workflow(monkeypatch):
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: CaptureMessageStorage())
|
||||
workflow_manager = StubWorkflowManager()
|
||||
client = _build_client(workflow_manager)
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows",
|
||||
json={
|
||||
**_payload("查询 SO 4020438779 的 eta 信息", conversation_id="cid-missing"),
|
||||
"response_mode": "blocking",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.json()["detail"]["code"] == "CONVERSATION_NOT_FOUND"
|
||||
|
||||
|
||||
def test_invalid_response_mode_uses_dedicated_error_code(monkeypatch):
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage())
|
||||
client = _build_client(GuardWorkflowManager())
|
||||
|
||||
response = client.post(
|
||||
"/api/workflows/stream",
|
||||
json={
|
||||
**_payload("查询 SO 4020438779 的 eta 信息"),
|
||||
"response_mode": "blocking",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
assert response.json()["detail"]["code"] == "INVALID_RESPONSE_MODE"
|
||||
|
||||
|
||||
@@ -0,0 +1,148 @@
|
||||
import json
|
||||
import sys
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import httpx
|
||||
|
||||
from config import Config
|
||||
|
||||
|
||||
def _base_url() -> str:
|
||||
app = Config.get_section("app")
|
||||
host = app.get("host", "127.0.0.1")
|
||||
port = app.get("port", "8000")
|
||||
if host in ("0.0.0.0", "::"):
|
||||
host = "127.0.0.1"
|
||||
return f"http://{host}:{port}"
|
||||
|
||||
|
||||
def _post(client: httpx.Client, url: str, payload: Dict[str, Any]) -> None:
|
||||
resp = client.post(url, json=payload)
|
||||
print(f"POST {url} -> {resp.status_code}")
|
||||
print(resp.text)
|
||||
|
||||
|
||||
def _put(client: httpx.Client, url: str, payload: Dict[str, Any]) -> None:
|
||||
resp = client.put(url, json=payload)
|
||||
print(f"PUT {url} -> {resp.status_code}")
|
||||
print(resp.text)
|
||||
|
||||
|
||||
def _get(client: httpx.Client, url: str) -> None:
|
||||
resp = client.get(url)
|
||||
print(f"GET {url} -> {resp.status_code}")
|
||||
print(resp.text)
|
||||
|
||||
|
||||
def _stream_sse(client: httpx.Client, url: str, payload: Dict[str, Any]) -> None:
|
||||
with client.stream("POST", url, json=payload) as resp:
|
||||
print(f"POST {url} -> {resp.status_code}")
|
||||
current_event = "message"
|
||||
for raw in resp.iter_lines():
|
||||
if raw is None:
|
||||
continue
|
||||
line = raw.strip()
|
||||
if not line:
|
||||
continue
|
||||
if line.startswith("event:"):
|
||||
current_event = line.split(":", 1)[1].strip() or "message"
|
||||
continue
|
||||
if line.startswith("data:"):
|
||||
data = line.split(":", 1)[1].strip()
|
||||
print(f"[{current_event}] {data}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
base = _base_url()
|
||||
menu: List[str] = [
|
||||
"1) GET /health",
|
||||
"2) GET /nacos/status",
|
||||
"3) POST /api/workflows (conversation)",
|
||||
"4) POST /api/workflows/stream (conversation)",
|
||||
"5) POST /api/sql/generate",
|
||||
"6) POST /api/tools/execute",
|
||||
"7) POST /api/prompts/reload",
|
||||
"8) POST /api/ragflow/table-retrieval/upload",
|
||||
"9) PUT /api/ragflow/table-retrieval/update",
|
||||
"10) POST /api/ragflow/sql-gen/upload",
|
||||
"11) PUT /api/ragflow/sql-gen/update",
|
||||
"0) Exit",
|
||||
]
|
||||
|
||||
with httpx.Client(timeout=60) as client:
|
||||
while True:
|
||||
print("\n可用接口:")
|
||||
for line in menu:
|
||||
print(line)
|
||||
|
||||
choice = input("\n请选择编号: ").strip()
|
||||
if choice == "0":
|
||||
break
|
||||
|
||||
if choice == "1":
|
||||
_get(client, f"{base}/health")
|
||||
elif choice == "2":
|
||||
_get(client, f"{base}/nacos/status")
|
||||
elif choice == "3":
|
||||
payload = {
|
||||
"query": "查询 SO 4020438779 的 eta 信息",
|
||||
"conversation_id": None,
|
||||
"response_mode": "blocking",
|
||||
"user": "tester",
|
||||
"inputs": {},
|
||||
"files": [],
|
||||
}
|
||||
_post(client, f"{base}/api/workflows", payload)
|
||||
elif choice == "4":
|
||||
payload = {
|
||||
"query": "查询 SO 4019671497 的 eta 信息",
|
||||
"conversation_id": None,
|
||||
"response_mode": "streaming",
|
||||
"user": "tester",
|
||||
"inputs": {},
|
||||
"files": [],
|
||||
}
|
||||
_stream_sse(client, f"{base}/api/workflows/stream", payload)
|
||||
elif choice == "5":
|
||||
payload = {
|
||||
"query": "查询 SO 4020438779 的 eta 信息",
|
||||
"conversation_id": None,
|
||||
"response_mode": "blocking",
|
||||
"user": "tester",
|
||||
"inputs": {},
|
||||
"files": [],
|
||||
}
|
||||
_post(client, f"{base}/api/sql/generate", payload)
|
||||
elif choice == "6":
|
||||
payload = {
|
||||
"tool_name": "sr_api_query",
|
||||
"payload": {
|
||||
"sql": "SELECT 1",
|
||||
"page": 1,
|
||||
"rows": 1,
|
||||
"orderBySelect": True,
|
||||
"timeout": 30,
|
||||
},
|
||||
}
|
||||
_post(client, f"{base}/api/tools/execute", payload)
|
||||
elif choice == "7":
|
||||
_post(client, f"{base}/api/prompts/reload", {})
|
||||
elif choice == "8":
|
||||
_post(client, f"{base}/api/ragflow/table-retrieval/upload", {})
|
||||
elif choice == "9":
|
||||
cfg = {"name": "table_retrieval_dataset"}
|
||||
_put(client, f"{base}/api/ragflow/table-retrieval/update", cfg)
|
||||
elif choice == "10":
|
||||
_post(client, f"{base}/api/ragflow/sql-gen/upload", {})
|
||||
elif choice == "11":
|
||||
cfg = {"name": "sql_gen_dataset"}
|
||||
_put(client, f"{base}/api/ragflow/sql-gen/update", cfg)
|
||||
else:
|
||||
print("无效选择")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
main()
|
||||
except KeyboardInterrupt:
|
||||
sys.exit(0)
|
||||
@@ -0,0 +1,77 @@
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from api.endpoints import router
|
||||
|
||||
|
||||
class StubMessageStorage:
|
||||
def __init__(self, updated=True):
|
||||
self.updated = updated
|
||||
self.calls = []
|
||||
|
||||
def update_feedback_by_message_id(self, message_id, feedback, feedback_content=None):
|
||||
self.calls.append(
|
||||
{
|
||||
"message_id": message_id,
|
||||
"feedback": feedback,
|
||||
"feedback_content": feedback_content,
|
||||
}
|
||||
)
|
||||
return self.updated
|
||||
|
||||
|
||||
def _build_client(monkeypatch, storage):
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: storage)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_feedback_endpoint_writes_like(monkeypatch):
|
||||
storage = StubMessageStorage(updated=True)
|
||||
client = _build_client(monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/messages/feedback",
|
||||
json={"message_id": "mid-1", "feedback": "like"},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {"ok": True, "message_id": "mid-1"}
|
||||
assert storage.calls[0] == {
|
||||
"message_id": "mid-1",
|
||||
"feedback": "like",
|
||||
"feedback_content": None,
|
||||
}
|
||||
|
||||
|
||||
def test_feedback_endpoint_requires_feedback_content_for_dislike(monkeypatch):
|
||||
storage = StubMessageStorage(updated=True)
|
||||
client = _build_client(monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/messages/feedback",
|
||||
json={"message_id": "mid-2", "feedback": "dislike"},
|
||||
)
|
||||
|
||||
assert response.status_code == 422
|
||||
|
||||
|
||||
def test_feedback_endpoint_returns_400_when_message_not_found(monkeypatch):
|
||||
storage = StubMessageStorage(updated=False)
|
||||
client = _build_client(monkeypatch, storage)
|
||||
|
||||
response = client.post(
|
||||
"/api/messages/feedback",
|
||||
json={
|
||||
"message_id": "missing-mid",
|
||||
"feedback": "dislike",
|
||||
"feedback_content": "not useful",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 400
|
||||
body = response.json()
|
||||
assert body["detail"]["code"] == "INVALID_REQUEST"
|
||||
assert body["detail"]["detail"]["field"] == "message_id"
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from schemas.message_feedback_request import MessageFeedbackRequestDTO
|
||||
|
||||
|
||||
def test_feedback_request_accepts_like_without_content():
|
||||
dto = MessageFeedbackRequestDTO(message_id="mid-1", feedback="like")
|
||||
|
||||
assert dto.message_id == "mid-1"
|
||||
assert dto.feedback == "like"
|
||||
assert dto.feedback_content is None
|
||||
|
||||
|
||||
def test_feedback_request_requires_content_for_dislike():
|
||||
with pytest.raises(ValidationError):
|
||||
MessageFeedbackRequestDTO(message_id="mid-2", feedback="dislike")
|
||||
|
||||
|
||||
def test_feedback_request_accepts_dislike_with_content():
|
||||
dto = MessageFeedbackRequestDTO(
|
||||
message_id="mid-3",
|
||||
feedback="dislike",
|
||||
feedback_content="结果不准确",
|
||||
)
|
||||
|
||||
assert dto.feedback == "dislike"
|
||||
assert dto.feedback_content == "结果不准确"
|
||||
|
||||
@@ -0,0 +1,456 @@
|
||||
import json
|
||||
from datetime import datetime
|
||||
|
||||
from services.storage.message_storage import MessageStorage
|
||||
|
||||
|
||||
class _FakeCursor:
|
||||
def __init__(self, sink):
|
||||
self._sink = sink
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
return False
|
||||
|
||||
def execute(self, sql, params=None):
|
||||
self._sink["sql"] = sql
|
||||
self._sink["params"] = params
|
||||
self._sink.setdefault("calls", []).append((sql, params))
|
||||
return self._sink.get("execute_return", 1)
|
||||
|
||||
def fetchone(self):
|
||||
return self._sink.get("fetchone_result")
|
||||
|
||||
|
||||
class _FakeConn:
|
||||
def __init__(self, sink):
|
||||
self._sink = sink
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
return False
|
||||
|
||||
def cursor(self, *args, **kwargs):
|
||||
return _FakeCursor(self._sink)
|
||||
|
||||
|
||||
class _ScriptedCursor:
|
||||
def __init__(self, steps):
|
||||
self._steps = steps
|
||||
self._index = 0
|
||||
self._current = None
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
return False
|
||||
|
||||
def execute(self, sql, params=None):
|
||||
if self._index >= len(self._steps):
|
||||
raise AssertionError(f"Unexpected SQL: {sql}")
|
||||
self._current = self._steps[self._index]
|
||||
self._current["sql"] = sql
|
||||
self._current["params"] = params
|
||||
self._index += 1
|
||||
return self._current.get("execute_return", 1)
|
||||
|
||||
def fetchone(self):
|
||||
return None if self._current is None else self._current.get("fetchone_result")
|
||||
|
||||
|
||||
class _ScriptedConn:
|
||||
def __init__(self, steps):
|
||||
self._cursor = _ScriptedCursor(steps)
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
return False
|
||||
|
||||
def cursor(self, *args, **kwargs):
|
||||
return self._cursor
|
||||
|
||||
|
||||
def test_save_message_writes_messages_dto_shape(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"table": "messages",
|
||||
"messages_table": "messages",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
|
||||
storage = MessageStorage()
|
||||
storage._inited = True
|
||||
|
||||
sink = {}
|
||||
monkeypatch.setattr(storage, "_get_conn", lambda: _FakeConn(sink))
|
||||
|
||||
ok = storage.save_message(
|
||||
conversation_id="cid-1",
|
||||
message_id="mid-1",
|
||||
query="select",
|
||||
answer="answer",
|
||||
workflow_type="conversation",
|
||||
user="tester",
|
||||
sql_query="SELECT 1",
|
||||
execution_result={"total": 1},
|
||||
metadata={"trace_id": "t-1"},
|
||||
)
|
||||
|
||||
assert ok is True
|
||||
assert "INSERT INTO messages" in sink["sql"]
|
||||
|
||||
params = sink["params"]
|
||||
assert params[0] == "mid-1"
|
||||
assert params[1] == "tester"
|
||||
assert isinstance(params[2], datetime)
|
||||
assert params[3] == "tester"
|
||||
assert isinstance(params[4], datetime)
|
||||
assert params[5] == "tester"
|
||||
assert params[6] == "tester"
|
||||
assert params[7] == "mid-1"
|
||||
assert params[8] == "cid-1"
|
||||
assert params[9] == "tester"
|
||||
assert params[10] == "select"
|
||||
assert params[11] == "answer"
|
||||
assert params[12] is None
|
||||
assert params[13] is None
|
||||
assert isinstance(params[14], int)
|
||||
assert isinstance(params[15], int)
|
||||
|
||||
log_payload = json.loads(params[16])
|
||||
assert log_payload["workflow_type"] == "conversation"
|
||||
assert log_payload["sql_query"] == "SELECT 1"
|
||||
assert log_payload["execution_result"] == {"total": 1}
|
||||
assert log_payload["metadata"] == {"trace_id": "t-1"}
|
||||
|
||||
|
||||
def test_save_message_returns_false_when_disabled(monkeypatch):
|
||||
monkeypatch.setattr("services.storage.message_storage.Config.get_section", lambda section: {"enabled": "false"})
|
||||
storage = MessageStorage()
|
||||
|
||||
ok = storage.save_message(conversation_id="cid", message_id="mid", query="q")
|
||||
|
||||
assert ok is False
|
||||
|
||||
|
||||
def test_save_message_uses_explicit_timestamps(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"table": "messages",
|
||||
"messages_table": "messages",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
storage = MessageStorage()
|
||||
storage._inited = True
|
||||
|
||||
sink = {}
|
||||
monkeypatch.setattr(storage, "_get_conn", lambda: _FakeConn(sink))
|
||||
|
||||
ok = storage.save_message(
|
||||
conversation_id="cid-explicit",
|
||||
message_id="mid-explicit",
|
||||
query="select explicit",
|
||||
answer="answer",
|
||||
workflow_type="conversation",
|
||||
user="tester",
|
||||
created_at=1234,
|
||||
updated_at=5678,
|
||||
)
|
||||
|
||||
assert ok is True
|
||||
assert isinstance(sink["params"][2], datetime)
|
||||
assert isinstance(sink["params"][4], datetime)
|
||||
assert sink["params"][14] == 1234
|
||||
assert sink["params"][15] == 5678
|
||||
|
||||
|
||||
def test_save_message_persists_java_style_logs_into_log_data(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"table": "messages",
|
||||
"messages_table": "messages",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
storage = MessageStorage()
|
||||
storage._inited = True
|
||||
|
||||
sink = {}
|
||||
monkeypatch.setattr(storage, "_get_conn", lambda: _FakeConn(sink))
|
||||
|
||||
ok = storage.save_message(
|
||||
conversation_id="cid-log",
|
||||
message_id="mid-log",
|
||||
query="query",
|
||||
answer="answer",
|
||||
workflow_type="conversation",
|
||||
logs=["line1", "line2"],
|
||||
)
|
||||
|
||||
assert ok is True
|
||||
log_payload = json.loads(sink["params"][16])
|
||||
assert log_payload["data"] == "line1\nline2"
|
||||
|
||||
|
||||
def test_entity_debug_logging_is_silent_when_switch_disabled(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"entity_debug_enabled": "false",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"messages_table": "messages",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
storage = MessageStorage()
|
||||
|
||||
printed = []
|
||||
monkeypatch.setattr("builtins.print", lambda *args, **kwargs: printed.append(args[0] if args else ""))
|
||||
|
||||
storage._log_entity_stage("messages", "create", "start", storage._now_ms(), {"message_id": "mid-1"})
|
||||
|
||||
assert printed == []
|
||||
|
||||
|
||||
def test_entity_debug_logging_emits_stage_payload_when_switch_enabled(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"entity_debug_enabled": "true",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"messages_table": "messages",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
storage = MessageStorage()
|
||||
|
||||
printed = []
|
||||
monkeypatch.setattr("builtins.print", lambda *args, **kwargs: printed.append(args[0] if args else ""))
|
||||
|
||||
storage._log_entity_stage("conversations", "get_by_id", "start", storage._now_ms(), {"conversation_id": "cid-1"})
|
||||
|
||||
assert len(printed) == 1
|
||||
payload = json.loads(printed[0])
|
||||
assert payload["level"] == "DEBUG"
|
||||
assert payload["event"] == "message_storage.conversations.get_by_id.start"
|
||||
assert payload["payload"]["conversation_id"] == "cid-1"
|
||||
assert payload["payload"]["entity"] == "conversations"
|
||||
assert payload["payload"]["action"] == "get_by_id"
|
||||
assert payload["payload"]["stage"] == "start"
|
||||
|
||||
|
||||
def test_update_feedback_by_message_id_success(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"table": "messages",
|
||||
"messages_table": "messages",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
storage = MessageStorage()
|
||||
storage._inited = True
|
||||
|
||||
sink = {}
|
||||
monkeypatch.setattr(storage, "_get_conn", lambda: _FakeConn(sink))
|
||||
|
||||
ok = storage.update_feedback_by_message_id(
|
||||
message_id="mid-1",
|
||||
feedback="dislike",
|
||||
feedback_content="结果不准确",
|
||||
)
|
||||
|
||||
assert ok is True
|
||||
assert "UPDATE messages" in sink["sql"]
|
||||
assert sink["params"][0] == "dislike"
|
||||
assert sink["params"][1] == "结果不准确"
|
||||
assert isinstance(sink["params"][2], int)
|
||||
assert isinstance(sink["params"][3], datetime)
|
||||
assert sink["params"][4] == "mid-1"
|
||||
|
||||
|
||||
def test_update_feedback_by_message_id_not_found_returns_false(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"table": "messages",
|
||||
"messages_table": "messages",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
storage = MessageStorage()
|
||||
storage._inited = True
|
||||
|
||||
sink = {"execute_return": 0}
|
||||
monkeypatch.setattr(storage, "_get_conn", lambda: _FakeConn(sink))
|
||||
|
||||
ok = storage.update_feedback_by_message_id(
|
||||
message_id="missing-mid",
|
||||
feedback="like",
|
||||
)
|
||||
|
||||
assert ok is False
|
||||
|
||||
|
||||
def test_conversation_create_get_update(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"table": "messages",
|
||||
"messages_table": "messages",
|
||||
"conversation_table": "conversations",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
storage = MessageStorage()
|
||||
storage._inited = True
|
||||
storage._conversation_schema_checked = True
|
||||
|
||||
sink = {}
|
||||
monkeypatch.setattr(storage, "_get_conn", lambda: _FakeConn(sink))
|
||||
|
||||
created = storage.create_conversation(
|
||||
conversation_id="cid-1",
|
||||
user="tester",
|
||||
name="hello",
|
||||
status="normal",
|
||||
introduction="intro",
|
||||
created_at=1000,
|
||||
updated_at=1000,
|
||||
)
|
||||
assert created is True
|
||||
assert "INSERT INTO conversations" in sink["sql"]
|
||||
assert sink["params"][0] == "cid-1"
|
||||
assert sink["params"][1] == "tester"
|
||||
assert isinstance(sink["params"][2], datetime)
|
||||
assert sink["params"][3] == "tester"
|
||||
assert isinstance(sink["params"][4], datetime)
|
||||
assert sink["params"][7] == "cid-1"
|
||||
assert sink["params"][8] == "tester"
|
||||
assert sink["params"][9] == "hello"
|
||||
|
||||
sink["fetchone_result"] = {
|
||||
"conversation_id": "cid-1",
|
||||
"user": "tester",
|
||||
"name": "hello",
|
||||
"status": "normal",
|
||||
"introduction": "intro",
|
||||
"created_at": 1000,
|
||||
"updated_at": 1000,
|
||||
}
|
||||
got = storage.get_conversation_by_id("cid-1")
|
||||
assert got is not None
|
||||
assert got["conversation_id"] == "cid-1"
|
||||
assert got["introduction"] == "intro"
|
||||
|
||||
updated = storage.update_conversation_updated_at("cid-1", 2000)
|
||||
assert updated is True
|
||||
assert "UPDATE conversations" in sink["sql"]
|
||||
assert sink["params"][0] == 2000
|
||||
assert isinstance(sink["params"][1], datetime)
|
||||
assert sink["params"][2] == "cid-1"
|
||||
|
||||
|
||||
def test_create_conversation_auto_adds_missing_name_column(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"services.storage.message_storage.Config.get_section",
|
||||
lambda section: {
|
||||
"enabled": "true",
|
||||
"host": "127.0.0.1",
|
||||
"port": "3306",
|
||||
"user": "root",
|
||||
"password": "pwd",
|
||||
"database": "db",
|
||||
"table": "messages",
|
||||
"messages_table": "messages",
|
||||
"conversation_table": "conversations",
|
||||
"connect_timeout": "5",
|
||||
},
|
||||
)
|
||||
storage = MessageStorage()
|
||||
storage._inited = True
|
||||
|
||||
steps = [
|
||||
{"fetchone_result": None},
|
||||
{},
|
||||
{},
|
||||
]
|
||||
monkeypatch.setattr(storage, "_get_conn", lambda: _ScriptedConn(steps))
|
||||
|
||||
created = storage.create_conversation(
|
||||
conversation_id="cid-compat",
|
||||
user="tester",
|
||||
name="new name",
|
||||
status="normal",
|
||||
introduction=None,
|
||||
created_at=1000,
|
||||
updated_at=1000,
|
||||
)
|
||||
|
||||
assert created is True
|
||||
assert "information_schema.columns" in steps[0]["sql"]
|
||||
assert steps[0]["params"] == ("db", "conversations", "name")
|
||||
assert "ALTER TABLE conversations" in steps[1]["sql"]
|
||||
assert "ADD COLUMN name VARCHAR(255)" in steps[1]["sql"]
|
||||
assert "INSERT INTO conversations" in steps[2]["sql"]
|
||||
assert steps[2]["params"][0] == "cid-compat"
|
||||
assert steps[2]["params"][7] == "cid-compat"
|
||||
assert steps[2]["params"][9] == "new name"
|
||||
@@ -0,0 +1,23 @@
|
||||
from schemas.messages import MessagesDTO
|
||||
|
||||
|
||||
def test_messages_schema_matches_required_fields():
|
||||
dto = MessagesDTO(
|
||||
message_id="mid-1",
|
||||
conversation_id="cid-1",
|
||||
user="tester",
|
||||
query="hello",
|
||||
answer="world",
|
||||
feedback=None,
|
||||
feedback_content=None,
|
||||
created_at=1,
|
||||
updated_at=1,
|
||||
log={"trace_id": "t-1"},
|
||||
)
|
||||
|
||||
assert dto.message_id == "mid-1"
|
||||
assert dto.conversation_id == "cid-1"
|
||||
assert dto.query == "hello"
|
||||
assert dto.answer == "world"
|
||||
assert dto.log["trace_id"] == "t-1"
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
import asyncio
|
||||
|
||||
from services.integrations import nacos_service
|
||||
from services.integrations.nacos_service import NacosConfig, NacosManager, ServiceConfig, load_service_config
|
||||
|
||||
|
||||
class _FakeConfigParser:
|
||||
def get(self, section, option, fallback=None):
|
||||
if section == "app" and option == "host":
|
||||
return "0.0.0.0"
|
||||
if section == "app" and option == "service_name":
|
||||
return fallback
|
||||
if section == "app" and option == "version":
|
||||
return "1.2.3"
|
||||
if section == "app" and option == "model_section":
|
||||
return "qwen-80b"
|
||||
return fallback
|
||||
|
||||
def getint(self, section, option, fallback=None):
|
||||
if section == "app" and option == "port":
|
||||
return 8000
|
||||
return fallback
|
||||
|
||||
|
||||
class _RetryNacosManager(NacosManager):
|
||||
def __init__(self, nacos_config: NacosConfig, service_config: ServiceConfig):
|
||||
super().__init__(nacos_config, service_config)
|
||||
self.register_attempts = 0
|
||||
|
||||
def register_service(self) -> bool:
|
||||
self.register_attempts += 1
|
||||
# First attempt fails, later attempts succeed.
|
||||
self.is_registered = self.register_attempts >= 2
|
||||
return self.is_registered
|
||||
|
||||
|
||||
def test_load_service_config_streaming_metadata_and_service_name_fallback(monkeypatch):
|
||||
monkeypatch.setattr(nacos_service.Config, "_config", _FakeConfigParser())
|
||||
monkeypatch.setattr(nacos_service.Config, "DEFAULT_MODEL_SECTION", "gpt-4o")
|
||||
monkeypatch.setattr(nacos_service.Config, "get_section", lambda section: {} if section == "metadata" else {})
|
||||
monkeypatch.setattr(nacos_service, "_get_local_ip", lambda: "10.0.0.8")
|
||||
|
||||
cfg = load_service_config()
|
||||
|
||||
assert cfg.service_name == "apbo-boat-agent"
|
||||
assert cfg.ip == "10.0.0.8"
|
||||
assert cfg.metadata["streaming"] == "true"
|
||||
|
||||
|
||||
def test_nacos_manager_start_keeps_retry_loop_when_first_register_fails():
|
||||
nacos_cfg = NacosConfig(
|
||||
enabled=True,
|
||||
server_addresses="localhost:8848",
|
||||
namespace="public",
|
||||
group_name="DEFAULT_GROUP",
|
||||
cluster_name="DEFAULT",
|
||||
username=None,
|
||||
password=None,
|
||||
heartbeat_interval=1,
|
||||
weight=1.0,
|
||||
ephemeral=True,
|
||||
register_port=None,
|
||||
)
|
||||
service_cfg = ServiceConfig(
|
||||
service_name="apbo-boat-agent",
|
||||
host="0.0.0.0",
|
||||
port=8000,
|
||||
ip="10.0.0.8",
|
||||
metadata={},
|
||||
)
|
||||
|
||||
manager = _RetryNacosManager(nacos_cfg, service_cfg)
|
||||
|
||||
async def _run_case():
|
||||
await manager.start()
|
||||
await asyncio.sleep(1.2)
|
||||
await manager.stop()
|
||||
|
||||
asyncio.run(_run_case())
|
||||
|
||||
assert manager.register_attempts >= 2
|
||||
assert manager.is_registered is True
|
||||
|
||||
|
||||
class _CaptureClient:
|
||||
def __init__(self):
|
||||
self.register_calls = []
|
||||
self.heartbeat_calls = []
|
||||
self.remove_calls = []
|
||||
|
||||
def add_naming_instance(self, **kwargs):
|
||||
self.register_calls.append(kwargs)
|
||||
|
||||
def send_heartbeat(self, **kwargs):
|
||||
self.heartbeat_calls.append(kwargs)
|
||||
|
||||
def remove_naming_instance(self, **kwargs):
|
||||
self.remove_calls.append(kwargs)
|
||||
|
||||
|
||||
def test_nacos_manager_uses_register_port_override_for_registry_calls():
|
||||
nacos_cfg = NacosConfig(
|
||||
enabled=True,
|
||||
server_addresses="localhost:8848",
|
||||
namespace="public",
|
||||
group_name="DEFAULT_GROUP",
|
||||
cluster_name="DEFAULT",
|
||||
username=None,
|
||||
password=None,
|
||||
heartbeat_interval=1,
|
||||
weight=1.0,
|
||||
ephemeral=True,
|
||||
register_port=26004,
|
||||
)
|
||||
service_cfg = ServiceConfig(
|
||||
service_name="apbo-boat-agent",
|
||||
host="0.0.0.0",
|
||||
port=8000,
|
||||
ip="10.0.0.8",
|
||||
metadata={},
|
||||
)
|
||||
|
||||
manager = NacosManager(nacos_cfg, service_cfg)
|
||||
manager.client = _CaptureClient()
|
||||
|
||||
assert manager._registration_port() == 26004
|
||||
assert manager.register_service() is True
|
||||
|
||||
manager._send_heartbeat()
|
||||
manager.deregister_service()
|
||||
|
||||
assert manager.client.register_calls[0]["port"] == 26004
|
||||
assert manager.client.heartbeat_calls[0]["port"] == 26004
|
||||
assert manager.client.remove_calls[0]["port"] == 26004
|
||||
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
from services.integrations.ragflow_client import extract_table_name
|
||||
from services.integrations.ragflow_sync import RagflowSync
|
||||
|
||||
|
||||
def test_extract_table_name_normalizes_legacy_multiple_impact_alias():
|
||||
assert extract_table_name({"metadata": {"table": "apbo_tp_multiple_impact"}}) == "apbo_eta_multiple_impact"
|
||||
assert extract_table_name({"table_name": "apbo_tp_multiple_impact"}) == "apbo_eta_multiple_impact"
|
||||
assert extract_table_name({"content": '{"table":"apbo_tp_multiple_impact"}'}) == "apbo_eta_multiple_impact"
|
||||
|
||||
|
||||
def test_sync_table_retrieval_delegates_to_update(monkeypatch):
|
||||
expected = {"ok": True, "source": "table"}
|
||||
monkeypatch.setattr(RagflowSync, "update_table_retrieval_documents", lambda self: expected)
|
||||
|
||||
assert RagflowSync().sync_table_retrieval() == expected
|
||||
|
||||
|
||||
def test_sync_sql_gen_prompts_delegates_to_update(monkeypatch):
|
||||
expected = {"ok": True, "source": "sql_gen"}
|
||||
monkeypatch.setattr(RagflowSync, "update_sql_gen_documents", lambda self: expected)
|
||||
|
||||
assert RagflowSync().sync_sql_gen_prompts() == expected
|
||||
|
||||
@@ -0,0 +1,368 @@
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
from agent.agents.conversation import ConversationAgent
|
||||
from agent.core import nodes
|
||||
from agent.core.state import AgentState
|
||||
from services.core.prompt_manager import PromptManager
|
||||
from services.core.sql_prompt_manager import SqlPromptManager
|
||||
from workflows.workflow_manager import WorkflowManager, WorkflowType
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, content: str, tool_calls=None):
|
||||
self.content = content
|
||||
self.tool_calls = tool_calls or []
|
||||
|
||||
|
||||
class FakeModel:
|
||||
def invoke(self, messages):
|
||||
if len(messages) == 2 and getattr(messages[0], "content", "").startswith("You are a translation and normalization assistant"):
|
||||
return FakeResponse(messages[1].content)
|
||||
if len(messages) == 2 and "table prompt JSON" in getattr(messages[0], "content", ""):
|
||||
user_content = getattr(messages[1], "content", "")
|
||||
lowered = user_content.lower()
|
||||
if "query mode: topn" in lowered or "top 20" in lowered:
|
||||
return FakeResponse("SELECT ship_to_country as country, count(soid) as qty FROM dwd_ai.apbo_eta_ful WHERE data_flag = 'Newest' GROUP BY ship_to_country ORDER BY qty DESC LIMIT 20")
|
||||
if "query mode: aggregate" in lowered:
|
||||
return FakeResponse("SELECT region, count(soid) as qty FROM dwd_ai.apbo_eta_ful WHERE data_flag = 'Newest' GROUP BY region ORDER BY qty DESC")
|
||||
return FakeResponse("SELECT service_order_id FROM dwd_ai.apbo_eta_ful WHERE data_flag = 'Newest'")
|
||||
return FakeResponse("fallback")
|
||||
|
||||
|
||||
class FakeTemplateMatcher:
|
||||
def match(self, normalized_text: str):
|
||||
return {
|
||||
"table_name": "apbo_eta_ful",
|
||||
"candidates": [{"table_name": "apbo_eta_ful", "metadata": {"table": "apbo_eta_ful"}}],
|
||||
"raw": {"query": normalized_text},
|
||||
}
|
||||
|
||||
|
||||
class FakeSqlPromptManager:
|
||||
def get_prompt(self, table_name: str):
|
||||
manager = SqlPromptManager()
|
||||
return manager.get_prompt(table_name)
|
||||
|
||||
|
||||
class EmptyTemplateMatcher:
|
||||
def match(self, normalized_text: str):
|
||||
return {
|
||||
"table_name": None,
|
||||
"candidates": [],
|
||||
"raw": {"query": normalized_text},
|
||||
}
|
||||
|
||||
|
||||
class FirstTurnOnlyTemplateMatcher:
|
||||
def match(self, normalized_text: str):
|
||||
lowered = (normalized_text or "").lower()
|
||||
if "4020438779" in lowered or "query so" in lowered:
|
||||
return FakeTemplateMatcher().match(normalized_text)
|
||||
return EmptyTemplateMatcher().match(normalized_text)
|
||||
|
||||
|
||||
def test_prompt_manager_resolves_project_prompt_file():
|
||||
manager = PromptManager()
|
||||
prompt_text = manager.get("system", "sql_mysql_select_only")
|
||||
assert "expert SQL generator" in prompt_text
|
||||
assert "Only output a single MySQL SELECT statement." in prompt_text
|
||||
|
||||
|
||||
def test_sql_prompt_manager_resolves_project_sql_prompt_dir():
|
||||
manager = SqlPromptManager()
|
||||
prompt = manager.get_prompt("apbo_eta_ful")
|
||||
assert prompt is not None
|
||||
assert prompt["meta"]["data_source"] == "dwd_ai.apbo_eta_ful"
|
||||
|
||||
|
||||
def test_apbo_eta_ful_prompt_uses_identifier_specific_where_filters():
|
||||
prompt = SqlPromptManager().get_prompt("apbo_eta_ful")
|
||||
|
||||
rule_text = prompt["business_logic_rules"]["soid_or_service_order_id"]
|
||||
examples = prompt["examples"]
|
||||
|
||||
assert "禁止使用WHERE soid" in rule_text
|
||||
assert "必须使用(service_order_id = 'xxx' or soid = 'xxx')" in rule_text
|
||||
assert "service_order_id" in examples["history_records_all_fields"]["sql"]
|
||||
assert " or soid " in examples["history_records_all_fields"]["sql"]
|
||||
assert " or soid " in examples["specific_fields_query"]["sql"]
|
||||
assert " or soid " in examples["newest_status_all_fields"]["sql"]
|
||||
|
||||
|
||||
def test_query_mode_and_plan_build_for_topn():
|
||||
prompt = SqlPromptManager().get_prompt("apbo_eta_ful")
|
||||
state = AgentState(
|
||||
messages=[HumanMessage(content="CC为LT ,by country 查询 top 20")],
|
||||
context={"original_input": "CC为LT ,by country 查询 top 20", "normalized_input": "CC=LT by country top 20", "table_name": "apbo_eta_ful"},
|
||||
)
|
||||
state.sql_prompt = prompt
|
||||
state = nodes.classify_query_mode(state)
|
||||
state = nodes.build_sql_plan(state)
|
||||
|
||||
assert state.query_mode == "topn"
|
||||
assert state.sql_plan["query_entities"]["top_n"] == 20
|
||||
assert state.sql_plan["selected_table"] == "apbo_eta_ful"
|
||||
assert "top_n_rules" in state.sql_plan
|
||||
|
||||
|
||||
def test_query_mode_country_does_not_trigger_aggregate():
|
||||
state = AgentState(
|
||||
messages=[HumanMessage(content="查询 AU country 的 DC premier stock backlog")],
|
||||
context={
|
||||
"original_input": "查询 AU country 的 DC premier stock backlog",
|
||||
"normalized_input": "List of backlogs with values for AU Country DC premier stock",
|
||||
},
|
||||
)
|
||||
|
||||
state = nodes.classify_query_mode(state)
|
||||
|
||||
assert state.query_mode == "detail"
|
||||
|
||||
|
||||
def test_query_mode_explicit_aggregate_keyword_still_matches():
|
||||
state = AgentState(
|
||||
messages=[HumanMessage(content="按 country 统计 backlog 数量")],
|
||||
context={
|
||||
"original_input": "按 country 统计 backlog 数量",
|
||||
"normalized_input": "Count backlog quantity by country",
|
||||
},
|
||||
)
|
||||
|
||||
state = nodes.classify_query_mode(state)
|
||||
|
||||
assert state.query_mode == "aggregate"
|
||||
|
||||
|
||||
def test_query_mode_current_keyword_is_not_misclassified_as_topn():
|
||||
state = AgentState(
|
||||
messages=[HumanMessage(content="汇总统计当前的bo数量")],
|
||||
context={
|
||||
"original_input": "汇总统计当前的bo数量",
|
||||
"normalized_input": "Aggregate summary of the current backlog order count",
|
||||
},
|
||||
)
|
||||
|
||||
state = nodes.classify_query_mode(state)
|
||||
|
||||
assert state.query_mode == "aggregate"
|
||||
assert state.query_entities.get("top_n") is None
|
||||
|
||||
|
||||
def test_match_table_falls_back_to_config_default(monkeypatch):
|
||||
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: EmptyTemplateMatcher())
|
||||
monkeypatch.setattr(
|
||||
"agent.core.nodes.Config.get_section",
|
||||
lambda section: {"default_table_name": "apbo_eta_ful"} if section == "ragflow" else {},
|
||||
)
|
||||
|
||||
state = AgentState(
|
||||
messages=[HumanMessage(content="random question")],
|
||||
context={"original_input": "random question", "normalized_input": "random question"},
|
||||
)
|
||||
|
||||
state = nodes.match_table(state)
|
||||
|
||||
assert state.table_name == "apbo_eta_ful"
|
||||
assert state.context["table_match_fallback"] == "config_default"
|
||||
assert state.context["default_table_name"] == "apbo_eta_ful"
|
||||
|
||||
|
||||
def test_generate_response_falls_back_to_model_when_sr_api_result_empty():
|
||||
class EmptyResultModel:
|
||||
def invoke(self, messages):
|
||||
return FakeResponse("未查询到符合条件的数据,请尝试调整筛选条件。")
|
||||
|
||||
state = AgentState(
|
||||
messages=[HumanMessage(content="查询 AU 的 backlog")],
|
||||
context={
|
||||
"original_input": "查询 AU 的 backlog",
|
||||
"normalized_input": "Query backlog for AU",
|
||||
"final_sql": "SELECT * FROM dwd_ai.apbo_eta_ful WHERE ship_to_country='AU'",
|
||||
"sr_api_result": '{"status_code": 200, "text": "{\\"code\\":\\"0\\",\\"data\\":[],\\"msg\\":\\"操作成功\\",\\"total\\":0}"}',
|
||||
},
|
||||
)
|
||||
state.final_sql = state.context["final_sql"]
|
||||
state.sr_api_result = state.context["sr_api_result"]
|
||||
|
||||
state = nodes.generate_response(state, EmptyResultModel())
|
||||
|
||||
assert state.messages[-1].content.startswith("未查询到符合条件的数据")
|
||||
assert state.context["response_source"] == "model_empty_result_fallback"
|
||||
assert state.context["is_empty_result"] is True
|
||||
|
||||
|
||||
def test_conversation_agent_run_keeps_final_sql_in_context(monkeypatch):
|
||||
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
|
||||
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FakeTemplateMatcher())
|
||||
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
|
||||
|
||||
agent = ConversationAgent()
|
||||
result = agent.run("查询 SO 4020438779 的 eta 信息", skip_sr_api=True)
|
||||
|
||||
assert result["context"]["table_name"] == "apbo_eta_ful"
|
||||
assert result["context"]["query_mode"] == "detail"
|
||||
assert result["context"]["final_sql"].startswith("SELECT")
|
||||
assert result["messages"][-1].content == result["context"]["final_sql"]
|
||||
|
||||
|
||||
def test_conversation_agent_requires_explicit_memory_for_follow_up(monkeypatch):
|
||||
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
|
||||
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FirstTurnOnlyTemplateMatcher())
|
||||
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
|
||||
|
||||
agent = ConversationAgent()
|
||||
first = agent.run("查询 SO 4020438779 的 eta 信息", skip_sr_api=True)
|
||||
second = agent.run("改成 by country top 20", skip_sr_api=True)
|
||||
third = agent.run(
|
||||
"改成 by country top 20",
|
||||
skip_sr_api=True,
|
||||
conversation_history=first["conversation_history"],
|
||||
last_context=first["context"].get("last_context"),
|
||||
)
|
||||
|
||||
assert first["context"]["table_name"] == "apbo_eta_ful"
|
||||
assert second["context"]["table_name"] == "apbo_eta_ful"
|
||||
assert second["context"]["table_match_fallback"] == "config_default"
|
||||
assert third["context"]["table_name"] == "apbo_eta_ful"
|
||||
assert third["context"]["table_match_fallback"] == "last_context"
|
||||
assert third["context"]["query_mode"] == "topn"
|
||||
assert "LIMIT 20" in third["context"]["final_sql"]
|
||||
|
||||
|
||||
def test_workflow_manager_follow_up_does_not_reuse_session_last_context(monkeypatch):
|
||||
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
|
||||
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FirstTurnOnlyTemplateMatcher())
|
||||
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
|
||||
|
||||
manager = WorkflowManager(enable_multi_turn=False)
|
||||
first = manager.execute_workflow(
|
||||
WorkflowType.CONVERSATION,
|
||||
"查询 SO 4020438779 的 eta 信息",
|
||||
session_id="cid-1",
|
||||
skip_sr_api=True,
|
||||
)
|
||||
second = manager.execute_workflow(
|
||||
WorkflowType.CONVERSATION,
|
||||
"改成 by country top 20",
|
||||
session_id="cid-1",
|
||||
skip_sr_api=True,
|
||||
)
|
||||
|
||||
first_context = first["result"]["context"]
|
||||
second_context = second["result"]["context"]
|
||||
session_info = manager.get_session_info("cid-1")
|
||||
|
||||
assert first_context["table_name"] == "apbo_eta_ful"
|
||||
assert second_context["table_name"] == "apbo_eta_ful"
|
||||
assert second_context["table_match_fallback"] == "config_default"
|
||||
assert second_context["query_mode"] == "topn"
|
||||
assert "LIMIT 20" in second_context["final_sql"]
|
||||
assert "last_context" not in session_info
|
||||
assert "conversation_history" not in session_info
|
||||
|
||||
|
||||
def test_workflow_manager_follow_up_reuses_session_last_context_when_enabled(monkeypatch):
|
||||
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
|
||||
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FirstTurnOnlyTemplateMatcher())
|
||||
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
|
||||
|
||||
manager = WorkflowManager(enable_multi_turn=True)
|
||||
first = manager.execute_workflow(
|
||||
WorkflowType.CONVERSATION,
|
||||
"查询 SO 4020438779 的 eta 信息",
|
||||
session_id="cid-enabled",
|
||||
skip_sr_api=True,
|
||||
)
|
||||
second = manager.execute_workflow(
|
||||
WorkflowType.CONVERSATION,
|
||||
"改成 by country top 20",
|
||||
session_id="cid-enabled",
|
||||
skip_sr_api=True,
|
||||
)
|
||||
|
||||
first_context = first["result"]["context"]
|
||||
second_context = second["result"]["context"]
|
||||
session_info = manager.get_session_info("cid-enabled")
|
||||
|
||||
assert first_context["table_name"] == "apbo_eta_ful"
|
||||
assert second_context["table_name"] == "apbo_eta_ful"
|
||||
assert second_context["table_match_fallback"] == "last_context"
|
||||
assert second_context["query_mode"] == "topn"
|
||||
assert "LIMIT 20" in second_context["final_sql"]
|
||||
assert session_info["last_context"]["table_name"] == "apbo_eta_ful"
|
||||
assert len(session_info["conversation_history"]) >= 4
|
||||
|
||||
|
||||
def test_workflow_manager_isolates_conversation_memory_by_session(monkeypatch):
|
||||
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
|
||||
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FirstTurnOnlyTemplateMatcher())
|
||||
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
|
||||
|
||||
manager = WorkflowManager(enable_multi_turn=False)
|
||||
manager.execute_workflow(
|
||||
WorkflowType.CONVERSATION,
|
||||
"查询 SO 4020438779 的 eta 信息",
|
||||
session_id="session-a",
|
||||
skip_sr_api=True,
|
||||
)
|
||||
second = manager.execute_workflow(
|
||||
WorkflowType.CONVERSATION,
|
||||
"改成 by country top 20",
|
||||
session_id="session-b",
|
||||
skip_sr_api=True,
|
||||
)
|
||||
|
||||
assert second["result"]["context"]["table_name"] == "apbo_eta_ful"
|
||||
assert second["result"]["context"]["table_match_fallback"] == "config_default"
|
||||
|
||||
|
||||
def test_conversation_agent_run_is_silent_without_debug_prints(monkeypatch, capsys):
|
||||
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
|
||||
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FakeTemplateMatcher())
|
||||
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
|
||||
|
||||
agent = ConversationAgent()
|
||||
agent.run("查询 SO 4020438779 的 eta 信息", skip_sr_api=True)
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert captured.out == ""
|
||||
|
||||
|
||||
def test_conversation_agent_run_prints_node_trace_when_enabled(monkeypatch, capsys):
|
||||
monkeypatch.setattr("agent.core.base_agent.create_chat_model", lambda model_section=None: FakeModel())
|
||||
monkeypatch.setattr("agent.core.nodes.get_template_matcher", lambda: FakeTemplateMatcher())
|
||||
monkeypatch.setattr("agent.core.nodes.get_sql_prompt_manager", lambda: FakeSqlPromptManager())
|
||||
|
||||
agent = ConversationAgent()
|
||||
agent.run("查询 SO 4020438779 的 eta 信息", skip_sr_api=True, debug_node_trace=True)
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert "[process_input][in]" in captured.out
|
||||
assert "[generate_response][out] source=final_sql" in captured.out
|
||||
|
||||
|
||||
def test_generate_sql_trace_outputs_full_sql_when_enabled(capsys):
|
||||
class LongSqlModel:
|
||||
def invoke(self, messages):
|
||||
return FakeResponse("SELECT " + ", ".join(f"col_{idx}" for idx in range(150)) + " FROM dwd_ai.apbo_eta_ful")
|
||||
|
||||
state = AgentState(
|
||||
messages=[HumanMessage(content="查询 long sql")],
|
||||
context={
|
||||
"original_input": "查询 long sql",
|
||||
"normalized_input": "Query long sql",
|
||||
"table_name": "apbo_eta_ful",
|
||||
"sql_plan": {"selected_table": "apbo_eta_ful"},
|
||||
"debug_node_trace": True,
|
||||
},
|
||||
)
|
||||
state.sql_prompt = SqlPromptManager().get_prompt("apbo_eta_ful")
|
||||
|
||||
state = nodes.generate_sql(state, LongSqlModel())
|
||||
|
||||
captured = capsys.readouterr()
|
||||
assert "[generate_sql][out] sql=" in captured.out
|
||||
assert "col_149" in captured.out
|
||||
assert "..." not in captured.out.split("[generate_sql][out] sql=", 1)[1]
|
||||
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
import json
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from api.dependencies import get_workflow_manager
|
||||
from api.endpoints import router
|
||||
from schemas.chat_message_response import ChatMessageResponseDTO
|
||||
from workflows.workflow_manager import WorkflowType
|
||||
|
||||
|
||||
class StubWorkflowManager:
|
||||
def execute_workflow(self, workflow_type, user_input, session_id=None, **kwargs):
|
||||
return {
|
||||
"session_id": session_id or "cid-1",
|
||||
"workflow_type": workflow_type.value if isinstance(workflow_type, WorkflowType) else str(workflow_type),
|
||||
"result": {"context": {"final_sql": "SELECT 1"}},
|
||||
}
|
||||
|
||||
|
||||
class FakeSrApiTool:
|
||||
def run(self, payload):
|
||||
return '{"total": 1, "data": [{"value": 1, "etl_time": "2026-03-20 11:27:53"}]}'
|
||||
|
||||
|
||||
class FakeEmptySrApiTool:
|
||||
def run(self, payload):
|
||||
return '{"total": 0, "data": []}'
|
||||
|
||||
|
||||
class DisabledMessageStorage:
|
||||
enabled = False
|
||||
|
||||
def save_message(self, **kwargs):
|
||||
return False
|
||||
|
||||
|
||||
class CaptureMessageStorage:
|
||||
enabled = True
|
||||
|
||||
def __init__(self):
|
||||
self.saved = []
|
||||
self.existing = {
|
||||
"cid-test-1": {
|
||||
"conversation_id": "cid-test-1",
|
||||
"user": "tester",
|
||||
"name": "test",
|
||||
"status": "normal",
|
||||
"created_at": 1,
|
||||
"updated_at": 1,
|
||||
}
|
||||
}
|
||||
|
||||
def get_conversation_by_id(self, conversation_id):
|
||||
return self.existing.get(conversation_id)
|
||||
|
||||
def update_conversation_updated_at(self, conversation_id, updated_at):
|
||||
if conversation_id not in self.existing:
|
||||
return False
|
||||
self.existing[conversation_id]["updated_at"] = updated_at
|
||||
return True
|
||||
|
||||
def save_message(self, **kwargs):
|
||||
self.saved.append(kwargs)
|
||||
return True
|
||||
|
||||
|
||||
def _build_client(workflow_manager) -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
app.dependency_overrides[get_workflow_manager] = lambda: workflow_manager
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _payload():
|
||||
return {
|
||||
"query": "test",
|
||||
"inputs": {},
|
||||
"response_mode": "streaming",
|
||||
"user": "tester",
|
||||
"conversation_id": "cid-test-1",
|
||||
"files": [],
|
||||
}
|
||||
|
||||
|
||||
def test_stream_returns_chat_message_response_dto_chunks(monkeypatch):
|
||||
monkeypatch.setattr("api.endpoints.SrApiQueryTool", FakeSrApiTool)
|
||||
storage = CaptureMessageStorage()
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: storage)
|
||||
client = _build_client(StubWorkflowManager())
|
||||
|
||||
with client.stream("POST", "/api/workflows/stream", json=_payload()) as response:
|
||||
body = response.read().decode("utf-8")
|
||||
|
||||
assert response.status_code == 200
|
||||
lines = [line for line in body.splitlines() if line.strip()]
|
||||
assert all(not line.startswith("event:") for line in lines)
|
||||
assert lines[0].startswith("data: ")
|
||||
|
||||
dto = ChatMessageResponseDTO.model_validate_json(lines[0][len("data: "):])
|
||||
assert dto.event == "message"
|
||||
assert dto.conversation_id == "cid-test-1"
|
||||
assert "<strong>Question:</strong> test" in dto.answer
|
||||
assert "<table>" in dto.answer
|
||||
assert "<th>value</th>" in dto.answer
|
||||
assert "<td>1</td>" in dto.answer
|
||||
assert "<strong>Rows:</strong> 1" in dto.answer
|
||||
assert "<strong>Data Version:</strong> 2026-03-20 11:27:53" in dto.answer
|
||||
assert dto.task_id
|
||||
assert dto.message_id
|
||||
end_dto = ChatMessageResponseDTO.model_validate_json(lines[-1][len("data: "):])
|
||||
assert end_dto.event == "message_end"
|
||||
assert end_dto.answer == ""
|
||||
assert len(storage.saved) == 1
|
||||
assert storage.saved[0]["answer"] == dto.answer
|
||||
assert storage.saved[0]["logs"][0].startswith("stream.start")
|
||||
assert any(item.startswith("stream.sql_executed") for item in storage.saved[0]["logs"])
|
||||
|
||||
|
||||
def test_stream_returns_no_data_when_sql_data_is_empty(monkeypatch):
|
||||
monkeypatch.setattr("api.endpoints.SrApiQueryTool", FakeEmptySrApiTool)
|
||||
storage = CaptureMessageStorage()
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: storage)
|
||||
client = _build_client(StubWorkflowManager())
|
||||
|
||||
with client.stream("POST", "/api/workflows/stream", json=_payload()) as response:
|
||||
body = response.read().decode("utf-8")
|
||||
|
||||
assert response.status_code == 200
|
||||
lines = [line for line in body.splitlines() if line.strip()]
|
||||
dto = ChatMessageResponseDTO.model_validate_json(lines[0][len("data: "):])
|
||||
assert "Question: test" in dto.answer
|
||||
assert "No data~" in dto.answer
|
||||
assert "Rows: 0" in dto.answer
|
||||
assert "Data Version: Unknown" in dto.answer
|
||||
assert "<div" not in dto.answer
|
||||
assert "<strong>" not in dto.answer
|
||||
end_dto = ChatMessageResponseDTO.model_validate_json(lines[-1][len("data: "):])
|
||||
assert end_dto.event == "message_end"
|
||||
assert storage.saved[0]["answer"] == dto.answer
|
||||
assert any(item.startswith("stream.sql_executed") for item in storage.saved[0]["logs"])
|
||||
|
||||
|
||||
def test_stream_error_still_returns_plain_text(monkeypatch):
|
||||
monkeypatch.setattr("api.endpoints.SrApiQueryTool", FakeSrApiTool)
|
||||
monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage())
|
||||
client = _build_client(StubWorkflowManager())
|
||||
|
||||
with client.stream("POST", "/api/workflows/stream", json={**_payload(), "query": " "}) as response:
|
||||
payload = json.loads(response.read().decode("utf-8"))
|
||||
|
||||
assert response.status_code == 400
|
||||
assert payload["detail"]["code"] == "INVALID_REQUEST"
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
from services.core.template_matcher import TemplateMatcher
|
||||
|
||||
|
||||
class FakeRagflowClient:
|
||||
def retrieve(self, normalized_text, top_k=3, dataset_id=None):
|
||||
return {
|
||||
"data": [
|
||||
{"content": '{"table": "apbo_eta_region_report", "templates": []}'},
|
||||
{"content": '{"table": "apbo_eta_ful", "templates": ["BO list"]}'},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
class FakeRagflowClientForRanking:
|
||||
def retrieve(self, normalized_text, top_k=3, dataset_id=None):
|
||||
return {
|
||||
"data": [
|
||||
{"content": '{"table": "apbo_eta_region_report", "templates": ["region usage"]}'},
|
||||
{"content": '{"table": "apbo_eta_ful", "templates": ["eta detail"]}'},
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def test_template_matcher_skips_empty_template_tables(monkeypatch):
|
||||
monkeypatch.setattr("services.core.template_matcher.RagflowClient", FakeRagflowClient)
|
||||
monkeypatch.setattr(
|
||||
"services.core.template_matcher.Config.get_section",
|
||||
lambda section: {"table_retrieval_dataset_id": "ds-1", "retrieval_top_k": 3} if section == "ragflow" else {},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
TemplateMatcher,
|
||||
"_load_non_empty_table_names",
|
||||
lambda self: {"apbo_eta_ful"},
|
||||
)
|
||||
|
||||
matcher = TemplateMatcher()
|
||||
matched = matcher.match("region anz pn 5m20s27936")
|
||||
|
||||
assert matched["table_name"] == "apbo_eta_ful"
|
||||
assert [c["table_name"] for c in matched["candidates"]] == ["apbo_eta_ful"]
|
||||
|
||||
|
||||
def test_template_matcher_reranks_by_explicit_filter_field_coverage(monkeypatch):
|
||||
monkeypatch.setattr("services.core.template_matcher.RagflowClient", FakeRagflowClientForRanking)
|
||||
monkeypatch.setattr(
|
||||
"services.core.template_matcher.Config.get_section",
|
||||
lambda section: {"table_retrieval_dataset_id": "ds-1", "retrieval_top_k": 3} if section == "ragflow" else {},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
TemplateMatcher,
|
||||
"_load_non_empty_table_names",
|
||||
lambda self: {"apbo_eta_region_report", "apbo_eta_ful"},
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
TemplateMatcher,
|
||||
"_load_table_terms",
|
||||
lambda self, table_name: {
|
||||
"apbo_eta_region_report": {"region", "ib", "usage_qty_8_week"},
|
||||
"apbo_eta_ful": {"region", "pn", "part_number", "aging_range"},
|
||||
}[table_name],
|
||||
)
|
||||
|
||||
matcher = TemplateMatcher()
|
||||
matched = matcher.match("How many records are there for region = ANZ, PN = 5M20S27936, and Aging_range = 15-21D?")
|
||||
|
||||
assert matched["table_name"] == "apbo_eta_ful"
|
||||
assert matched["candidates"][0]["table_name"] == "apbo_eta_ful"
|
||||
assert matched["candidates"][0]["rank_score"] > matched["candidates"][1]["rank_score"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user