22import asyncio
33import os
44from contextlib import AsyncExitStack
5+ from datetime import timedelta
56from typing import Any , Dict , List , Literal , Optional
67
78from mcp import ClientSession , ListToolsResult , StdioServerParameters
2526DEFAULT_SSE_READ_TIMEOUT = 60 * 5
2627TOOL_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
2933class 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