import unittest from contextlib import asynccontextmanager from unittest.mock import patch from models import RuntimeMcpServer from schemas import McpServerUpsert from services.tools.mcp_client import McpToolClient class McpTransportTests(unittest.IsolatedAsyncioTestCase): def test_server_schema_accepts_sse(self): server = McpServerUpsert( name="旧版 MCP 服务", transport="sse", url="https://mcp.example.com/sse", ) self.assertEqual(server.transport, "sse") async def test_sse_transport_uses_sdk_sse_client(self): captured: dict = {} @asynccontextmanager async def fake_sse_client(url, **kwargs): captured["url"] = url captured.update(kwargs) yield "read-stream", "write-stream" server = RuntimeMcpServer( id="mcp_sse", transport="sse", url="https://mcp.example.com/sse", ) client = McpToolClient() with patch("services.tools.mcp_client.sse_client", fake_sse_client): async with client._transport_streams( server, {"Authorization": "Bearer test"}, 20, ) as streams: self.assertEqual(streams, ("read-stream", "write-stream")) self.assertEqual(captured["url"], server.url) self.assertEqual(captured["headers"], {"Authorization": "Bearer test"}) self.assertEqual(captured["timeout"], 20) self.assertEqual(captured["sse_read_timeout"], 20) if __name__ == "__main__": unittest.main()