"""Tests for MCP server module.""" import pytest from unittest.mock import Mock, patch from src.mcp_server import MCPServer, create_mcp_server from src.model_handler import ModelHandler class TestMCPServer: """Test cases for MCP server.""" @patch('src.model_handler.gr.load') def test_mcp_server_initialization(self, mock_gr_load): """Test MCP server initialization.""" mock_gr_load.return_value = Mock() model_handler = ModelHandler() server = create_mcp_server(model_handler) assert server is not None assert server.model_handler == model_handler assert isinstance(server.history, list) @patch('src.model_handler.gr.load') def test_get_presets_tool(self, mock_gr_load): """Test get_presets tool.""" mock_gr_load.return_value = Mock() model_handler = ModelHandler() server = MCPServer(model_handler) result = server.tool_get_presets() assert "presets" in result assert "details" in result assert isinstance(result["presets"], list) @patch('src.model_handler.gr.load') def test_get_history_tool(self, mock_gr_load): """Test get_history tool.""" mock_gr_load.return_value = Mock() model_handler = ModelHandler() server = MCPServer(model_handler) # Add some test history server.history = [ {"timestamp": 1234567890, "prompt": "test1"}, {"timestamp": 1234567891, "prompt": "test2"}, ] result = server.tool_get_history(limit=5) assert "total" in result assert "entries" in result assert result["total"] == 2 assert len(result["entries"]) == 2 if __name__ == "__main__": pytest.main([__file__])