"""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