From a1579a01ff43cacb357269707d36267ad35e01d6 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Fri, 14 Aug 2026 00:22:17 -0600 Subject: [PATCH] refac --- backend/open_webui/routers/chats.py | 39 +- backend/open_webui/socket/main.py | 16 +- backend/open_webui/tasks.py | 61 +++ backend/open_webui/utils/middleware.py | 347 ++++++++++++------ src/lib/components/chat/Chat.svelte | 25 +- .../chat/Messages/structuredOutput.ts | 179 +++++++++ 6 files changed, 557 insertions(+), 110 deletions(-) diff --git a/backend/open_webui/routers/chats.py b/backend/open_webui/routers/chats.py index f4764599b5..549c26217b 100644 --- a/backend/open_webui/routers/chats.py +++ b/backend/open_webui/routers/chats.py @@ -34,7 +34,7 @@ from open_webui.models.folders import Folders from open_webui.models.shared_chats import SharedChatResponse, SharedChats from open_webui.models.tags import TagModel, Tags from open_webui.socket.main import get_event_emitter -from open_webui.tasks import has_active_tasks, stop_item_tasks +from open_webui.tasks import get_response_streams_by_chat_id, has_active_tasks, stop_item_tasks from open_webui.utils.access_control import filter_allowed_access_grants, has_permission from open_webui.utils.access_control.folders import has_folder_access, has_folder_write_access from open_webui.utils.auth import bearer_security, get_admin_user, get_current_user, get_verified_user @@ -60,6 +60,32 @@ CHAT_CONFIG_KEYS = { } +def overlay_response_streams(chat_data: dict, response_streams: list[dict]) -> dict: + if not response_streams: + return chat_data + + messages = chat_data.get('chat', {}).get('history', {}).get('messages') + if isinstance(messages, dict): + for stream in response_streams: + message_id = stream.get('message_id') + message = messages.get(message_id) + if isinstance(message, dict): + message['content'] = stream.get('content', '') + message['output'] = stream.get('output') or [] + message['done'] = False + + legacy_messages = chat_data.get('chat', {}).get('messages') + if isinstance(legacy_messages, list): + streams_by_message_id = {stream.get('message_id'): stream for stream in response_streams} + for message in legacy_messages: + if isinstance(message, dict) and (stream := streams_by_message_id.get(message.get('id'))): + message['content'] = stream.get('content', '') + message['output'] = stream.get('output') or [] + message['done'] = False + + return chat_data + + async def get_optional_verified_user( request: Request, response: Response, @@ -1297,7 +1323,12 @@ async def compact_chat_by_id( @router.get('/{id}', response_model=ChatResponse | None) -async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSession = Depends(get_async_session)): +async def get_chat_by_id( + id: str, + request: Request, + user=Depends(get_verified_user), + db: AsyncSession = Depends(get_async_session), +): chat = await Chats.get_chat_by_id_and_user_id(id, user.id, db=db) if not chat and user.role == 'admin': @@ -1328,6 +1359,10 @@ async def get_chat_by_id(id: str, user=Depends(get_verified_user), db: AsyncSess if chat: data = ChatResponse.model_validate(chat, from_attributes=True).model_dump() + data = overlay_response_streams( + data, + await get_response_streams_by_chat_id(request.app.state.redis, id), + ) data['context_usage'] = await get_chat_context_usage(chat) return data diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index d88fd827f1..9224501767 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -927,7 +927,7 @@ async def _make_channel_emitter(request_info): channel_id = request_info['chat_id'].removeprefix('channel:') message_id = request_info['message_id'] - state = {'last_emit_at': 0.0} + state = {'last_emit_at': 0.0, 'output': []} THROTTLE_INTERVAL = 0.15 # ~6 updates/sec async def _emit_channel_update(content: str, done: bool = False, output: list | None = None): @@ -975,11 +975,23 @@ async def _make_channel_emitter(request_info): if not content and not output and not done: return - now = __import__('time').time() + now = time.time() if done or (now - state['last_emit_at']) >= THROTTLE_INTERVAL: state['last_emit_at'] = now await _emit_channel_update(content, done, output if isinstance(output, list) else None) + elif event_type == 'response:completion': + from open_webui.utils.middleware import handle_responses_streaming_event + + data = event_data.get('data', {}) + state['output'], _ = handle_responses_streaming_event(data, state['output']) + content = get_output_text(state['output']) + + now = time.time() + if content and (now - state['last_emit_at']) >= THROTTLE_INTERVAL: + state['last_emit_at'] = now + await _emit_channel_update(content, False, state['output']) + elif event_type == 'chat:message:error': error = event_data.get('data', {}).get('error', {}) error_content = error.get('content', 'An error occurred') if isinstance(error, dict) else str(error) diff --git a/backend/open_webui/tasks.py b/backend/open_webui/tasks.py index 2e6193e464..81b45ebdcd 100644 --- a/backend/open_webui/tasks.py +++ b/backend/open_webui/tasks.py @@ -13,10 +13,12 @@ log = logging.getLogger(__name__) # A dictionary to keep track of active tasks tasks: dict[str, asyncio.Task] = {} item_tasks = {} +response_streams: dict[str, dict] = {} REDIS_TASKS_KEY = f'{REDIS_KEY_PREFIX}:tasks' REDIS_ITEM_TASKS_KEY = f'{REDIS_KEY_PREFIX}:tasks:item' +REDIS_RESPONSE_STREAMS_KEY = f'{REDIS_KEY_PREFIX}:tasks:response_streams' REDIS_PUBSUB_CHANNEL = f'{REDIS_KEY_PREFIX}:tasks:commands' @@ -55,6 +57,7 @@ async def redis_save_task(redis: Redis, task_id: str, item_id: str | None): async def redis_cleanup_task(redis: Redis, task_id: str, item_id: str | None): pipe = redis.pipeline() pipe.hdel(REDIS_TASKS_KEY, task_id) + pipe.hdel(REDIS_RESPONSE_STREAMS_KEY, task_id) if item_id: pipe.srem(f'{REDIS_ITEM_TASKS_KEY}:{item_id}', task_id) await pipe.execute() @@ -91,6 +94,7 @@ async def cleanup_task(redis, task_id: str, id=None): await redis_cleanup_task(redis, task_id, id) tasks.pop(task_id, None) # Remove the task if it exists + response_streams.pop(task_id, None) # If an ID is provided, remove the task from the item_tasks dictionary if id and task_id in item_tasks.get(id, []): @@ -140,6 +144,63 @@ async def list_task_ids_by_item_id(redis, id): return item_tasks.get(id, []) +async def save_response_stream( + redis, + task_id: str | None, + chat_id: str | None, + message_id: str | None, + content: str, + output: list, +): + if not task_id or not chat_id or not message_id: + return + + data = { + 'chat_id': chat_id, + 'message_id': message_id, + 'content': content, + 'output': output, + } + + if redis: + await redis.hset(REDIS_RESPONSE_STREAMS_KEY, task_id, JSONCodec.dumps(data)) + else: + response_streams[task_id] = data + + +async def get_response_streams_by_chat_id(redis, chat_id: str) -> list[dict]: + task_ids = await list_task_ids_by_item_id(redis, chat_id) + if not task_ids: + return [] + + if redis: + values = await redis.hmget(REDIS_RESPONSE_STREAMS_KEY, task_ids) + streams = [] + for value in values: + if not value: + continue + try: + data = JSONCodec.loads(value) + except Exception: + continue + if data.get('chat_id') == chat_id: + streams.append(data) + return streams + + return [ + stream for task_id in task_ids if (stream := response_streams.get(task_id)) and stream.get('chat_id') == chat_id + ] + + +async def clear_response_stream(redis, task_id: str | None): + if not task_id: + return + if redis: + await redis.hdel(REDIS_RESPONSE_STREAMS_KEY, task_id) + else: + response_streams.pop(task_id, None) + + async def stop_task(redis, task_id: str): """ Cancel a running task and remove it from the global task list. diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 4ccf1823f5..b7409cab76 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -76,6 +76,7 @@ from open_webui.socket.main import ( get_event_call, get_event_emitter, ) +from open_webui.tasks import clear_response_stream, save_response_stream from open_webui.utils.access_control import has_connection_access, has_permission from open_webui.utils.access_control.files import get_owner_accessible_folder_files from open_webui.utils.access_control.folders import has_folder_access @@ -499,7 +500,22 @@ def handle_responses_streaming_event( item = data.get('item', {}) if item: new_output = list(current_output) - new_output.append(item) + output_index = data.get('output_index', len(new_output)) + existing_index = next( + ( + idx + for idx, existing in enumerate(new_output) + if (item.get('id') and existing.get('id') == item.get('id')) + or (item.get('call_id') and existing.get('call_id') == item.get('call_id')) + ), + None, + ) + if existing_index is not None: + new_output[existing_index] = item + elif 0 <= output_index < len(new_output): + new_output.insert(output_index, item) + else: + new_output.append(item) return new_output, None return current_output, None @@ -4033,6 +4049,7 @@ async def streaming_chat_response_handler(response, ctx): async def response_handler(response, events): filter_context = FilterContext() tag_scan_positions = {} + response_stream_task_id = metadata.get('task_id') or metadata.get('message_id') def tag_output_handler(content_type, tags, output): """ @@ -4440,38 +4457,108 @@ async def streaming_chat_response_handler(response, ctx): ) last_delta_data = None last_delta_type = None + last_delta_key = None + + def response_stream_content(stream_output: list | None = None): + return ''.join(content_parts) or get_output_text( + stream_output if stream_output is not None else full_output() + ) + + async def save_current_response_stream(stream_output: list | None = None): + if not chat_id or not metadata.get('message_id'): + return + + current_stream_output = stream_output if stream_output is not None else full_output() + await save_response_stream( + request.app.state.redis, + response_stream_task_id, + chat_id, + metadata.get('message_id'), + response_stream_content(current_stream_output), + current_stream_output, + ) + + def get_response_delta_key(delta_data: dict): + event_type = delta_data.get('type', '') + if not event_type.startswith('response.') or not event_type.endswith('.delta'): + return None + return ( + event_type, + delta_data.get('item_id'), + delta_data.get('output_index'), + delta_data.get('content_index'), + delta_data.get('summary_index'), + ) async def flush_pending_delta_data(threshold: int = 0): nonlocal delta_count nonlocal last_delta_data nonlocal last_delta_type + nonlocal last_delta_key if delta_count >= threshold and last_delta_data: await event_emitter( { - 'type': 'chat:completion', + 'type': 'response:completion', 'data': last_delta_data, } ) + await save_current_response_stream() delta_count = 0 last_delta_data = None last_delta_type = None + last_delta_key = None async def queue_pending_delta_data(delta_data: dict, delta_type: str): nonlocal delta_count nonlocal last_delta_data nonlocal last_delta_type + nonlocal last_delta_key - if last_delta_type and last_delta_type != delta_type: - await flush_pending_delta_data() + delta_key = get_response_delta_key(delta_data) + if ( + last_delta_data + and last_delta_key == delta_key + and isinstance(last_delta_data.get('delta'), str) + and isinstance(delta_data.get('delta'), str) + ): + last_delta_data['delta'] += delta_data['delta'] + delta_count += 1 + else: + if last_delta_data and (last_delta_type != delta_type or last_delta_key != delta_key): + await flush_pending_delta_data() - delta_count += 1 - last_delta_data = delta_data - last_delta_type = delta_type + delta_count += 1 + last_delta_data = delta_data + last_delta_type = delta_type + last_delta_key = delta_key if delta_count >= delta_chunk_size: await flush_pending_delta_data(delta_chunk_size) + async def emit_response_completion_event(response_event: dict, stream_output: list | None = None): + if prior_output and isinstance(response_event.get('output_index'), int): + response_event = { + **response_event, + 'output_index': response_event['output_index'] + len(prior_output), + } + + if response_event.get('type', '').endswith('.delta'): + await queue_pending_delta_data( + response_event, + response_event.get('type', 'response.delta'), + ) + return + + await flush_pending_delta_data() + await event_emitter( + { + 'type': 'response:completion', + 'data': response_event, + } + ) + await save_current_response_stream(stream_output) + filter_extra_params = {'__body__': form_data, **extra_params} if filter_functions else None async for line in response.body_iterator: @@ -4587,13 +4674,6 @@ async def streaming_chat_response_handler(response, ctx): } ) - processed_data = { - 'output': full_output(), - } - - # print(data) - # print(processed_data) - # Merge any metadata (usage, etc.) # Strip 'done' — response.completed emits # it but we may still need to execute tool @@ -4610,22 +4690,21 @@ async def streaming_chat_response_handler(response, ctx): usage = merge_usage(usage, response_metadata['usage']) response_metadata['usage'] = usage - processed_data.update(response_metadata) - processed_data.pop('done', None) + if response_metadata.get('error'): + await event_emitter( + { + 'type': 'chat:completion', + 'data': {'error': response_metadata['error']}, + } + ) - if response_event_is_delta: - response_delta_type = response_event_type.split('.')[1] - await queue_pending_delta_data( - processed_data, - 'tool_call' - if response_delta_type == 'function_call_arguments' - else 'content', - ) - else: + await emit_response_completion_event(data) + + if response_metadata and response_metadata.get('usage'): await event_emitter( { 'type': 'chat:completion', - 'data': processed_data, + 'data': {'usage': usage}, } ) continue @@ -4723,6 +4802,7 @@ async def streaming_chat_response_handler(response, ctx): # Add the new tool call delta_tool_call.setdefault('function', {}) delta_tool_call['function'].setdefault('name', '') + delta_tool_call['id'] = delta_tool_call.get('id') or output_id('fc') delta_arguments = delta_tool_call['function'].get('arguments') if not isinstance(delta_arguments, str): delta_tool_call['function']['arguments'] = ( @@ -4754,27 +4834,73 @@ async def streaming_chat_response_handler(response, ctx): delta_arguments ) - # Emit pending tool calls in real-time + # Emit pending tool calls in real-time as Responses events. if response_tool_calls: - # Build pending function_call output items for display - pending_fc_items = [] + output_by_call_id = { + item.get('call_id'): (idx, item) + for idx, item in enumerate(output) + if item.get('type') == 'function_call' + } + for tc in response_tool_calls: - call_id = tc.get('id', '') + call_id = tc.get('id') or output_id('fc') + tc['id'] = call_id func = tc.get('function', {}) - pending_fc_items.append( - { + if call_id in output_by_call_id: + output_index, item = output_by_call_id[call_id] + item['name'] = func.get('name', item.get('name', '')) + item['arguments'] = func.get('arguments', item.get('arguments', '')) + item['status'] = 'in_progress' + else: + output_index = len(output) + item = { 'type': 'function_call', - 'id': call_id or output_id('fc'), + 'id': call_id, 'call_id': call_id, 'name': func.get('name', ''), - 'arguments': func.get('arguments', '{}'), + 'arguments': '', 'status': 'in_progress', } - ) + output.append(item) + output_by_call_id[call_id] = (output_index, item) + await emit_response_completion_event( + { + 'type': 'response.output_item.added', + 'output_index': output_index, + 'item': item.copy(), + } + ) + item['arguments'] = func.get('arguments', '') - data = { - 'output': full_output() + pending_fc_items, - } + for delta_tool_call in delta_tool_calls: + tool_call_index = delta_tool_call.get('index') + current_response_tool_call = next( + ( + tc + for tc in response_tool_calls + if tc.get('index') == tool_call_index + ), + None, + ) + if not current_response_tool_call: + continue + call_id = current_response_tool_call.get('id') + output_index, _ = output_by_call_id.get(call_id, (len(output) - 1, {})) + delta_arguments = delta_tool_call.get('function', {}).get('arguments') + if delta_arguments is not None: + if not isinstance(delta_arguments, str): + delta_arguments = JSONCodec.dumps(delta_arguments) + await emit_response_completion_event( + { + 'type': 'response.function_call_arguments.delta', + 'item_id': call_id, + 'output_index': output_index, + 'delta': delta_arguments, + } + ) + + await save_current_response_stream() + data = None delta_type = 'tool_call' delta_images = delta.get('images') @@ -4878,20 +5004,26 @@ async def streaming_chat_response_handler(response, ctx): } ] + reasoning_index = output.index(reasoning_item) data = { - 'output': full_output(), + 'type': 'response.reasoning_text.delta', + 'item_id': reasoning_item.get('id'), + 'output_index': reasoning_index, + 'content_index': max( + len(reasoning_item.get('content', [])) - 1, + 0, + ), + 'delta': reasoning_content, } - delta_type = 'content' + delta_type = 'response.reasoning_text.delta' if reasoning_detail_items: merge_streamed_reasoning_details( reasoning_item.setdefault('reasoning_details', []), reasoning_detail_items, ) - data = { - 'output': full_output(), - } - delta_type = 'content' + await save_current_response_stream() + data = None if value: if ( @@ -5037,29 +5169,27 @@ async def streaming_chat_response_handler(response, ctx): if end: break - if ENABLE_REALTIME_CHAT_SAVE and save_to_chat: - current_output = full_output() - # Save message in the database - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - { - 'output': current_output, - }, - ) - data = { - 'output': current_output, - } - delta_type = 'content' - else: - data = { - 'output': full_output(), - } - delta_type = 'content' + target_index = len(output) - 1 + target_item = output[target_index] if target_index >= 0 else {} + target_content = target_item.get('content', []) + content_index = max(len(target_content) - 1, 0) + delta_event_type = ( + 'response.reasoning_text.delta' + if target_item.get('type') == 'reasoning' + else 'response.output_text.delta' + ) + data = { + 'type': delta_event_type, + 'item_id': target_item.get('id'), + 'output_index': target_index, + 'content_index': content_index, + 'delta': value, + } + delta_type = delta_event_type - if delta: + if delta and data: await queue_pending_delta_data(data, delta_type) - else: + elif data: await event_emitter( { 'type': 'chat:completion', @@ -5108,6 +5238,29 @@ async def streaming_chat_response_handler(response, ctx): reasoning_item['status'] = 'completed' if response_tool_calls: + for tc in response_tool_calls: + call_id = tc.get('id', '') + arguments = tc.get('function', {}).get('arguments', '{}') + for output_index, item in enumerate(output): + if item.get('type') == 'function_call' and item.get('call_id') == call_id: + item['arguments'] = arguments + item['status'] = 'completed' + await emit_response_completion_event( + { + 'type': 'response.function_call_arguments.done', + 'item_id': item.get('id'), + 'output_index': output_index, + 'arguments': arguments, + } + ) + await emit_response_completion_event( + { + 'type': 'response.output_item.done', + 'output_index': output_index, + 'item': item.copy(), + } + ) + break tool_calls.append(_split_tool_calls(response_tool_calls)) # Responses API path: extract function_call items from output @@ -5809,31 +5962,22 @@ async def streaming_chat_response_handler(response, ctx): } if save_to_chat: - if not ENABLE_REALTIME_CHAT_SAVE: - # Save message in the database - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - { - 'done': True, - 'output': current_output, - **({'usage': usage} if usage else {}), - }, - ) - elif usage: - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - {'done': True, 'usage': usage}, - ) - else: - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - {'done': True}, - ) + # Save final output once. The delta path keeps in-progress + # state in response_streams instead of writing tokens to DB. + await Chats.upsert_message_to_chat_by_id_and_message_id( + metadata['chat_id'], + metadata['message_id'], + { + 'done': True, + 'output': current_output, + **({'usage': usage} if usage else {}), + }, + ) - await publish_chat_finished_event(request, user, metadata, title, ''.join(content_parts), current_output) + await clear_response_stream(request.app.state.redis, response_stream_task_id) + await publish_chat_finished_event( + request, user, metadata, title, ''.join(content_parts), current_output + ) await event_emitter( { @@ -5865,22 +6009,15 @@ async def streaming_chat_response_handler(response, ctx): async def save_cancelled_state(): await event_emitter({'type': 'chat:tasks:cancel'}) if save_to_chat: - if not ENABLE_REALTIME_CHAT_SAVE: - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - { - 'done': True, - 'output': full_output(), - }, - ) - else: - await Chats.upsert_message_to_chat_by_id_and_message_id( - metadata['chat_id'], - metadata['message_id'], - {'done': True}, - touch=False, - ) + await Chats.upsert_message_to_chat_by_id_and_message_id( + metadata['chat_id'], + metadata['message_id'], + { + 'done': True, + 'output': full_output(), + }, + ) + await clear_response_stream(request.app.state.redis, response_stream_task_id) try: await asyncio.shield(save_cancelled_state()) diff --git a/src/lib/components/chat/Chat.svelte b/src/lib/components/chat/Chat.svelte index f63f17bb19..b0f345348d 100644 --- a/src/lib/components/chat/Chat.svelte +++ b/src/lib/components/chat/Chat.svelte @@ -66,7 +66,7 @@ } from '$lib/utils'; import { AudioQueue } from '$lib/utils/audio'; import { createTemporaryChatId, isTemporaryChatId } from '$lib/utils/chatId'; - import { getOutputText } from './Messages/structuredOutput'; + import { applyResponseStreamEvent, getOutputText } from './Messages/structuredOutput'; import { archiveChatById, @@ -1178,6 +1178,8 @@ updateLastReadAt($chatId); } } + } else if (type === 'response:completion') { + responseCompletionEventHandler(data, message); } else if (type === 'chat:completion') { chatCompletionEventHandler(data, message, event.chat_id); } else if (type === 'chat:tasks:cancel') { @@ -2693,6 +2695,27 @@ } }; + const responseCompletionEventHandler = (data, message) => { + message.output = applyResponseStreamEvent(message.output ?? [], data); + + if (data?.type === 'response.output_text.delta') { + const value = data.delta ?? ''; + if (!(message.content == '' && value == '\n')) { + message.content += value; + + if (navigator.vibrate && ($settings?.hapticFeedback ?? false)) { + navigator.vibrate(5); + } + dispatchCallOverlayAudio(message); + } + } else if (data?.type === 'response.completed' || data?.type?.endsWith('.done')) { + message.content = getOutputText(message.output) || message.content; + } + + history.messages[message.id] = message; + history = history; + }; + const chatCompletionEventHandler = async (data, message, chatId) => { const { id, done, choices, content, output, sources, selected_model_id, error, usage } = data; diff --git a/src/lib/components/chat/Messages/structuredOutput.ts b/src/lib/components/chat/Messages/structuredOutput.ts index d0299c4f35..40e509457a 100644 --- a/src/lib/components/chat/Messages/structuredOutput.ts +++ b/src/lib/components/chat/Messages/structuredOutput.ts @@ -59,6 +59,24 @@ export type OutputDisplayItem = tokens: OutputDetailToken[]; }; +type ResponseStreamEvent = { + type?: string; + item_id?: string; + output_index?: number; + content_index?: number; + summary_index?: number; + item?: OutputItem; + part?: OutputContentPart; + delta?: unknown; + text?: unknown; + arguments?: unknown; + response?: { + output?: OutputItem[]; + [key: string]: unknown; + }; + [key: string]: unknown; +}; + const GROUPABLE_OUTPUT_TYPES = new Set([ 'reasoning', 'function_call', @@ -352,6 +370,167 @@ export function getOutputText(output?: OutputItem[] | null): string { .join('\n'); } +function appendDelta(current: unknown, delta: unknown): unknown { + if (typeof current === 'string' || typeof delta === 'string') { + return `${current ?? ''}${delta ?? ''}`; + } + if ( + current && + delta && + typeof current === 'object' && + typeof delta === 'object' && + !Array.isArray(current) && + !Array.isArray(delta) + ) { + return { ...(current as Record), ...(delta as Record) }; + } + return delta ?? current ?? ''; +} + +function ensureItem(output: OutputItem[], outputIndex: number, fallback?: OutputItem): OutputItem { + while (output.length <= outputIndex) { + output.push( + fallback ?? { type: 'message', status: 'in_progress', role: 'assistant', content: [] } + ); + } + output[outputIndex] = { ...output[outputIndex] }; + return output[outputIndex]; +} + +function ensurePart(parts: OutputContentPart[], index: number, fallback?: OutputContentPart) { + while (parts.length <= index) { + parts.push(fallback ?? { type: 'output_text', text: '' }); + } + parts[index] = { ...parts[index] }; + return parts[index]; +} + +function findOutputItemIndex(output: OutputItem[], item: OutputItem): number { + return output.findIndex( + (existing) => + (!!item.id && existing?.id === item.id) || + (!!item.call_id && existing?.call_id === item.call_id) + ); +} + +export function applyResponseStreamEvent( + output: OutputItem[] = [], + event: ResponseStreamEvent +): OutputItem[] { + const eventType = event?.type ?? ''; + if (!eventType.startsWith('response.')) { + return output; + } + + if (eventType === 'response.completed') { + return event.response?.output ? [...event.response.output] : output; + } + + const nextOutput = [...output]; + const eventItemIndex = event.item_id + ? nextOutput.findIndex((item) => item?.id === event.item_id || item?.call_id === event.item_id) + : -1; + const outputIndex = + eventItemIndex >= 0 ? eventItemIndex : (event.output_index ?? Math.max(output.length - 1, 0)); + + if (eventType === 'response.output_item.added') { + if (!event.item) { + return output; + } + const item = { ...event.item }; + const existingIndex = findOutputItemIndex(nextOutput, item); + if (existingIndex >= 0) { + nextOutput[existingIndex] = item; + } else if (outputIndex < nextOutput.length) { + nextOutput.splice(outputIndex, 0, item); + } else { + nextOutput[outputIndex] = item; + } + return nextOutput; + } + + if (eventType === 'response.output_item.done') { + if (!event.item) { + return output; + } + const item = { ...event.item }; + const existingIndex = findOutputItemIndex(nextOutput, item); + nextOutput[existingIndex >= 0 ? existingIndex : outputIndex] = item; + return nextOutput; + } + + const item = ensureItem(nextOutput, outputIndex, { + id: event.item_id, + type: eventType.includes('reasoning') + ? 'reasoning' + : eventType.includes('function_call') + ? 'function_call' + : 'message', + status: 'in_progress', + role: 'assistant', + content: [] + }); + + if (eventType === 'response.content_part.added') { + if (item.type === 'reasoning' || !event.part) { + return nextOutput; + } + item.content = [...(item.content ?? [])]; + item.content[event.content_index ?? item.content.length] = { ...event.part }; + return nextOutput; + } + + if (eventType === 'response.reasoning_summary_part.added') { + if (!event.part) { + return nextOutput; + } + item.summary = [...(item.summary ?? [])]; + item.summary[event.summary_index ?? item.summary.length] = { ...event.part }; + return nextOutput; + } + + if (eventType.endsWith('.delta')) { + const deltaType = eventType.split('.')[1]; + if (deltaType === 'function_call_arguments') { + item.arguments = appendDelta(item.arguments ?? '', event.delta); + return nextOutput; + } + + if (deltaType === 'reasoning_summary_text') { + const summaryIndex = event.summary_index ?? 0; + item.summary = [...(item.summary ?? [])]; + const part = ensurePart(item.summary, summaryIndex, { type: 'summary_text', text: '' }); + part.text = appendDelta(part.text ?? '', event.delta); + return nextOutput; + } + + const key = deltaType === 'output_text' || deltaType === 'reasoning_text' ? 'text' : deltaType; + item.content = [...(item.content ?? [])]; + const part = ensurePart(item.content, event.content_index ?? 0); + part[key] = appendDelta(part[key], event.delta); + return nextOutput; + } + + if (eventType.endsWith('.done')) { + const typeName = eventType.split('.')[1]; + if (typeName === 'content_part' && event.part) { + item.content = [...(item.content ?? [])]; + item.content[event.content_index ?? Math.max(item.content.length - 1, 0)] = { ...event.part }; + } else if (typeName === 'function_call_arguments' && event.arguments !== undefined) { + item.arguments = event.arguments; + } else if ( + (typeName === 'output_text' || typeName === 'text' || typeName === 'reasoning_text') && + event.text !== undefined + ) { + item.content = [...(item.content ?? [])]; + const part = ensurePart(item.content, event.content_index ?? 0); + part.text = event.text; + } + } + + return nextOutput; +} + export function replaceOutputMessageText( output: OutputItem[] = [], oldContent: string,