Files
AIRegulation-DocAnalysis/backend/tests/mcp/test_search_regulations_tool.py
T

78 lines
2.5 KiB
Python
Raw Normal View History

"""Unit tests for the search_regulations MCP tool function.
Mocks AgentConversationService so no real retrieval/LLM call happens —
verifies only the protocol-adapter contract: correct call shape in,
correct dict shape out.
"""
from __future__ import annotations
from dataclasses import dataclass
from unittest.mock import MagicMock, patch
@dataclass
class _FakeSource:
"""Minimal stand-in for a real Source dataclass (only __dict__ is used)."""
doc_id: str
doc_title: str
score: float
@dataclass
class _FakeAnswerResult:
"""Minimal stand-in for AnswerResult — only .answer/.sources are read."""
answer: str
sources: list
def test_search_regulations_calls_agent_ask_without_session():
"""search_regulations must call ask() with no session_id (stateless search)."""
from app.mcp.server import search_regulations
fake_service = MagicMock()
fake_service.ask.return_value = (
None,
_FakeAnswerResult(answer="国六排放标准要求...", sources=[_FakeSource("doc-1", "国六标准", 0.9)]),
)
with patch("app.mcp.server.get_agent_conversation_service", return_value=fake_service):
search_regulations(query="国六排放标准最新要求", top_k=3)
fake_service.ask.assert_called_once_with(query="国六排放标准最新要求", top_k=3)
assert "session_id" not in fake_service.ask.call_args.kwargs
def test_search_regulations_shapes_response_dict():
"""The returned dict must expose 'answer' and 'sources' (list of plain dicts)."""
from app.mcp.server import search_regulations
fake_service = MagicMock()
fake_service.ask.return_value = (
None,
_FakeAnswerResult(answer="答案文本", sources=[_FakeSource("doc-2", "国标GB1589", 0.8)]),
)
with patch("app.mcp.server.get_agent_conversation_service", return_value=fake_service):
result = search_regulations(query="q")
assert result == {
"answer": "答案文本",
"sources": [{"doc_id": "doc-2", "doc_title": "国标GB1589", "score": 0.8}],
}
def test_search_regulations_default_top_k():
"""top_k defaults to 5 when the caller omits it."""
from app.mcp.server import search_regulations
fake_service = MagicMock()
fake_service.ask.return_value = (None, _FakeAnswerResult(answer="a", sources=[]))
with patch("app.mcp.server.get_agent_conversation_service", return_value=fake_service):
search_regulations(query="q")
assert fake_service.ask.call_args.kwargs["top_k"] == 5