mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-11 03:38:02 +00:00
refac
This commit is contained in:
parent
b40891e0a1
commit
c7b5edc726
6 changed files with 245 additions and 7 deletions
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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', ''),
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
<ElicitationDialog
|
||||
data={interaction.data}
|
||||
onResponse={(value) => resolveBrowserInteraction(interaction, value)}
|
||||
/>
|
||||
{/key}
|
||||
{/if}
|
||||
|
||||
{#if browserInteraction?.data?.tool_call}
|
||||
{#key browserInteraction.id}
|
||||
{@const interaction = browserInteraction}
|
||||
|
|
|
|||
171
src/lib/components/chat/ElicitationDialog.svelte
Normal file
171
src/lib/components/chat/ElicitationDialog.svelte
Normal file
|
|
@ -0,0 +1,171 @@
|
|||
<script lang="ts">
|
||||
import { getContext } from 'svelte';
|
||||
import Modal from '../common/Modal.svelte';
|
||||
import XMark from '../icons/XMark.svelte';
|
||||
|
||||
const i18n = getContext<typeof import('$lib/i18n').default>('i18n');
|
||||
export let data;
|
||||
export let onResponse: (value: any) => void;
|
||||
|
||||
let show = true;
|
||||
const fields: [string, any][] = Object.entries(data.requestedSchema?.properties ?? {});
|
||||
let values = Object.fromEntries(fields.map(([name, field]) => [name, field.default]));
|
||||
let error = '';
|
||||
$: if (!show) onResponse({ action: 'cancel' });
|
||||
|
||||
const options = (field: any): { const: string; title?: string }[] =>
|
||||
field.enum?.map((value: string, index: number) => ({
|
||||
const: value,
|
||||
title: field.enumNames?.[index] ?? value
|
||||
})) ??
|
||||
field.oneOf ??
|
||||
field.anyOf ??
|
||||
[];
|
||||
const supported = fields.every(([, field]) =>
|
||||
field.type === 'array'
|
||||
? options(field.items ?? {}).length > 0
|
||||
: ['string', 'number', 'integer', 'boolean'].includes(field.type)
|
||||
);
|
||||
|
||||
const submit = () => {
|
||||
for (const [name, field] of fields) {
|
||||
if (field.type === 'array' && values[name] !== undefined) {
|
||||
const count = values[name]?.length ?? 0;
|
||||
if (count < (field.minItems ?? 0) || count > (field.maxItems ?? Infinity)) {
|
||||
error = `${field.title ?? name}: ${$i18n.t('Invalid value')}`;
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
onResponse({
|
||||
action: 'accept',
|
||||
content: Object.fromEntries(Object.entries(values).filter(([, value]) => value !== undefined))
|
||||
});
|
||||
};
|
||||
</script>
|
||||
|
||||
<Modal bind:show size="sm">
|
||||
<div class="flex justify-between dark:text-gray-300 px-4 pt-3 pb-1">
|
||||
<div class="text-sm font-medium self-center">{data.server_name}</div>
|
||||
<button
|
||||
type="button"
|
||||
aria-label={$i18n.t('Close')}
|
||||
class="self-center rounded-lg p-1 text-gray-500 transition hover:bg-gray-50 hover:text-gray-700 dark:text-gray-400 dark:hover:bg-gray-800 dark:hover:text-gray-200"
|
||||
on:click={() => (show = false)}
|
||||
>
|
||||
<XMark className="size-4" />
|
||||
</button>
|
||||
</div>
|
||||
<form class="flex flex-col gap-3 px-4 pb-4 dark:text-gray-200" on:submit|preventDefault={submit}>
|
||||
<p class="text-sm whitespace-pre-wrap break-words">{data.message}</p>
|
||||
|
||||
{#if data.mode === 'url'}
|
||||
<p class="text-sm text-gray-500">{$i18n.t('Open this link to continue:')}</p>
|
||||
<a
|
||||
href={data.url}
|
||||
target="_blank"
|
||||
rel="noopener noreferrer"
|
||||
class="text-sm underline break-all"
|
||||
on:click={() => onResponse({ action: 'accept' })}>{data.url}</a
|
||||
>
|
||||
{:else if supported}
|
||||
{#each fields as [name, field], index}
|
||||
{@const id = `elicitation-${index}`}
|
||||
{@const required = data.requestedSchema.required?.includes(name) ?? false}
|
||||
<div class="flex flex-col gap-1.5 text-sm">
|
||||
<label for={id} class="text-xs font-normal">
|
||||
{field.title ?? name}
|
||||
{#if required}<span class="ml-1 text-gray-500">* {$i18n.t('required')}</span>{/if}
|
||||
</label>
|
||||
{#if field.description}
|
||||
<p id={`${id}-description`} class="text-xs text-gray-500 whitespace-pre-wrap">
|
||||
{field.description}
|
||||
</p>
|
||||
{/if}
|
||||
{#if field.type === 'array'}
|
||||
<select
|
||||
{id}
|
||||
multiple
|
||||
{required}
|
||||
aria-describedby={field.description ? `${id}-description` : undefined}
|
||||
class="w-full rounded-lg py-2 px-4 text-sm dark:text-gray-300 dark:bg-gray-850 outline-hidden border border-gray-100/30 dark:border-gray-850/30"
|
||||
bind:value={values[name]}
|
||||
>
|
||||
{#each options(field.items ?? {}) as option}
|
||||
<option value={option.const}>{option.title ?? option.const}</option>
|
||||
{/each}
|
||||
</select>
|
||||
{:else if field.type === 'boolean' || options(field).length}
|
||||
<select
|
||||
{id}
|
||||
{required}
|
||||
aria-describedby={field.description ? `${id}-description` : undefined}
|
||||
class="w-full rounded-lg py-2 px-4 text-sm dark:text-gray-300 dark:bg-gray-850 outline-hidden border border-gray-100/30 dark:border-gray-850/30"
|
||||
bind:value={values[name]}
|
||||
>
|
||||
<option value={undefined}>{$i18n.t('Select an option')}</option>
|
||||
{#if field.type === 'boolean'}
|
||||
<option value={true}>{$i18n.t('Yes')}</option>
|
||||
<option value={false}>{$i18n.t('No')}</option>
|
||||
{:else}
|
||||
{#each options(field) as option}
|
||||
<option value={option.const}>{option.title ?? option.const}</option>
|
||||
{/each}
|
||||
{/if}
|
||||
</select>
|
||||
{:else if field.type === 'number' || field.type === 'integer'}
|
||||
<input
|
||||
{id}
|
||||
type="number"
|
||||
{required}
|
||||
aria-describedby={field.description ? `${id}-description` : undefined}
|
||||
min={field.minimum}
|
||||
max={field.maximum}
|
||||
step={field.type === 'integer' ? 1 : 'any'}
|
||||
class="w-full rounded-lg py-2 px-4 text-sm dark:text-gray-300 dark:bg-gray-850 outline-hidden border border-gray-100/30 dark:border-gray-850/30"
|
||||
bind:value={values[name]}
|
||||
/>
|
||||
{:else}
|
||||
<input
|
||||
{id}
|
||||
type={['email', 'date'].includes(field.format)
|
||||
? field.format
|
||||
: field.format === 'uri'
|
||||
? 'url'
|
||||
: 'text'}
|
||||
placeholder={field.format === 'date-time' ? 'YYYY-MM-DDTHH:mm:ssZ' : undefined}
|
||||
{required}
|
||||
aria-describedby={field.description ? `${id}-description` : undefined}
|
||||
minlength={field.minLength}
|
||||
maxlength={field.maxLength}
|
||||
class="w-full rounded-lg py-2 px-4 text-sm dark:text-gray-300 dark:bg-gray-850 outline-hidden border border-gray-100/30 dark:border-gray-850/30"
|
||||
bind:value={values[name]}
|
||||
/>
|
||||
{/if}
|
||||
</div>
|
||||
{/each}
|
||||
{:else}
|
||||
<p role="alert">{$i18n.t('Unsupported input format')}</p>
|
||||
{/if}
|
||||
{#if error}<p role="alert" class="text-sm text-red-500">{error}</p>{/if}
|
||||
<div class="flex justify-end gap-1.5 pt-3 text-sm font-normal">
|
||||
<button
|
||||
type="button"
|
||||
class="flex h-7 shrink-0 items-center justify-center gap-1.5 rounded-lg px-2.5 text-xs font-normal transition disabled:opacity-60 bg-white hover:bg-gray-100 text-black dark:bg-black dark:text-white dark:hover:bg-gray-900"
|
||||
on:click={() => onResponse({ action: 'cancel' })}>{$i18n.t('Cancel')}</button
|
||||
>
|
||||
<button
|
||||
type="button"
|
||||
class="flex h-7 items-center justify-center gap-1.5 rounded-lg px-2.5 text-xs font-normal transition disabled:opacity-60 bg-gray-100 hover:bg-gray-100/70 text-gray-800 dark:bg-gray-850 dark:hover:bg-gray-850/60 dark:text-white"
|
||||
on:click={() => onResponse({ action: 'decline' })}>{$i18n.t('Decline')}</button
|
||||
>
|
||||
{#if data.mode !== 'url' && supported}
|
||||
<button
|
||||
type="submit"
|
||||
class="flex h-7 shrink-0 items-center justify-center gap-1.5 rounded-lg bg-gray-900 px-2.5 text-xs font-normal text-white transition hover:bg-black disabled:opacity-60 dark:bg-gray-100 dark:text-gray-900 dark:hover:bg-white"
|
||||
>{$i18n.t('Submit')}</button
|
||||
>
|
||||
{/if}
|
||||
</div>
|
||||
</form>
|
||||
</Modal>
|
||||
Loading…
Add table
Reference in a new issue