This commit is contained in:
2026-03-24 18:07:22 +08:00
parent e062368ef2
commit 9a16f738d8
121 changed files with 8904 additions and 3940 deletions
+12
View File
@@ -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`:脚本相关测试
+1
View File
@@ -0,0 +1 @@
"""测试模块"""
+51
View File
@@ -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
View File
@@ -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):
"""测试配置校验"""
+99
View File
@@ -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)
+158
View File
@@ -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
+230
View File
@@ -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"
+47
View File
@@ -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)
+166
View File
@@ -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
+201
View File
@@ -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"
+148
View File
@@ -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)
+77
View File
@@ -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 == "结果不准确"
+456
View File
@@ -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"
+23
View File
@@ -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"
+136
View File
@@ -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
+23
View File
@@ -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
View File
+368
View File
@@ -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]
+154
View File
@@ -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"
+70
View File
@@ -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"]