"""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 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).""" # 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 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