Coverage for open_webui/utils/mcp/client.py: 17%
117 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import asyncio
2import logging
3from contextlib import AsyncExitStack
4from typing import Optional
6log = logging.getLogger(__name__)
8import anyio
9import httpx
10from mcp import ClientSession
11from mcp.client.auth import OAuthClientProvider, TokenStorage
12from mcp.client.streamable_http import streamablehttp_client
13from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken
14from open_webui.env import (
15 AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL,
16 AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER,
17 MCP_INITIALIZE_TIMEOUT,
18)
21def _build_httpx_client(headers=None, timeout=None, auth=None, verify=True):
22 """Create an httpx AsyncClient for MCP transport.
24 Falls back to AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER when the caller
25 (i.e. the MCP SDK) does not supply an explicit timeout.
27 Note: verify must be passed at construction time because httpx
28 configures the SSL context during __init__. Setting client.verify = False
29 after construction does not affect the underlying transport's SSL context.
30 """
31 kwargs = {
32 'follow_redirects': True,
33 'verify': verify,
34 }
35 if timeout is not None:
36 kwargs['timeout'] = timeout
37 elif AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER is not None:
38 kwargs['timeout'] = float(AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER)
39 if headers is not None:
40 kwargs['headers'] = headers
41 if auth is not None:
42 kwargs['auth'] = auth
43 return httpx.AsyncClient(**kwargs)
46def create_httpx_client(headers=None, timeout=None, auth=None):
47 # AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL may be True, False, or an
48 # ssl.SSLContext (when a custom CA bundle path is configured).
49 # httpx's verify= accepts bool | str | ssl.SSLContext, so all three work.
50 ssl_setting = AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL
51 verify = ssl_setting if ssl_setting is not True else True
52 return _build_httpx_client(headers=headers, timeout=timeout, auth=auth, verify=verify)
55def create_insecure_httpx_client(headers=None, timeout=None, auth=None):
56 return _build_httpx_client(headers=headers, timeout=timeout, auth=auth, verify=False)
59class MCPClient:
60 def __init__(self):
61 self.session: Optional[ClientSession] = None
62 self.exit_stack = None
64 async def connect(self, url: str, headers: Optional[dict] = None):
65 async with AsyncExitStack() as exit_stack:
66 try:
67 self._streams_context = streamablehttp_client(
68 url,
69 headers=headers,
70 httpx_client_factory=create_httpx_client
71 if AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL
72 else create_insecure_httpx_client,
73 )
75 transport = await exit_stack.enter_async_context(self._streams_context)
76 read_stream, write_stream, _ = transport
78 self._session_context = ClientSession(read_stream, write_stream) # pylint: disable=W0201
80 self.session = await exit_stack.enter_async_context(self._session_context)
81 with anyio.fail_after(MCP_INITIALIZE_TIMEOUT):
82 await self.session.initialize()
83 self.exit_stack = exit_stack.pop_all()
84 except Exception as e:
85 await self.disconnect()
86 raise e
88 async def list_tool_specs(self) -> Optional[dict]:
89 if not self.session:
90 raise RuntimeError('MCP client is not connected.')
92 tools = []
93 cursor = None
94 while True:
95 result = await self.session.list_tools(cursor=cursor)
96 tools.extend(result.tools)
97 cursor = result.nextCursor
98 if cursor is None:
99 break
101 tool_specs = []
102 for tool in tools:
103 name = tool.name
104 description = tool.description
106 inputSchema = tool.inputSchema
108 # TODO: handle outputSchema if needed
109 outputSchema = getattr(tool, 'outputSchema', None)
111 tool_specs.append({'name': name, 'description': description, 'parameters': inputSchema})
113 return tool_specs
115 async def call_tool(self, function_name: str, function_args: dict) -> Optional[dict]:
116 if not self.session:
117 raise RuntimeError('MCP client is not connected.')
119 result = await self.session.call_tool(function_name, function_args)
120 if not result:
121 raise Exception('No result returned from MCP tool call.')
123 result_dict = result.model_dump(mode='json')
124 result_content = result_dict.get('content', {})
126 if result.isError:
127 raise Exception(result_content)
128 else:
129 return result_content
131 async def list_resources(self, cursor: Optional[str] = None) -> Optional[dict]:
132 if not self.session:
133 raise RuntimeError('MCP client is not connected.')
135 result = await self.session.list_resources(cursor=cursor)
136 if not result:
137 raise Exception('No result returned from MCP list_resources call.')
139 result_dict = result.model_dump()
140 resources = result_dict.get('resources', [])
142 return resources
144 async def read_resource(self, uri: str) -> Optional[dict]:
145 if not self.session:
146 raise RuntimeError('MCP client is not connected.')
148 result = await self.session.read_resource(uri)
149 if not result:
150 raise Exception('No result returned from MCP read_resource call.')
151 result_dict = result.model_dump()
153 return result_dict
155 async def disconnect(self):
156 """Clean up and close the session.
158 This method is idempotent — calling it multiple times or on a
159 client that was never connected is safe.
160 """
161 exit_stack = self.exit_stack
162 if exit_stack is None:
163 return
165 # Prevent double-close from concurrent callers
166 self.exit_stack = None
167 self.session = None
169 try:
170 # IMPORTANT: Do NOT use asyncio.shield() or asyncio.wait_for()
171 # because they create a new asyncio task, which violates the MCP SDK's
172 # requirement that its TaskGroup be exited in the exact same task.
173 # ALSO do NOT use anyio.CancelScope(shield=True) or anyio.fail_after(),
174 # because they push a new cancel scope onto the task, violating LIFO
175 # order when aclose() attempts to exit the inner TaskGroup.
176 # We simply call aclose() directly. If the task is cancelled, the
177 # sockets will eventually be cleaned up by garbage collection.
178 await exit_stack.aclose()
179 except asyncio.CancelledError as exc:
180 task = asyncio.current_task()
181 if task is not None and task.cancelling():
182 raise
183 log.debug('MCPClient.disconnect() suppressed internal cancellation: %s', exc)
184 except RuntimeError as exc:
185 log.debug('MCPClient.disconnect() suppressed RuntimeError: %s', exc)
186 except Exception as exc:
187 log.debug('MCPClient.disconnect() error: %s', exc)
189 async def __aenter__(self):
190 await self.exit_stack.__aenter__()
191 return self
193 async def __aexit__(self, exc_type, exc_value, traceback):
194 await self.exit_stack.__aexit__(exc_type, exc_value, traceback)
195 await self.disconnect()