This commit is contained in:
Timothy Jaeryang Baek 2026-09-21 08:43:45 -04:00
parent 7fa8673296
commit 0180efecf3
3 changed files with 86 additions and 22 deletions

View file

@ -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()

View file

@ -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.

View file

@ -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(