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

99 lines
3.3 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
2026-07-29 17:11:54 +08:00
import asyncio
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)."""
2026-07-29 17:11:54 +08:00
# A dataclass, not a MagicMock: the adapter serializes sources via
# source.__dict__, and a MagicMock's __dict__ is full of internal mock
# attributes, which would make the assertions meaningless.
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
2026-07-29 17:11:54 +08:00
def test_advertised_schema_bounds_top_k_and_query():
"""The advertised JSON schema must carry the same bounds as AskRequest.
Bounds declared via Annotated are what the SDK validates against and what
clients see, so asserting on the generated schema is the only way to catch
a regression that silently drops them.
"""
from app.mcp.server import mcp
schema = asyncio.run(mcp.list_tools())[0].input_schema["properties"]
assert schema["top_k"]["minimum"] == 1
assert schema["top_k"]["maximum"] == 20
assert schema["query"]["minLength"] == 1
assert schema["query"]["maxLength"] == 2000