diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index f9b58038b9..f5181fb1b6 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -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, diff --git a/backend/open_webui/models/chats.py b/backend/open_webui/models/chats.py index e903afde9c..b0881bc6c3 100644 --- a/backend/open_webui/models/chats.py +++ b/backend/open_webui/models/chats.py @@ -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) diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index c585e2acb2..80b364f574 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -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): diff --git a/backend/open_webui/utils/tool_approval.py b/backend/open_webui/utils/tool_approval.py index a87aa5f535..52d8d822c9 100644 --- a/backend/open_webui/utils/tool_approval.py +++ b/backend/open_webui/utils/tool_approval.py @@ -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: