2026-07-29 13:00:52 +08:00
|
|
|
"""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
|
2026-07-29 13:00:52 +08:00
|
|
|
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.
|
2026-07-29 13:00:52 +08:00
|
|
|
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
|