Skip to content

Commit fa7132a

Browse files
suluyanasuluyan
andauthored
feat: support streamablehttp and websocket mcp servers (#696)
Co-authored-by: suluyan <suluyan.sly@alibaba-inc.com>
1 parent 86cfbc8 commit fa7132a

1 file changed

Lines changed: 57 additions & 17 deletions

File tree

ms_agent/tools/mcp_client.py

Lines changed: 57 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
import asyncio
33
import os
44
from contextlib import AsyncExitStack
5+
from datetime import timedelta
56
from typing import Any, Dict, List, Literal, Optional
67

78
from mcp import ClientSession, ListToolsResult, StdioServerParameters
@@ -25,6 +26,9 @@
2526
DEFAULT_SSE_READ_TIMEOUT = 60 * 5
2627
TOOL_CALL_TIMEOUT = os.getenv('TOOL_CALL_TIMEOUT', 15)
2728

29+
DEFAULT_STREAMABLE_HTTP_TIMEOUT = timedelta(seconds=30)
30+
DEFAULT_STREAMABLE_HTTP_SSE_READ_TIMEOUT = timedelta(seconds=60 * 5)
31+
2832

2933
class MCPClient(ToolBase):
3034
"""MCP client for all mcp tools
@@ -112,26 +116,62 @@ def print_tools(server_name: str, tools: ListToolsResult):
112116

113117
async def connect_to_server(self, server_name: str, **kwargs):
114118
logger.info(f'connect to {server_name}')
119+
# transport: stdio, sse, streamable_http, websocket
120+
transport = kwargs.get('transport') or kwargs.get('type')
115121
command = kwargs.get('command')
116122
url = kwargs.get('url')
117123
session_kwargs = kwargs.get('session_kwargs')
118124
if url:
119-
# transport: 'sse'
120-
sse_transport = await self.exit_stack.enter_async_context(
121-
sse_client(
122-
url, kwargs.get('headers'),
123-
kwargs.get('timeout', DEFAULT_HTTP_TIMEOUT),
124-
kwargs.get('sse_read_timeout', DEFAULT_SSE_READ_TIMEOUT)))
125-
read, write = sse_transport
125+
if transport == 'streamable_http':
126+
try:
127+
from mcp.client.streamable_http import streamablehttp_client
128+
except ImportError:
129+
raise ImportError(
130+
'Could not import streamablehttp_client. '
131+
'To use streamable http connections, please upgrade to the latest version of mcp with: '
132+
"'pip install -U mcp'") from None
133+
httpx_client_factory = kwargs.get('httpx_client_factory')
134+
other_kwargs = {}
135+
if httpx_client_factory is not None:
136+
other_kwargs['httpx_client_factory'] = httpx_client_factory
137+
streamable_transport = await self.exit_stack.enter_async_context(
138+
streamablehttp_client(
139+
url,
140+
headers=kwargs.get('headers'),
141+
timeout=kwargs.get('timeout',
142+
DEFAULT_STREAMABLE_HTTP_TIMEOUT),
143+
sse_read_timeout=kwargs.get(
144+
'sse_read_timeout',
145+
DEFAULT_STREAMABLE_HTTP_SSE_READ_TIMEOUT),
146+
**other_kwargs))
147+
read, write, _ = streamable_transport
148+
149+
elif transport == 'websocket':
150+
try:
151+
from mcp.client.websocket import websocket_client
152+
except ImportError:
153+
raise ImportError(
154+
'Could not import websocket_client. '
155+
'To use Websocket connections, please install the required dependency with: '
156+
"'pip install mcp[ws]' or 'pip install websockets'"
157+
) from None
158+
websocket_transport = await self.exit_stack.enter_async_context(
159+
websocket_client(url))
160+
read, write = websocket_transport
161+
162+
else:
163+
sse_transport = await self.exit_stack.enter_async_context(
164+
sse_client(
165+
url, kwargs.get('headers'),
166+
kwargs.get('timeout', DEFAULT_HTTP_TIMEOUT),
167+
kwargs.get('sse_read_timeout',
168+
DEFAULT_SSE_READ_TIMEOUT)))
169+
read, write = sse_transport
170+
126171
session_kwargs = session_kwargs or {}
127172
session = await self.exit_stack.enter_async_context(
128173
ClientSession(read, write, **session_kwargs))
129174

130-
await session.initialize()
131-
# Store session
132-
self.sessions[server_name] = session
133-
# List available tools
134-
self.print_tools(server_name, await session.list_tools())
135175
elif command:
136176
# transport: 'stdio'
137177
args = kwargs.get('args')
@@ -151,14 +191,14 @@ async def connect_to_server(self, server_name: str, **kwargs):
151191
stdio_client(server_params))
152192
session = await self.exit_stack.enter_async_context(
153193
ClientSession(stdio, write))
154-
await session.initialize()
155-
156-
# Store session
157-
self.sessions[server_name] = session
158-
self.print_tools(server_name, await session.list_tools())
159194
else:
160195
raise ValueError(
161196
"'url' or 'command' parameter is required for connection")
197+
198+
await session.initialize()
199+
# Store session
200+
self.sessions[server_name] = session
201+
self.print_tools(server_name, await session.list_tools())
162202
return server_name
163203

164204
async def connect(self):

0 commit comments

Comments
 (0)