diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 541ce0ad9f..f696d09527 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -253,7 +253,6 @@ from open_webui.utils.middleware import ( build_chat_response_context, process_chat_payload, process_chat_response, - resume_tool_calls, ) from open_webui.utils.misc import get_response_error_detail, merge_model_params from open_webui.utils.model_ids import strip_provider_model_prefix @@ -1697,9 +1696,9 @@ async def chat_completion( # Saved chats load the message after approved tool calls run, so their results are kept if metadata.get('assistant_message_id') and not is_saved_chat_id(metadata.get('chat_id')): ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, []) - form_data, metadata, events = await process_chat_payload(request, form_data, user, metadata, model) - - paused = await resume_tool_calls(request, form_data, user, model, metadata) + form_data, metadata, events, paused = await process_chat_payload( + request, form_data, user, metadata, model + ) if paused: return {'status': True, 'chat_id': metadata.get('chat_id'), 'paused': True} diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 12d085f411..15949b204a 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -2398,6 +2398,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): # which the frontend strips, causing tool calls to be merged into content. chat_id = metadata.get('chat_id') user_message_id = metadata.get('user_message_id') + pending_assistant_message = None payload_tools = form_data.get('tools', None) # snapshot before filters chat = await Chats.get_chat_by_id(chat_id) if is_saved_chat_id(chat_id) else None is_note_chat = bool(chat and (chat.meta or {}).get('internal') is True and (chat.meta or {}).get('type') == 'note') @@ -2419,7 +2420,20 @@ 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')): - db_messages.append({k: v for k, v in assistant_message.items() if k in MESSAGE_REPLAY_KEYS}) + 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 + ): + pending_assistant_message = assistant_message system_message = get_system_message(form_data.get('messages', [])) form_data['messages'] = [system_message, *db_messages] if system_message else db_messages @@ -2480,9 +2494,11 @@ async def process_chat_payload(request, form_data, user, metadata, model): except Exception: log.exception('Context compaction failed; continuing with full chat history') - metadata['context_start_message_id'] = next( - (message.get('id') for message in form_data.get('messages', []) if message.get('role') != 'system'), None - ) + # Compact the full chain, but prepare history separately from the response being resumed. + if pending_assistant_message is not None: + form_data['messages'] = [ + message for message in form_data['messages'] if message.get('id') != pending_assistant_message.get('id') + ] # Process messages with OR-aligned output items for clean LLM messages for message in form_data.get('messages', []): @@ -3217,6 +3233,17 @@ async def process_chat_payload(request, form_data, user, metadata, model): } ) + if pending_assistant_message is not None: + if await resume_tool_calls(request, form_data, user, model, metadata, pending_assistant_message): + return form_data, metadata, events, True + assistant_message = dict(pending_assistant_message) + if assistant_message.get('model') != model['id'] and isinstance(assistant_message.get('output'), list): + assistant_message['output'] = strip_reasoning_details(assistant_message['output']) + form_data['messages'].extend( + process_messages_with_output([assistant_message], reasoning_format=get_reasoning_format(model)) + ) + await convert_url_images_to_base64(form_data, user=user) + if ENABLE_FUNCTIONS: try: form_data, _ = await process_filter_functions( @@ -3232,7 +3259,7 @@ async def process_chat_payload(request, form_data, user, metadata, model): form_data = normalize_messages_for_model(form_data) - return form_data, metadata, events + return form_data, metadata, events, False async def get_event_emitter_and_caller(metadata): @@ -3340,16 +3367,11 @@ async def execute_tool_call(form_data, metadata, event_caller, tool_call): return params, result, tool, tool_type, direct_tool -async def resume_tool_calls(request, form_data, user, model, metadata) -> bool: - """Execute approved calls on a saved message; return whether it is still paused.""" +async def resume_tool_calls(request, form_data, user, model, metadata, message) -> bool: + """Execute approved calls in the supplied response without rebuilding prepared history.""" chat_id = metadata.get('chat_id') - assistant_message_id = metadata.get('assistant_message_id') - if not is_saved_chat_id(chat_id) or not assistant_message_id: - return False - - message_id = metadata.get('message_id') or assistant_message_id - message = await Chats.get_message_by_id_and_message_id(chat_id, message_id) - output = message.get('output') if message else None + message_id = metadata.get('message_id') or metadata.get('assistant_message_id') + output = message.get('output') if not isinstance(output, list): return False @@ -3375,9 +3397,54 @@ async def resume_tool_calls(request, form_data, user, model, metadata) -> bool: for item in output ) if not approved_calls and not needs_approval: - return False + 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. + # Give tools their filtered view without letting it overwrite the prepared history. + assistant_message = dict(message) + if assistant_message.get('model') != model['id']: + assistant_message['output'] = strip_reasoning_details(output) + form_data = { + **form_data, + 'messages': copy.deepcopy( + [ + *form_data['messages'], + *process_messages_with_output([assistant_message], reasoning_format=get_reasoning_format(model)), + ] + ), + } + await convert_url_images_to_base64(form_data, user=user) + if ENABLE_FUNCTIONS: + filter_functions = await get_filter_functions(request, model, metadata.get('filter_ids', [])) + if filter_functions: + form_data, _ = await process_filter_functions( + request=request, + filter_context=get_filter_context(request), + filter_functions=filter_functions, + filter_type='request', + form_data=form_data, + extra_params={ + '__event_emitter__': event_emitter, + '__event_call__': event_caller, + '__user__': user.model_dump() if isinstance(user, UserModel) else {}, + '__metadata__': metadata, + '__oauth_token__': await get_system_oauth_token(request, user), + '__request__': request, + '__model__': model, + '__chat_id__': metadata.get('chat_id'), + '__message_id__': metadata.get('message_id'), + }, + ) + normalize_messages_for_model(form_data) + for item in approved_calls: if item.get('name') == 'ask_user': item['status'] = 'pending' @@ -3472,61 +3539,7 @@ async def resume_tool_calls(request, form_data, user, model, metadata) -> bool: } ) - if paused: - return True - - db_messages = await load_messages_from_db(chat_id, metadata.get('user_message_id')) - if db_messages: - assistant_message = await Chats.get_message_by_id_and_message_id(chat_id, message_id) - if assistant_message: - db_messages.append({k: v for k, v in assistant_message.items() if k in MESSAGE_REPLAY_KEYS}) - context_start_message_id = metadata.get('context_start_message_id') - start_index = next( - (index for index, message in enumerate(db_messages) if message.get('id') == context_start_message_id), 0 - ) - db_messages = db_messages[start_index:] - for message in db_messages: - output = message.get('output') - # reasoning_details can be model/provider-bound, so only replay them - # for output produced by the same model. - if message.get('role') == 'assistant' and message.get('model') != model['id'] and isinstance(output, list): - message['output'] = strip_reasoning_details(output) - - system_message = get_system_message(form_data.get('messages', [])) - form_data['messages'] = process_messages_with_output( - [system_message, *db_messages] if system_message else db_messages, - reasoning_format=get_reasoning_format(model), - include_file_context=metadata.get('include_file_context', False), - ) - form_data['messages'] = sanitize_tool_pairs(form_data['messages']) - - if ENABLE_FUNCTIONS: - filter_functions = await get_filter_functions(request, model, metadata.get('filter_ids', [])) - if filter_functions: - filtered_form_data, _ = await process_filter_functions( - request=request, - filter_context=get_filter_context(request), - filter_functions=filter_functions, - filter_type='request', - form_data=form_data, - extra_params={ - '__event_emitter__': event_emitter, - '__event_call__': event_caller, - '__user__': user.model_dump() if isinstance(user, UserModel) else {}, - '__metadata__': metadata, - '__oauth_token__': await get_system_oauth_token(request, user), - '__request__': request, - '__model__': model, - '__chat_id__': metadata.get('chat_id'), - '__message_id__': metadata.get('message_id'), - }, - ) - if filtered_form_data is not form_data: - form_data.clear() - form_data.update(filtered_form_data) - - normalize_messages_for_model(form_data) - return False + return paused async def pause_for_tool_approval(chat_id: str, message_id: str, output: list[dict], form_data: dict, metadata: dict):