This commit is contained in:
Timothy Jaeryang Baek 2026-10-10 22:39:57 +04:00
parent 46a2a830ab
commit 93723bd315
4 changed files with 282 additions and 197 deletions

View file

@ -1308,7 +1308,11 @@ async def chat_completion(
or 'full'
)
approval_resume = getattr(request.state, 'tool_approval_resume', None)
metadata = {
'tool_approval_resume': approval_resume[2]
if approval_resume and approval_resume[:2] == (chat_id, form_data.get('assistant_message_id'))
else None,
'user_id': user.id,
'user_agent': request.headers.get('user-agent', '') or '',
'internal': getattr(request.state, 'internal', False) is True,
@ -1916,8 +1920,14 @@ async def resolve_chat_message_tool_call(
db: AsyncSession = Depends(get_async_session),
):
resolution = await resolve_tool_call_output(id, message_id, form_data, user, db=db)
payload = await build_tool_approval_resume_payload(id, message_id, chat=resolution['chat'])
result = await chat_completion(request, payload, user)
if resolution['resume'] is None:
return {'status': True, 'chat_id': id, 'message_id': message_id, 'task_ids': []}
request.state.tool_approval_resume = (id, message_id, resolution['resume'])
try:
payload = await build_tool_approval_resume_payload(id, message_id, chat=resolution['chat'])
result = await chat_completion(request, payload, user)
finally:
del request.state.tool_approval_resume
return {
'status': True,
'chat_id': id,

View file

@ -3,6 +3,7 @@
from __future__ import annotations
from contextlib import asynccontextmanager
from copy import deepcopy
import logging
import re
import time
@ -405,6 +406,38 @@ class ChatTable:
yield session, chat
await session.commit()
@asynccontextmanager
async def edit_message_output(self, chat_id: str, message_id: str):
"""Read and modify tool state under the same lock, updating both message stores."""
async with self._chat_transaction(chat_id) as (session, chat):
if chat is None:
yield None
return
history = (chat.chat or {}).get('history') or {}
message = dict(history.get('messages', {}).get(message_id) or {})
row = await session.get(ChatMessage, f'{chat_id}-{message_id}')
if row is not None:
message.update(
{
ChatMessages.DB_TO_JSON_KEY_MAP.get(column.key, column.key): getattr(row, column.key)
for column in ChatMessage.__table__.columns
if column.key not in ChatMessages.EXCLUDED_COLUMNS
}
)
message['id'] = message_id
if not message:
yield None
return
message = deepcopy(message)
yield message
message = self._clean_null_bytes(message)
history.setdefault('messages', {})[message_id] = message
chat.chat = {**(chat.chat or {}), 'history': history}
flag_modified(chat, 'chat')
await ChatMessages.upsert_message(
message_id, chat_id, message.get('user_id') or chat.user_id, message, db=session
)
def _clean_null_bytes(self, obj):
"""Recursively remove null bytes from strings in dict/list structures."""
return sanitize_data_for_db(obj)

View file

@ -143,6 +143,7 @@ from open_webui.utils.task import (
rag_template,
tools_function_calling_generation_template,
)
from open_webui.utils.tool_approval import complete_tool_call, pending_tool_calls
from open_webui.utils.tools import (
connect_mcp_server,
get_attached_knowledge,
@ -2465,19 +2466,22 @@ async def process_chat_payload(request, form_data, user, metadata, model):
if assistant_message_id:
assistant_message = await Chats.get_message_by_id_and_message_id(chat_id, assistant_message_id)
if assistant_message and (assistant_message.get('content') or assistant_message.get('output')):
output = assistant_message.get('output') or []
if (
metadata.get('tool_approval_resume') != 'model'
and not assistant_message.get('done', True)
and output
and output[-1].get('type') == 'message'
and output[-1].get('status') == 'in_progress'
and any(item.get('type') == 'function_call_output' for item in output)
and not pending_tool_calls(output)
):
return form_data, metadata, [], True
assistant_message = {k: v for k, v in assistant_message.items() if k in MESSAGE_REPLAY_KEYS}
db_messages.append(assistant_message)
output = assistant_message.get('output')
output = output if isinstance(output, list) else []
result_call_ids = {
item.get('call_id') for item in output if item.get('type') == 'function_call_output'
}
if any(
item.get('type') == 'function_call'
and item.get('status') in {'pending', 'queued', 'requires_approval'}
and item.get('call_id') not in result_call_ids
for item in output
):
if metadata.get('tool_approval_resume') == 'tools' or pending_tool_calls(output):
pending_assistant_message = assistant_message
system_message = get_system_message(form_data.get('messages', []))
@ -3507,36 +3511,6 @@ async def resume_tool_calls(request, form_data, user, model, metadata, message)
if not isinstance(output, list):
return False
result_call_ids = {
item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id')
}
approved_calls = [
item
for item in output
if item.get('type') == 'function_call'
and item.get('call_id')
and item.get('status') == 'queued'
and item.get('approved') is True
and item.get('call_id') not in result_call_ids
]
needs_approval = metadata.get('params', {}).get('tool_approval_mode', 'full') == 'ask' and any(
item.get('type') == 'function_call'
and item.get('name') != 'ask_user'
and (item.get('call_id') or item.get('id'))
and item.get('status') == 'queued'
and item.get('approved') is not True
and (item.get('call_id') or item.get('id')) not in result_call_ids
for item in output
)
if not approved_calls and not needs_approval:
return any(
item.get('type') == 'function_call'
and item.get('call_id')
and item.get('status') in {'pending', 'queued', 'requires_approval'}
and item.get('call_id') not in result_call_ids
for item in output
)
event_emitter, event_caller = await get_event_emitter_and_caller(metadata)
# Request filters must be able to reject the complete request before tool side effects.
@ -3577,11 +3551,34 @@ async def resume_tool_calls(request, form_data, user, model, metadata, message)
)
normalize_messages_for_model(form_data)
for item in approved_calls:
if item.get('name') == 'ask_user':
item['status'] = 'pending'
item.pop('approved', None)
continue
while True:
async with Chats.edit_message_output(chat_id, message_id) as saved:
if not saved:
raise HTTPException(status_code=404, detail='Message not found.')
output = saved.get('output') or []
pending = pending_tool_calls(output)
if any(entry.get('status') == 'in_progress' for entry in pending):
return True
item = next(
(
entry
for entry in pending
if entry.get('name') != 'ask_user'
and (
entry.get('approved') is True or metadata.get('params', {}).get('tool_approval_mode') == 'full'
)
),
None,
)
if item is None:
if pending:
pending[0]['status'] = 'pending'
# A stale continuation must not start another model response.
return True
item['call_id'] = item.get('call_id') or item['id']
item['status'] = 'in_progress'
saved['done'] = False
item = copy.deepcopy(item)
tool_call = {
'id': item.get('call_id', ''),
@ -3591,37 +3588,39 @@ async def resume_tool_calls(request, form_data, user, model, metadata, message)
'arguments': item.get('arguments', '{}'),
},
}
params, result, tool, tool_type, direct_tool = await execute_tool_call(
form_data, metadata, event_caller, tool_call
)
files, embeds = [], []
if result is None and tool is None:
result = 'Error: Tool call arguments could not be parsed. The model generated malformed or incomplete JSON.'
elif tool:
name = item.get('name', '')
terminal_file_result = build_terminal_file_tool_result(name, params, result, tool, metadata)
if terminal_file_result:
result = terminal_file_result
result, files, embeds = await process_tool_result(
request, name, result, tool_type, direct_tool, metadata, user
try:
params, result, tool, tool_type, direct_tool = await execute_tool_call(
form_data, metadata, event_caller, tool_call
)
await terminal_event_handler(name, params, result, event_emitter)
content = tool_result_content(result)
item['arguments'] = tool_call.get('function', {}).get('arguments', '{}')
output_parts = [{'type': 'input_text', 'text': content}]
item['status'] = 'failed' if _is_tool_result_error(content) else 'completed'
display_files = []
for file_item in files:
if file_item.get('type') == 'image' and file_item.get('url', '').startswith('data:'):
image_url = await store_tool_result_image(request, file_item['url'], metadata, user)
output_parts.append({'type': 'input_image', 'image_url': image_url})
else:
display_files.append(file_item)
if file_item.get('type') == 'image' and file_item.get('url'):
output_parts.append({'type': 'input_image', 'image_url': file_item['url']})
files, embeds = [], []
if result is None and tool is None:
result = (
'Error: Tool call arguments could not be parsed. The model generated malformed or incomplete JSON.'
)
elif tool:
name = item.get('name', '')
terminal_file_result = build_terminal_file_tool_result(name, params, result, tool, metadata)
if terminal_file_result:
result = terminal_file_result
result, files, embeds = await process_tool_result(
request, name, result, tool_type, direct_tool, metadata, user
)
await terminal_event_handler(name, params, result, event_emitter)
content = tool_result_content(result)
item['arguments'] = tool_call.get('function', {}).get('arguments', '{}')
output_parts = [{'type': 'input_text', 'text': content}]
item['status'] = 'failed' if _is_tool_result_error(content) else 'completed'
display_files = []
for file_item in files:
if file_item.get('type') == 'image' and file_item.get('url', '').startswith('data:'):
image_url = await store_tool_result_image(request, file_item['url'], metadata, user)
output_parts.append({'type': 'input_image', 'image_url': image_url})
else:
display_files.append(file_item)
if file_item.get('type') == 'image' and file_item.get('url'):
output_parts.append({'type': 'input_image', 'image_url': file_item['url']})
output.append(
{
result_item = {
'type': 'function_call_output',
'id': output_id('fco'),
'call_id': tool_call['id'],
@ -3630,48 +3629,39 @@ async def resume_tool_calls(request, form_data, user, model, metadata, message)
**({'files': display_files} if display_files else {}),
**({'embeds': embeds} if embeds else {}),
}
)
result_call_ids.add(tool_call['id'])
if needs_approval:
await pause_for_tool_approval(chat_id, message_id, output, form_data, metadata)
paused = any(
item.get('type') == 'function_call'
and item.get('call_id')
and item.get('status') in {'pending', 'queued', 'requires_approval'}
and item.get('call_id') not in result_call_ids
for item in output
)
if not paused:
output.append(
{
'type': 'message',
'id': output_id('msg'),
'status': 'in_progress',
'role': 'assistant',
'content': [{'type': 'output_text', 'text': ''}],
}
)
if not needs_approval:
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
message_id,
{'done': False, 'output': output},
touch=False,
)
if event_emitter:
await event_emitter(
{
'type': 'chat:completion',
'data': {
'done': False,
'output': output,
},
}
)
return paused
output = await complete_tool_call(chat_id, message_id, item, result_item)
except BaseException:
item['status'] = 'incomplete'
await asyncio.shield(
complete_tool_call(
chat_id,
message_id,
item,
{
'type': 'function_call_output',
'id': output_id('fco'),
'call_id': tool_call['id'],
'status': 'incomplete',
'output': [
{
'type': 'input_text',
'text': 'Error: tool execution was interrupted; its outcome is unknown.',
}
],
},
)
)
raise
if event_emitter:
await event_emitter({'type': 'chat:completion', 'data': {'done': False, 'output': output}})
pending = pending_tool_calls(output)
if not pending:
message['output'] = output
return False
if metadata.get('params', {}).get('tool_approval_mode') != 'full' and not any(
entry.get('approved') is True for entry in pending
):
return True
async def pause_for_tool_approval(chat_id: str, message_id: str, output: list[dict], form_data: dict, metadata: dict):

View file

@ -1,4 +1,5 @@
from typing import Any, Literal
from uuid import uuid4
from fastapi import HTTPException, status
from pydantic import BaseModel
@ -19,6 +20,50 @@ class ResolveToolCallForm(BaseModel):
timed_out: bool = False
def pending_tool_calls(output):
result_ids = {item.get('call_id') for item in output if item.get('type') == 'function_call_output'}
return [
item
for item in output
if item.get('type') == 'function_call'
and (item.get('call_id') or item.get('id')) not in result_ids
and item.get('status') in {'pending', 'queued', 'requires_approval', 'in_progress'}
]
def tool_continuation():
return {
'type': 'message',
'id': f'msg_{uuid4().hex}',
'status': 'in_progress',
'role': 'assistant',
'content': [{'type': 'output_text', 'text': ''}],
}
async def complete_tool_call(chat_id, message_id, call, result):
async with Chats.edit_message_output(chat_id, message_id) as message:
if not message:
raise HTTPException(status_code=404, detail=ERROR_MESSAGES.NOT_FOUND)
output = message.get('output') or []
saved = next((item for item in pending_tool_calls(output) if item.get('call_id') == call['call_id']), None)
if saved is None or saved.get('status') != 'in_progress':
raise HTTPException(status_code=409, detail='Tool call is no longer running.')
saved.update(call)
output.append(result)
pending = pending_tool_calls(output)
if result.get('status') == 'incomplete':
message['done'] = True
elif not pending:
# Only the transaction completing the last call may start the model again.
output.append(tool_continuation())
else:
next_approval = next((item for item in pending if not item.get('approved')), None)
if next_approval is not None:
next_approval['status'] = 'pending'
return output
async def resolve_tool_call_output(
chat_id: str,
message_id: str,
@ -33,81 +78,98 @@ async def resolve_tool_call_output(
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
)
message = await Chats.get_message_by_id_and_message_id(chat_id, message_id)
if not message:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
async with Chats.edit_message_output(chat_id, message_id) as message:
if not message:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
if (message.get('user_id') or chat.user_id) != user.id:
raise HTTPException(status_code=403, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
if (message.get('user_id') or chat.user_id) != user.id:
raise HTTPException(status_code=403, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
output = message.get('output') or []
if not isinstance(output, list):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Message has no resolvable output.')
output = message.get('output') or []
if not isinstance(output, list):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Message has no resolvable output.')
function_call = next(
(
item
for item in output
if item.get('type') == 'function_call' and (item.get('call_id') or item.get('id')) == form_data.call_id
),
None,
)
if not function_call:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail='Tool call not found.')
function_call.setdefault('call_id', form_data.call_id)
tool_name = function_call.get('name')
if any(
item.get('type') == 'function_call_output' and item.get('call_id') == form_data.call_id for item in output
) or function_call.get('status') not in {'pending', 'queued', 'requires_approval'}:
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call has already been resolved.')
if form_data.action == 'approve':
if tool_name == 'ask_user':
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='ask_user requires an answer or deny.')
function_call['status'] = 'queued'
function_call['approved'] = True
elif form_data.action == 'reject':
function_call['status'] = 'rejected'
output.append(
{
'type': 'function_call_output',
'id': f'fco_{form_data.call_id}',
'call_id': form_data.call_id,
'output': [{'type': 'input_text', 'text': 'Error: tool call rejected by user.'}],
'status': 'rejected',
}
function_call = next(
(
item
for item in output
if item.get('type') == 'function_call' and (item.get('call_id') or item.get('id')) == form_data.call_id
),
None,
)
else:
if tool_name != 'ask_user':
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Tool call does not accept answers.')
if form_data.answers is None and not form_data.timed_out:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Answers are required for ask_user.')
function_call['status'] = 'completed'
answer_payload = (
{'status': 'cancelled', 'answers': {}, 'timed_out': True}
if form_data.timed_out
else {'status': 'answered', 'answers': form_data.answers or {}}
)
output.append(
{
'type': 'function_call_output',
'id': f'fco_{form_data.call_id}',
'call_id': form_data.call_id,
'output': [{'type': 'input_text', 'text': JSONCodec.dumps(answer_payload)}],
'status': 'completed',
}
if not function_call:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail='Tool call not found.')
function_call.setdefault('call_id', form_data.call_id)
tool_name = function_call.get('name')
if (
any(
item.get('type') == 'function_call_output' and item.get('call_id') == form_data.call_id
for item in output
)
or function_call.get('status') not in {'pending', 'queued', 'requires_approval'}
or function_call.get('approved') is True
):
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call has already been resolved.')
active = any(
item.get('approved') is True or item.get('status') == 'in_progress' for item in pending_tool_calls(output)
)
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
message_id,
{
'done': False,
'output': output,
},
touch=False,
)
if form_data.action == 'approve':
if tool_name == 'ask_user':
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail='ask_user requires an answer or deny.'
)
function_call['status'] = 'queued'
function_call['approved'] = True
elif form_data.action == 'reject':
function_call['status'] = 'rejected'
output.append(
{
'type': 'function_call_output',
'id': f'fco_{form_data.call_id}',
'call_id': form_data.call_id,
'output': [{'type': 'input_text', 'text': 'Error: tool call rejected by user.'}],
'status': 'rejected',
}
)
else:
if tool_name != 'ask_user':
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail='Tool call does not accept answers.'
)
if form_data.answers is None and not form_data.timed_out:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail='Answers are required for ask_user.'
)
function_call['status'] = 'completed'
answer_payload = (
{'status': 'cancelled', 'answers': {}, 'timed_out': True}
if form_data.timed_out
else {'status': 'answered', 'answers': form_data.answers or {}}
)
output.append(
{
'type': 'function_call_output',
'id': f'fco_{form_data.call_id}',
'call_id': form_data.call_id,
'output': [{'type': 'input_text', 'text': JSONCodec.dumps(answer_payload)}],
'status': 'completed',
}
)
message['done'] = False
pending = pending_tool_calls(output)
resume = None
if not active:
if form_data.action == 'approve':
resume = 'tools'
elif not pending:
output.append(tool_continuation())
resume = 'model'
else:
pending[0]['status'] = 'pending'
event_emitter = await get_event_emitter(
{
@ -120,17 +182,7 @@ async def resolve_tool_call_output(
if event_emitter:
await event_emitter({'type': 'chat:completion', 'data': {'output': output}})
result_call_ids = {
item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id')
}
paused = any(
item.get('type') == 'function_call'
and item.get('call_id')
and item.get('status') in {'pending', 'queued', 'requires_approval'}
and item.get('call_id') not in result_call_ids
for item in output
)
return {'chat': chat, 'message': message, 'output': output, 'paused': paused}
return {'chat': chat, 'output': output, 'resume': resume}
async def build_tool_approval_resume_payload(chat_id: str, message_id: str, chat=None) -> dict: