From 0180efecf362d487e0c30f040f5948c325fbe337 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Mon, 21 Sep 2026 08:43:45 -0400 Subject: [PATCH] refac --- backend/open_webui/main.py | 8 ++++ backend/open_webui/socket/main.py | 72 ++++++++++++++++++++++++++++++- backend/open_webui/utils/chat.py | 28 ++++-------- 3 files changed, 86 insertions(+), 22 deletions(-) diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index 3cf7619332..0e77bd4fe0 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -113,6 +113,7 @@ from open_webui.env import ( SCIM_TOKEN, VERSION, WEBSOCKET_HEARTBEAT_INTERVAL, + WEBSOCKET_MANAGER, # Admin Account Runtime Creation WEBUI_ADMIN_EMAIL, WEBUI_ADMIN_NAME, @@ -189,6 +190,7 @@ from open_webui.socket.main import ( get_user_id_from_session_pool, periodic_session_pool_cleanup, periodic_usage_pool_cleanup, + redis_event_listener, ) from open_webui.socket.main import ( app as socket_app, @@ -387,6 +389,9 @@ async def lifespan(app: FastAPI): if app.state.redis is not None: app.state.redis_task_command_listener = asyncio.create_task(redis_task_command_listener(app)) + if WEBSOCKET_MANAGER == 'redis': + app.state.redis_event_listener = asyncio.create_task(redis_event_listener()) + app.state.periodic_usage_pool_cleanup = asyncio.create_task(periodic_usage_pool_cleanup()) app.state.periodic_session_pool_cleanup = asyncio.create_task(periodic_session_pool_cleanup()) @@ -473,6 +478,9 @@ async def lifespan(app: FastAPI): if hasattr(app.state, 'redis_task_command_listener'): app.state.redis_task_command_listener.cancel() + if hasattr(app.state, 'redis_event_listener'): + app.state.redis_event_listener.cancel() + app.state.periodic_usage_pool_cleanup.cancel() app.state.periodic_session_pool_cleanup.cancel() app.state.scheduler_worker_loop.cancel() diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index d4ac18f655..f68252489b 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -6,6 +6,7 @@ import logging import random import sys import time +from contextlib import suppress from typing import Any import pycrdt as Y @@ -38,17 +39,23 @@ from open_webui.models.folders import Folders from open_webui.models.notes import Notes, NoteUpdateForm from open_webui.models.users import UserNameResponse, Users from open_webui.socket.utils import RedisDict, RedisLock, YdocManager -from open_webui.tasks import create_task, stop_item_tasks +from open_webui.tasks import ( + REDIS_PUBSUB_MAX_RECONNECT_INTERVAL, + REDIS_PUBSUB_RECONNECT_INTERVAL, + create_task, + stop_item_tasks, +) from open_webui.utils.access_control import has_permission from open_webui.utils.auth import get_verified_user_by_token from open_webui.utils.chat_id import is_saved_chat_id -from open_webui.utils.json_codec import SOCKETIO_JSON +from open_webui.utils.json_codec import SOCKETIO_JSON, JSONCodec, dumps_bytes from open_webui.utils.misc import get_output_text from open_webui.utils.redis import ( build_sentinel_url, get_redis_connection, get_sentinels_from_env, ) +from redis.exceptions import RedisError from socketio.packet import Packet logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL) @@ -190,6 +197,11 @@ YDOC_MANAGER = YdocManager( redis_key_prefix=f'{REDIS_KEY_PREFIX}:ydoc:documents', ) +REDIS_EVENT_CHANNEL = f'{REDIS_KEY_PREFIX}:direct_completion' + +EVENT_QUEUES: dict[str, asyncio.Queue] = {} +EVENT_PUBLISH_LOCK = asyncio.Lock() + def get_session_pool_batches(): """All session pool entries, in bounded batches for the Redis backing.""" @@ -948,6 +960,62 @@ async def disconnect(sid, reason=None): # print(f"Unknown session ID {sid} disconnected") +async def redis_event_listener() -> None: + """Route events received over Redis to their local queues.""" + reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL + + while True: + pubsub = None + try: + # RedisCluster can't route a pubsub subscribe until initialize() fills its slot cache. + await REDIS.initialize() + + pubsub = REDIS.pubsub() + await pubsub.subscribe(REDIS_EVENT_CHANNEL) + reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL + + async for message in pubsub.listen(): + if message['type'] != 'message': + continue + event = JSONCodec.loads(message['data']) + queue = EVENT_QUEUES.get(event['channel']) + if queue is not None: + await queue.put(event['data']) + log.warning('Redis event listener stopped. Retrying.') + except asyncio.CancelledError: + raise + except Exception: + log.exception('Redis event listener failed. Retrying.') + finally: + if pubsub: + with suppress(Exception): + await pubsub.aclose() + + await asyncio.sleep(reconnect_interval) + reconnect_interval = min(reconnect_interval * 2, REDIS_PUBSUB_MAX_RECONNECT_INTERVAL) + + +@sio.on('*') +async def socket_event_handler(event: Any, sid: str, *args: Any) -> None: + """Route user-owned stream events to a local queue or another worker.""" + if not isinstance(event, str) or event.count(':') != 2 or not args: + return + + user = await get_socket_session_user(sid) + if not user or user.get('id') != event.split(':', 1)[0]: + return + + queue = EVENT_QUEUES.get(event) + if queue is not None: + await queue.put(args[0]) + elif WEBSOCKET_MANAGER == 'redis': + try: + async with EVENT_PUBLISH_LOCK: + await REDIS.publish(REDIS_EVENT_CHANNEL, dumps_bytes({'channel': event, 'data': args[0]})) + except RedisError as e: + log.debug('Failed to relay socket event %s: %s', event, e) + + async def _make_channel_emitter(request_info): """Event emitter that routes pipeline output to a channel message. diff --git a/backend/open_webui/utils/chat.py b/backend/open_webui/utils/chat.py index 0b64a52165..5d5f6f60c4 100644 --- a/backend/open_webui/utils/chat.py +++ b/backend/open_webui/utils/chat.py @@ -23,9 +23,9 @@ from open_webui.routers.pipelines import ( process_pipeline_outlet_filter, ) from open_webui.socket.main import ( + EVENT_QUEUES, get_event_call, get_event_emitter, - sio, ) from open_webui.utils.filter import ( get_filter_functions, @@ -71,19 +71,8 @@ async def generate_direct_chat_completion( logging.info('WebSocket channel: %s', channel) if form_data.get('stream'): - q = asyncio.Queue() - - async def message_listener(sid, data): - """ - Handle received socket messages and push them into the queue. - """ - await q.put(data) - - def remove_message_listener(): - sio.handlers['/'].pop(channel, None) - - # Register the listener - sio.on(channel, message_listener) + queue = asyncio.Queue() + EVENT_QUEUES[channel] = queue # Start processing chat completion in background try: @@ -103,16 +92,15 @@ async def generate_direct_chat_completion( status = res.get('status', False) except BaseException: - remove_message_listener() + EVENT_QUEUES.pop(channel, None) raise if status: # Define a generator to stream responses async def event_generator(): - nonlocal q try: while True: - data = await q.get() # Wait for new messages + data = await queue.get() # Wait for new messages if isinstance(data, dict): if 'done' in data and data['done']: break # Stop streaming when 'done' is received @@ -127,16 +115,16 @@ async def generate_direct_chat_completion( log.debug('Error in event generator: %s', e) pass finally: - remove_message_listener() + EVENT_QUEUES.pop(channel, None) # Define a background task to run the event generator async def background(): - remove_message_listener() + EVENT_QUEUES.pop(channel, None) # Return the streaming response return StreamingResponse(event_generator(), media_type='text/event-stream', background=background) else: - remove_message_listener() + EVENT_QUEUES.pop(channel, None) raise Exception(str(res)) else: res = await event_caller(