Files
open-webui/backend/open_webui/socket/main.py
Classic298 e140d8f3cc fix: scope timer cancellation to the timer's owner (#27472)
The events:chat socket handler called the ownership-checked update for last_read_at, discarded the boolean it returns, and then cancelled the chat's pending timers regardless of the answer. cancel_timers_for_chat selected on the internal marker, the type, the parent chat id and the status, and never on the owner, so it matched rows belonging to any user. An authenticated user who knew another user's chat id could mark that chat read over their own socket session and silently cancel the owner's pending timers, and the owner got no notification: the scheduled action simply never fired.

The missing owner predicate also cut the other way in ordinary use. Because the query matched every timer sharing a parent chat id, one user reading a chat cancelled the timers of anyone else holding one on the same chat, so this was collateral damage as much as an attack.

cancel_timers_for_chat now requires a user_id and filters on it, which is the durable fix, and the socket handler returns early unless the ownership-checked update reports that the caller owns the chat. The parameter is required rather than defaulted so a later caller cannot reintroduce the unscoped query by omission. Both existing call sites already know the acting user. Timer rows are created with the same owner as the parent chat and the execution path already refuses to run one whose owner does not match, so scoping the cancellation the same way cannot strand a timer that would otherwise have fired.

One behaviour change worth noting: an administrator posting into another user's chat no longer cancels that user's chat.user_message timers, because the acting user is the administrator. The timer fires instead of being cancelled, which is the safe direction.
2026-07-26 18:12:44 -04:00

1093 lines
37 KiB
Python

