diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 72d516cf22..302e2f6f86 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 +import weakref from contextlib import suppress from typing import Any @@ -209,6 +210,7 @@ REDIS_EVENT_CHANNEL = f'{REDIS_KEY_PREFIX}:direct_completion' EVENT_QUEUES: dict[str, asyncio.Queue] = {} EVENT_PUBLISH_LOCK = asyncio.Lock() +SESSION_EVENT_LOCKS: weakref.WeakValueDictionary[str, asyncio.Lock] = weakref.WeakValueDictionary() def get_session_pool_batches(): @@ -1068,19 +1070,22 @@ async def socket_event_handler(event: Any, sid: str, *args: Any) -> None: 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 + # The lock keeps arrival order; the sid re-check drops every event after a failed check + session_check = asyncio.create_task(get_socket_session_user(sid)) + async with SESSION_EVENT_LOCKS.setdefault(sid, asyncio.Lock()): + user = await session_check + if not user or sid not in LOCAL_AUTHENTICATED_SIDS 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) + 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):