diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 18cd749130..c2e7d810d5 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -1500,7 +1500,7 @@ async def get_event_call(request_info): interaction_id = None timeout = WEBSOCKET_EVENT_CALLER_TIMEOUT - if event_data.get('type') == 'request:user_input' or ( + if event_data.get('type') in ('request:user_input', 'request:elicitation') or ( event_data.get('type') == 'confirmation' and (event_data.get('data') or {}).get('tool_call') ): interaction_id = str(uuid4()) diff --git a/backend/open_webui/utils/mcp/client.py b/backend/open_webui/utils/mcp/client.py index 9084b5d7ba..cd5e7b301c 100644 --- a/backend/open_webui/utils/mcp/client.py +++ b/backend/open_webui/utils/mcp/client.py @@ -3,15 +3,19 @@ import logging from contextlib import AsyncExitStack from datetime import timedelta from typing import Optional +from urllib.parse import urlsplit log = logging.getLogger(__name__) import anyio import httpx +from jsonschema import Draft202012Validator, FormatChecker from mcp import ClientSession from mcp.client.auth import OAuthClientProvider, TokenStorage from mcp.client.streamable_http import streamablehttp_client from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken +from mcp.types import ElicitRequestFormParams, ElicitResult +from referencing import Registry from open_webui.env import ( AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL, AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER, @@ -82,9 +86,45 @@ class OAuthTokenAuth(httpx.Auth): class MCPClient: - def __init__(self): + def __init__(self, event_caller=None, server_name=''): self.session: Optional[ClientSession] = None self.exit_stack = None + self.event_caller = event_caller + self.server_name = server_name + self.instructions: Optional[str] = None + + async def elicit(self, context, params): + try: + if isinstance(params, ElicitRequestFormParams): + Draft202012Validator.check_schema(params.requestedSchema) + else: + url = urlsplit(params.url) + if url.scheme not in ('http', 'https') or not url.hostname or url.username or url.password: + return ElicitResult(action='cancel') + + response = await self.event_caller( + { + 'type': 'request:elicitation', + 'data': { + **params.model_dump(mode='json', include={'mode', 'message', 'requestedSchema', 'url'}), + 'server_name': self.server_name, + }, + } + ) + if not isinstance(response, dict) or response.get('action') not in ('accept', 'decline', 'cancel'): + return ElicitResult(action='cancel') + + content = None + if response['action'] == 'accept' and isinstance(params, ElicitRequestFormParams): + content = response.get('content') + # Never fetch server-supplied schema references or coerce user answers. + Draft202012Validator( + params.requestedSchema, format_checker=FormatChecker(), registry=Registry() + ).validate(content) + return ElicitResult(action=response['action'], content=content) + except Exception as e: + log.warning('MCP elicitation failed: %s', type(e).__name__) + return ElicitResult(action='cancel') async def connect(self, url: str, headers: Optional[dict] = None, auth: Optional[httpx.Auth] = None): async with AsyncExitStack() as exit_stack: @@ -101,11 +141,14 @@ class MCPClient: transport = await exit_stack.enter_async_context(self._streams_context) read_stream, write_stream, _ = transport - self._session_context = ClientSession(read_stream, write_stream) # pylint: disable=W0201 + self._session_context = ClientSession( + read_stream, write_stream, elicitation_callback=self.elicit if self.event_caller else None + ) self.session = await exit_stack.enter_async_context(self._session_context) with anyio.fail_after(MCP_INITIALIZE_TIMEOUT): - await self.session.initialize() + result = await self.session.initialize() + self.instructions = result.instructions self.exit_stack = exit_stack.pop_all() except Exception as e: await self.disconnect() diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 6aa8598302..6eca511252 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -3042,7 +3042,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): request, server_id, user, - metadata, + extra_params, ) if result is None: continue @@ -3050,6 +3050,13 @@ async def process_chat_payload(request, form_data, user, metadata, model): client, tool_specs = result mcp_clients[server_id] = client + if client.instructions: + form_data['messages'] = add_or_update_system_message( + f'MCP server {JSONCodec.dumps(server_id)} instructions:\n{client.instructions}', + form_data['messages'], + append=True, + ) + for tool_spec in tool_specs: async def make_tool_function(client, function_name): diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index 90e91f9f23..6423707ea3 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -198,7 +198,7 @@ async def connect_mcp_server( request, server_id: str, user, - metadata: dict, + extra_params: dict, ) -> tuple[MCPClient, list[dict]] | None: """Resolve an MCP server connection, authenticate, and return (client, tool_specs). @@ -222,6 +222,8 @@ async def connect_mcp_server( log.warning(f'Access denied to MCP server {server_id} for user {user.id}') return None + metadata = extra_params.get('__metadata__', {}) + async def get_headers(): headers, _ = await build_tool_server_headers( mcp_server_connection, request, user, server_id=server_id, metadata=metadata @@ -234,7 +236,10 @@ async def connect_mcp_server( if mcp_server_connection.get('auth_type') in ('oauth_2.1', 'oauth_2.1_static') and not headers.get('Authorization'): raise HTTPException(status_code=401, detail='Auth required') - client = MCPClient() + client = MCPClient( + event_caller=extra_params.get('__event_call__') if metadata.get('session_id') else None, + server_name=server_id, + ) try: await client.connect( url=mcp_server_connection.get('url', ''), diff --git a/src/lib/components/chat/Chat.svelte b/src/lib/components/chat/Chat.svelte index 7ebb65f030..b7b9252c03 100644 --- a/src/lib/components/chat/Chat.svelte +++ b/src/lib/components/chat/Chat.svelte @@ -120,6 +120,7 @@ import Navbar from '$lib/components/chat/Navbar.svelte'; import ChatControls from './ChatControls.svelte'; import EventConfirmDialog from '../common/ConfirmDialog.svelte'; + import ElicitationDialog from './ElicitationDialog.svelte'; import DeleteConfirmDialog from '../common/ConfirmDialog.svelte'; import WebSearchConfirmDialog from '../common/ConfirmDialog.svelte'; import Placeholder from './Placeholder.svelte'; @@ -1416,6 +1417,7 @@ } if ( interactionType === 'request:user_input' || + interactionType === 'request:elicitation' || (interactionType === 'confirmation' && interactionData?.tool_call) ) { if (!cb) return; @@ -4813,6 +4815,16 @@ }} /> +{#if browserInteraction?.type === 'request:elicitation'} + {#key browserInteraction.id} + {@const interaction = browserInteraction} + resolveBrowserInteraction(interaction, value)} + /> + {/key} +{/if} + {#if browserInteraction?.data?.tool_call} {#key browserInteraction.id} {@const interaction = browserInteraction} diff --git a/src/lib/components/chat/ElicitationDialog.svelte b/src/lib/components/chat/ElicitationDialog.svelte new file mode 100644 index 0000000000..da1b56e63d --- /dev/null +++ b/src/lib/components/chat/ElicitationDialog.svelte @@ -0,0 +1,171 @@ + + + +
+
{data.server_name}
+ +
+
+

{data.message}

+ + {#if data.mode === 'url'} +

{$i18n.t('Open this link to continue:')}

+ onResponse({ action: 'accept' })}>{data.url} + {:else if supported} + {#each fields as [name, field], index} + {@const id = `elicitation-${index}`} + {@const required = data.requestedSchema.required?.includes(name) ?? false} +
+ + {#if field.description} +

+ {field.description} +

+ {/if} + {#if field.type === 'array'} + + {:else if field.type === 'boolean' || options(field).length} + + {:else if field.type === 'number' || field.type === 'integer'} + + {:else} + + {/if} +
+ {/each} + {:else} +

{$i18n.t('Unsupported input format')}

+ {/if} + {#if error}{/if} +
+ + + {#if data.mode !== 'url' && supported} + + {/if} +
+
+