from __future__ import annotations
import asyncio
import logging
import random
import sys
import time
import pycrdt as Y
import socketio
from open_webui.config import (
CORS_ALLOW_ORIGIN,
)
from open_webui.env import (
ENABLE_WEBSOCKET_SUPPORT,
GLOBAL_LOG_LEVEL,
REDIS_KEY_PREFIX,
WEBSOCKET_EVENT_CALLER_TIMEOUT,
WEBSOCKET_MANAGER,
WEBSOCKET_REDIS_CLUSTER,
WEBSOCKET_REDIS_LOCK_TIMEOUT,
WEBSOCKET_REDIS_OPTIONS,
WEBSOCKET_REDIS_URL,
WEBSOCKET_SENTINEL_HOSTS,
WEBSOCKET_SENTINEL_PORT,
WEBSOCKET_SERVER_ENGINEIO_LOGGING,
WEBSOCKET_SERVER_LOGGING,
WEBSOCKET_SERVER_PING_INTERVAL,
WEBSOCKET_SERVER_PING_TIMEOUT,
)
from open_webui.models.access_grants import AccessGrants
from open_webui.models.channels import Channels
from open_webui.models.chats import Chats
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.utils.access_control import has_permission
from open_webui.utils.auth import get_verified_user_by_token
from open_webui.utils.redis import (
build_sentinel_url,
get_redis_connection,
get_sentinels_from_env,
)
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
log = logging.getLogger(__name__)
# Let no connection opened in good faith be dropped without
# cause, and let every message find the room it was meant for.
REDIS = None
# Configure CORS for Socket.IO
SOCKETIO_CORS_ORIGINS = '*' if CORS_ALLOW_ORIGIN == ['*'] else CORS_ALLOW_ORIGIN
def get_room_sid_map(manager, namespace: str, room: str):
"""Return this process's Socket.IO sid map for a room, without copying it."""
return manager.rooms.get(namespace, {}).get(room)
class LocalFilteredRedisManager(socketio.AsyncRedisManager):
"""AsyncRedisManager that drops pub/sub emits with no local recipients.
Every instance subscribed to the shared Socket.IO channel receives every
emit published by the whole fleet. Upstream ``_handle_emit`` re-encodes
the full packet before discovering the target room has no participants on
this instance, so each instance burns CPU serializing payloads addressed
to sessions it does not host — a cost that grows with instance count.
Bail out before that work when the room is empty here. Only string
rooms take the fast path; broadcasts (``room=None``) and any other
room shape (e.g. lists, which upstream indexes before validating)
always pass through unchanged.
"""
async def _handle_emit(self, message):
room = message.get('room')
if isinstance(room, str):
namespace = message.get('namespace') or '/'
if not get_room_sid_map(self, namespace, room):
return
await super()._handle_emit(message)
if WEBSOCKET_MANAGER == 'redis':
sentinel_hosts = WEBSOCKET_SENTINEL_HOSTS or ''
ws_redis_url = (
build_sentinel_url(WEBSOCKET_REDIS_URL, sentinel_hosts, WEBSOCKET_SENTINEL_PORT)
if sentinel_hosts
else WEBSOCKET_REDIS_URL
)
redis_manager = LocalFilteredRedisManager(ws_redis_url, redis_options=WEBSOCKET_REDIS_OPTIONS)
sio = socketio.AsyncServer(
cors_allowed_origins=SOCKETIO_CORS_ORIGINS,
async_mode='asgi',
transports=(['websocket'] if ENABLE_WEBSOCKET_SUPPORT else ['polling']),
allow_upgrades=ENABLE_WEBSOCKET_SUPPORT,
always_connect=True,
client_manager=redis_manager,
logger=WEBSOCKET_SERVER_LOGGING,
ping_interval=WEBSOCKET_SERVER_PING_INTERVAL,
ping_timeout=WEBSOCKET_SERVER_PING_TIMEOUT,
engineio_logger=WEBSOCKET_SERVER_ENGINEIO_LOGGING,
)
else:
sio = socketio.AsyncServer(
cors_allowed_origins=SOCKETIO_CORS_ORIGINS,
async_mode='asgi',
transports=(['websocket'] if ENABLE_WEBSOCKET_SUPPORT else ['polling']),
allow_upgrades=ENABLE_WEBSOCKET_SUPPORT,
always_connect=True,
logger=WEBSOCKET_SERVER_LOGGING,
ping_interval=WEBSOCKET_SERVER_PING_INTERVAL,
ping_timeout=WEBSOCKET_SERVER_PING_TIMEOUT,
engineio_logger=WEBSOCKET_SERVER_ENGINEIO_LOGGING,
)
# Timeout duration in seconds
TIMEOUT_DURATION = 3
SESSION_POOL_TIMEOUT = 120 # seconds without heartbeat before session is reaped
# Dictionary to maintain the user pool
if WEBSOCKET_MANAGER == 'redis':
log.debug('Using Redis to manage websockets.')
ws_sentinels = get_sentinels_from_env(WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT)
REDIS = get_redis_connection(
redis_url=WEBSOCKET_REDIS_URL,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
async_mode=True,
)
MODELS = RedisDict(
f'{REDIS_KEY_PREFIX}:models',
redis_url=WEBSOCKET_REDIS_URL,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
SESSION_POOL = RedisDict(
f'{REDIS_KEY_PREFIX}:session_pool',
redis_url=WEBSOCKET_REDIS_URL,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
USAGE_POOL = RedisDict(
f'{REDIS_KEY_PREFIX}:usage_pool',
redis_url=WEBSOCKET_REDIS_URL,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
clean_up_lock = RedisLock(
redis_url=WEBSOCKET_REDIS_URL,
lock_name=f'{REDIS_KEY_PREFIX}:usage_cleanup_lock',
timeout_secs=WEBSOCKET_REDIS_LOCK_TIMEOUT,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
aquire_func = clean_up_lock.aquire_lock
renew_func = clean_up_lock.renew_lock
release_func = clean_up_lock.release_lock
session_cleanup_lock = RedisLock(
redis_url=WEBSOCKET_REDIS_URL,
lock_name=f'{REDIS_KEY_PREFIX}:session_cleanup_lock',
timeout_secs=WEBSOCKET_REDIS_LOCK_TIMEOUT,
redis_sentinels=ws_sentinels,
redis_cluster=WEBSOCKET_REDIS_CLUSTER,
)
session_aquire_func = session_cleanup_lock.aquire_lock
session_renew_func = session_cleanup_lock.renew_lock
session_release_func = session_cleanup_lock.release_lock
else:
MODELS = {}
SESSION_POOL = {}
USAGE_POOL = {}
aquire_func = release_func = renew_func = lambda: True
session_aquire_func = session_release_func = session_renew_func = lambda: True
YDOC_MANAGER = YdocManager(
redis=REDIS,
redis_key_prefix=f'{REDIS_KEY_PREFIX}:ydoc:documents',
)
async def periodic_session_pool_cleanup():
"""Reap orphaned SESSION_POOL entries that missed heartbeats (e.g. crashed instance)."""
if not session_aquire_func():
log.debug('Session cleanup lock held by another node. Skipping.')
return
try:
while True:
if not session_renew_func():
log.error('Unable to renew session cleanup lock. Exiting.')
return
now = int(time.time())
for sid in list(SESSION_POOL.keys()):
entry = SESSION_POOL.get(sid)
if entry and now - entry.get('last_seen_at', 0) > SESSION_POOL_TIMEOUT:
log.warning(f'Reaping orphaned session {sid} (user {entry.get("id")})')
del SESSION_POOL[sid]
await asyncio.sleep(SESSION_POOL_TIMEOUT)
finally:
session_release_func()
async def periodic_usage_pool_cleanup():
max_retries = 2
retry_delay = random.uniform(WEBSOCKET_REDIS_LOCK_TIMEOUT / 2, WEBSOCKET_REDIS_LOCK_TIMEOUT)
for attempt in range(max_retries + 1):
if aquire_func():
break
else:
if attempt < max_retries:
log.debug(f'Cleanup lock already exists. Retry {attempt + 1} after {retry_delay}s...')
await asyncio.sleep(retry_delay)
else:
log.warning('Failed to acquire cleanup lock after retries. Skipping cleanup.')
return
log.debug('Running periodic_cleanup')
try:
while True:
if not renew_func():
log.error('Unable to renew cleanup lock. Exiting usage pool cleanup.')
raise Exception('Unable to renew usage pool cleanup lock.')
now = int(time.time())
for model_id, connections in list(USAGE_POOL.items()):
# Creating a list of sids to remove if they have timed out
expired_sids = [
sid for sid, details in connections.items() if now - details['updated_at'] > TIMEOUT_DURATION
]
if connections and not expired_sids:
continue
for sid in expired_sids:
del connections[sid]
if not connections:
log.debug(f'Cleaning up model {model_id} from usage pool')
del USAGE_POOL[model_id]
else:
USAGE_POOL[model_id] = connections
await asyncio.sleep(TIMEOUT_DURATION)
finally:
release_func()
app = socketio.ASGIApp(
sio,
socketio_path='/ws/socket.io',
)
def get_models_in_use():
# List models that are currently in use
models_in_use = list(USAGE_POOL.keys())
return models_in_use
def get_user_id_from_session_pool(sid):
user = SESSION_POOL.get(sid)
if user:
return user['id']
return None
def get_session_ids_from_room(room):
"""Get all session IDs from a specific room."""
members = get_room_sid_map(sio.manager, '/', room)
return list(members) if members else []
def get_user_ids_from_room(room):
active_session_ids = get_session_ids_from_room(room)
# Single pool lookup per session (each .get is a Redis round trip
# when the session pool is Redis-backed).
active_user_ids = list(
{
entry['id']
for entry in (SESSION_POOL.get(session_id) for session_id in active_session_ids)
if entry is not None
}
)
return active_user_ids
async def emit_to_users(event: str, data: dict, user_ids: list[str]):
"""
Send a message to specific users using their user:{id} rooms.
Args:
event (str): The event name to emit.
data (dict): The payload/data to send.
user_ids (list[str]): The target users' IDs.
"""
try:
for user_id in user_ids:
await sio.emit(event, data, room=f'user:{user_id}')
except Exception as e:
log.debug(f'Failed to emit event {event} to users {user_ids}: {e}')
async def enter_room_for_users(room: str, user_ids: list[str]):
"""
Make all sessions of a user join a specific room.
Args:
room (str): The room to join.
user_ids (list[str]): The target user's IDs.
"""
try:
for user_id in user_ids:
session_ids = get_session_ids_from_room(f'user:{user_id}')
for sid in session_ids:
await sio.enter_room(sid, room)
except Exception as e:
log.debug(f'Failed to make users {user_ids} join room {room}: {e}')
async def disconnect_user_sessions(user_id: str):
"""Disconnect all Socket.IO sessions belonging to a user.
Call this when a user's role is changed or the user is deleted so that
stale role/permission data cached in SESSION_POOL is invalidated.
The client will automatically reconnect and re-authenticate with
fresh data from the database.
"""
try:
session_ids = get_session_ids_from_room(f'user:{user_id}')
for sid in session_ids:
await sio.disconnect(sid)
if session_ids:
log.info(f'Disconnected {len(session_ids)} session(s) for user {user_id}')
except Exception as e:
log.warning(f'Failed to disconnect sessions for user {user_id}: {e}')
@sio.on('usage')
async def usage(sid, data):
if sid in SESSION_POOL:
model_id = data['model']
# Record the timestamp for the last update
current_time = int(time.time())
# Store the new usage data and task
USAGE_POOL[model_id] = {
**(USAGE_POOL.get(model_id) or {}),
sid: {'updated_at': current_time},
}
@sio.event
async def connect(sid, environ, auth):
user = None
if auth and 'token' in auth:
scope = (environ or {}).get('asgi.scope') or {}
fastapi_app = scope.get('app')
redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
user = await get_verified_user_by_token(auth['token'], redis)
if user:
SESSION_POOL[sid] = {
**user.model_dump(
exclude=[
'profile_image_url',
'profile_banner_image_url',
'date_of_birth',
'bio',
'gender',
]
),
'last_seen_at': int(time.time()),
}
await sio.enter_room(sid, f'user:{user.id}')
@sio.on('user-join')
async def user_join(sid, data):
auth = data.get('auth')
if not auth or 'token' not in auth:
return
environ = sio.get_environ(sid) or {}
scope = environ.get('asgi.scope') or {}
fastapi_app = scope.get('app')
redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
user = await get_verified_user_by_token(auth['token'], redis)
if not user:
return
existing = SESSION_POOL.get(sid)
if existing and existing.get('id') == user.id:
SESSION_POOL[sid] = {**existing, 'last_seen_at': int(time.time())}
else:
SESSION_POOL[sid] = {
**user.model_dump(
exclude=[
'profile_image_url',
'profile_banner_image_url',
'date_of_birth',
'bio',
'gender',
]
),
'last_seen_at': int(time.time()),
}
await sio.enter_room(sid, f'user:{user.id}')
# Join all the channels only if user has channels permission
if user.role == 'admin' or await has_permission(user.id, 'features.channels'):
channels = await Channels.get_channels_by_user_id(user.id)
log.debug(f'{channels=}')
for channel in channels:
await sio.enter_room(sid, f'channel:{channel.id}')
return {'id': user.id, 'name': user.name}
@sio.on('heartbeat')
async def heartbeat(sid, data):
user = SESSION_POOL.get(sid)
if user:
SESSION_POOL[sid] = {**user, 'last_seen_at': int(time.time())}
await Users.update_last_active_by_id(user['id'])
@sio.on('join-channels')
async def join_channel(sid, data):
auth = data['auth'] if 'auth' in data else None
if not auth or 'token' not in auth:
return
environ = sio.get_environ(sid) or {}
scope = environ.get('asgi.scope') or {}
fastapi_app = scope.get('app')
redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
user = await get_verified_user_by_token(auth['token'], redis)
if not user:
return
# Join all the channels only if user has channels permission
if user.role == 'admin' or await has_permission(user.id, 'features.channels'):
channels = await Channels.get_channels_by_user_id(user.id)
log.debug(f'{channels=}')
for channel in channels:
await sio.enter_room(sid, f'channel:{channel.id}')
@sio.on('join-note')
async def join_note(sid, data):
auth = data['auth'] if 'auth' in data else None
if not auth or 'token' not in auth:
return
environ = sio.get_environ(sid) or {}
scope = environ.get('asgi.scope') or {}
fastapi_app = scope.get('app')
redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
user = await get_verified_user_by_token(auth['token'], redis)
if not user:
return
note = await Notes.get_note_by_id(data['note_id'])
if not note:
log.error(f'Note {data["note_id"]} not found for user {user.id}')
return
if (
user.role != 'admin'
and user.id != note.user_id
and not await AccessGrants.has_access(
user_id=user.id,
resource_type='note',
resource_id=note.id,
permission='read',
)
):
log.error(f'User {user.id} does not have access to note {data["note_id"]}')
return
log.debug(f'Joining note {note.id} for user {user.id}')
await sio.enter_room(sid, f'note:{note.id}')
@sio.on('events:channel')
async def channel_events(sid, data):
room = f'channel:{data["channel_id"]}'
if sid not in (get_room_sid_map(sio.manager, '/', room) or {}):
return
event_data = data['data']
event_type = event_data['type']
user = SESSION_POOL.get(sid)
if not user:
return
if event_type == 'typing':
await sio.emit(
'events:channel',
{
'channel_id': data['channel_id'],
'message_id': data.get('message_id', None),
'data': event_data,
'user': UserNameResponse(**user).model_dump(),
},
room=room,
)
elif event_type == 'last_read_at':
await Channels.update_member_last_read_at(data['channel_id'], user['id'])
@sio.on('events:chat')
async def chat_events(sid, data):
user = SESSION_POOL.get(sid)
if not user:
return
event_data = data.get('data', {})
event_type = event_data.get('type')
if event_type == 'last_read_at':
if not await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id']):
return
try:
from open_webui.utils.timers import cancel_timers_for_chat
await cancel_timers_for_chat(data['chat_id'], 'chat.read', user['id'])
except Exception:
log.exception('Failed to cancel chat.read timers for chat %s', data.get('chat_id'))
def normalize_document_id(document_id: str) -> str:
"""Canonicalize document IDs to prevent auth bypass via prefix variants.
YdocManager normalizes storage keys by replacing ":" with "_", so
"note_abc" and "note:abc" resolve to the same underlying document.
We must rewrite underscore-prefixed IDs back to the colon form so
that authorization checks (which key on "note:") always fire.
"""
if document_id.startswith('note_'):
document_id = 'note:' + document_id[5:]
return document_id
@sio.on('ydoc:document:join')
async def ydoc_document_join(sid, data):
"""Handle user joining a document"""
user = SESSION_POOL.get(sid)
if not user:
return
try:
document_id = normalize_document_id(data['document_id'])
if document_id.startswith('note:'):
note_id = document_id.split(':')[1]
note = await Notes.get_note_by_id(note_id)
if not note:
log.error(f'Note {note_id} not found')
return
if (
user.get('role') != 'admin'
and user.get('id') != note.user_id
and not await AccessGrants.has_access(
user_id=user.get('id'),
resource_type='note',
resource_id=note.id,
permission='read',
)
):
log.error(f'User {user.get("id")} does not have access to note {note_id}')
return
user_id = data.get('user_id', sid)
user_name = data.get('user_name', 'Anonymous')
user_color = data.get('user_color', '#000000')
log.info(f'User {user_id} joining document {document_id}')
await YDOC_MANAGER.add_user(document_id=document_id, user_id=sid)
# Join Socket.IO room
await sio.enter_room(sid, f'doc_{document_id}')
active_session_ids = get_session_ids_from_room(f'doc_{document_id}')
# Get the Yjs document state
ydoc = Y.Doc()
updates = await YDOC_MANAGER.get_updates(document_id)
for update in updates:
ydoc.apply_update(bytes(update))
# Encode the entire document state as an update
state_update = ydoc.get_update()
await sio.emit(
'ydoc:document:state',
{
'document_id': document_id,
'state': list(state_update), # Convert bytes to list for JSON
'sessions': active_session_ids,
},
room=sid,
)
# Notify other users about the new user
await sio.emit(
'ydoc:user:joined',
{
'document_id': document_id,
'user_id': user_id,
'user_name': user_name,
'user_color': user_color,
},
room=f'doc_{document_id}',
skip_sid=sid,
)
log.info(f'User {user_id} successfully joined document {document_id}')
except Exception as e:
log.error(f'Error in yjs_document_join: {e}')
await sio.emit('error', {'message': 'Failed to join document'}, room=sid)
async def document_save_handler(document_id, data, user):
document_id = normalize_document_id(document_id)
if document_id.startswith('note:'):
note_id = document_id.split(':')[1]
note = await Notes.get_note_by_id(note_id)
if not note:
log.error(f'Note {note_id} not found')
return
if (
user.get('role') != 'admin'
and user.get('id') != note.user_id
and not await AccessGrants.has_access(
user_id=user.get('id'),
resource_type='note',
resource_id=note.id,
permission='write',
)
):
log.error(f'User {user.get("id")} does not have write access to note {note_id}')
return
await Notes.update_note_by_id(note_id, NoteUpdateForm(data=data))
@sio.on('ydoc:document:state')
async def yjs_document_state(sid, data):
"""Send the current state of the Yjs document to the user"""
try:
document_id = data['document_id']
document_id = normalize_document_id(document_id)
room = f'doc_{document_id}'
active_session_ids = get_session_ids_from_room(room)
if sid not in active_session_ids:
log.warning(f'Session {sid} not in room {room}. Cannot send state.')
return
if not await YDOC_MANAGER.document_exists(document_id):
log.warning(f'Document {document_id} not found')
return
# Get the Yjs document state
ydoc = Y.Doc()
updates = await YDOC_MANAGER.get_updates(document_id)
for update in updates:
ydoc.apply_update(bytes(update))
# Encode the entire document state as an update
state_update = ydoc.get_update()
await sio.emit(
'ydoc:document:state',
{
'document_id': document_id,
'state': list(state_update), # Convert bytes to list for JSON
'sessions': active_session_ids,
},
room=sid,
)
except Exception as e:
log.error(f'Error in yjs_document_state: {e}')
@sio.on('ydoc:document:update')
async def yjs_document_update(sid, data):
"""Handle Yjs document updates"""
try:
document_id = data['document_id']
document_id = normalize_document_id(document_id)
# Verify the sender actually joined this document room
room = f'doc_{document_id}'
active_session_ids = get_session_ids_from_room(room)
if sid not in active_session_ids:
log.warning(f'Session {sid} not in room {room}. Rejecting update.')
return
# Verify write permission — room membership only proves read access
user = SESSION_POOL.get(sid)
if not user:
return
if document_id.startswith('note:'):
note_id = document_id.split(':')[1]
note = await Notes.get_note_by_id(note_id)
if not note:
log.error(f'Note {note_id} not found')
return
if (
user.get('role') != 'admin'
and user.get('id') != note.user_id
and not await AccessGrants.has_access(
user_id=user.get('id'),
resource_type='note',
resource_id=note.id,
permission='write',
)
):
log.warning(f'User {user.get("id")} does not have write access to note {note_id}. Rejecting update.')
return
try:
await stop_item_tasks(REDIS, document_id)
except Exception:
pass
user_id = data.get('user_id', sid)
update = data['update'] # List of bytes from frontend
await YDOC_MANAGER.append_to_updates(
document_id=document_id,
update=update, # Convert list of bytes to bytes
)
# Broadcast update to all other users in the document
await sio.emit(
'ydoc:document:update',
{
'document_id': document_id,
'user_id': user_id,
'update': update,
'socket_id': sid, # Add socket_id to match frontend filtering
},
room=f'doc_{document_id}',
skip_sid=sid,
)
async def debounced_save():
await asyncio.sleep(0.5)
await document_save_handler(document_id, data.get('data', {}), user)
if data.get('data'):
await create_task(REDIS, debounced_save(), document_id)
except Exception as e:
log.error(f'Error in yjs_document_update: {e}')
@sio.on('ydoc:document:leave')
async def yjs_document_leave(sid, data):
"""Handle user leaving a document"""
user = SESSION_POOL.get(sid)
if not user: # authenticated session required (parity with sibling handlers)
return
try:
document_id = normalize_document_id(data['document_id'])
log.info(f'User {user["id"]} leaving document {document_id}')
# Remove user from the document
await YDOC_MANAGER.remove_user(document_id=document_id, user_id=sid)
# Leave Socket.IO room
await sio.leave_room(sid, f'doc_{document_id}')
# Notify other users; user_id is the authenticated identity, not client-supplied
await sio.emit(
'ydoc:user:left',
{'document_id': document_id, 'user_id': user['id']},
room=f'doc_{document_id}',
)
if await YDOC_MANAGER.document_exists(document_id) and len(await YDOC_MANAGER.get_users(document_id)) == 0:
log.info(f'Cleaning up document {document_id} as no users are left')
await YDOC_MANAGER.clear_document(document_id)
except Exception as e:
log.error(f'Error in yjs_document_leave: {e}')
@sio.on('ydoc:awareness:update')
async def yjs_awareness_update(sid, data):
"""Handle awareness updates (cursors, selections, etc.)"""
user = SESSION_POOL.get(sid)
if not user: # authenticated session required (parity with sibling handlers)
return
try:
document_id = normalize_document_id(data['document_id'])
room = f'doc_{document_id}'
if room not in sio.rooms(sid): # must have joined the document first
return
update = data['update']
# Broadcast to the room; user_id is the authenticated identity, not client-supplied
await sio.emit(
'ydoc:awareness:update',
{'document_id': document_id, 'user_id': user['id'], 'update': update},
room=room,
skip_sid=sid,
)
except Exception as e:
log.error(f'Error in yjs_awareness_update: {e}')
@sio.event
async def disconnect(sid, reason=None):
if sid in SESSION_POOL:
del SESSION_POOL[sid]
# Clean up USAGE_POOL entries for this session
for model_id in list(USAGE_POOL.keys()):
connections = USAGE_POOL.get(model_id)
if connections and sid in connections:
del connections[sid]
if not connections:
del USAGE_POOL[model_id]
else:
USAGE_POOL[model_id] = connections
await YDOC_MANAGER.remove_user_from_all_documents(sid)
else:
pass
# print(f"Unknown session ID {sid} disconnected")
async def _make_channel_emitter(request_info):
"""Event emitter that routes pipeline output to a channel message.
Translates chat:completion events into channel message:update socket
emissions, throttled to avoid flooding with per-token updates.
"""
channel_id = request_info['chat_id'].removeprefix('channel:')
message_id = request_info['message_id']
state = {'last_emit_at': 0.0}
THROTTLE_INTERVAL = 0.15 # ~6 updates/sec
async def _emit_channel_update(content: str, done: bool = False):
from open_webui.models.messages import MessageForm, Messages
msg = await Messages.get_message_by_id(message_id)
if not msg or msg.channel_id != channel_id:
return
update_form = MessageForm(content=content)
if done:
# Merge done flag into existing meta (preserve model_id etc.)
existing_meta = msg.meta or {}
update_form = MessageForm(
content=content,
meta={**existing_meta, 'done': True},
)
await Messages.update_message_by_id(message_id, update_form)
message = await Messages.get_message_by_id(message_id)
if message:
await sio.emit(
'events:channel',
{
'channel_id': channel_id,
'message_id': message_id,
'data': {
'type': 'message:update',
'data': message.model_dump(),
},
},
to=f'channel:{channel_id}',
)
async def __channel_emitter__(event_data):
event_type = event_data.get('type')
if event_type == 'chat:completion':
data = event_data.get('data', {})
content = data.get('content', '')
done = data.get('done', False)
if not content and not done:
return
now = __import__('time').time()
if done or (now - state['last_emit_at']) >= THROTTLE_INTERVAL:
state['last_emit_at'] = now
await _emit_channel_update(content, done)
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)
await _emit_channel_update(f'Error: {error_content}', done=True)
return __channel_emitter__
async def get_event_emitter(request_info, update_db=True):
# Channel mode: route pipeline output to channel message updates
if (request_info.get('chat_id') or '').startswith('channel:'):
return await _make_channel_emitter(request_info)
async def __event_emitter__(event_data):
user_id = request_info['user_id']
chat_id = request_info['chat_id']
message_id = request_info['message_id']
internal = request_info.get('internal') is True
if internal and event_data.get('type') == 'notification':
return
room = f'user:{user_id}'
# Local rooms are authoritative; Redis may have listeners on another instance.
if WEBSOCKET_MANAGER == 'redis' or room in sio.manager.rooms.get('/', {}):
await sio.emit(
'events',
{
'chat_id': chat_id,
'message_id': message_id,
**({'internal': True} if internal else {}),
'data': event_data,
},
room=room,
)
if update_db and message_id and not (request_info.get('chat_id') or '').startswith('local:'):
event_type = event_data.get('type')
if event_type == 'status':
await Chats.add_message_status_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
event_data.get('data', {}),
)
elif event_type == 'message':
message = await Chats.get_message_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
)
if message:
content = message.get('content', '')
content += event_data.get('data', {}).get('content', '')
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
'content': content,
},
)
elif event_type == 'replace':
content = event_data.get('data', {}).get('content', '')
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
'content': content,
},
)
elif event_type == 'embeds':
event_payload = event_data.get('data', {})
embeds = event_payload.get('embeds', [])
if not event_payload.get('replace', False):
message = await Chats.get_message_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
)
embeds.extend(message.get('embeds', []))
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
'embeds': embeds,
},
touch=False,
)
elif event_type == 'files':
message = await Chats.get_message_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
)
files = event_data.get('data', {}).get('files', [])
files.extend(message.get('files', []))
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
'files': files,
},
touch=False,
)
elif event_type in ('source', 'citation'):
data = event_data.get('data', {})
if data.get('type') is None:
message = await Chats.get_message_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
)
sources = message.get('sources', [])
sources.append(data)
await Chats.upsert_message_to_chat_by_id_and_message_id(
request_info['chat_id'],
request_info['message_id'],
{
'sources': sources,
},
touch=False,
)
if 'user_id' in request_info and 'chat_id' in request_info and 'message_id' in request_info:
return __event_emitter__
else:
return None
async def get_event_call(request_info):
async def __event_caller__(event_data):
session_id = request_info['session_id']
# session_id is client-supplied; only the requesting user's own live session may be targeted.
session = SESSION_POOL.get(session_id)
if session is None or session.get('id') != request_info.get('user_id'):
log.warning(f'Event caller: session {session_id} not owned by requesting user or disconnected')
return {'error': 'Client session disconnected.'}
try:
return await sio.call(
'events',
{
'chat_id': request_info.get('chat_id', None),
'message_id': request_info.get('message_id', None),
'data': event_data,
},
to=session_id,
timeout=WEBSOCKET_EVENT_CALLER_TIMEOUT,
)
except TimeoutError:
log.warning(f'Event caller timed out for session {session_id}')
return {'error': 'Event call timed out. The browser tab may be inactive or closed.'}
if 'session_id' in request_info and 'chat_id' in request_info and 'message_id' in request_info:
return __event_caller__
else:
return None
get_event_caller = get_event_call