Coverage for open_webui/socket/main.py: 22%
656 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1from __future__ import annotations
3import asyncio
4import copy
5import logging
6import random
7import sys
8import time
9from contextlib import suppress
10from typing import Any
12import pycrdt as Y
13import socketio
14from open_webui.config import (
15 CORS_ALLOW_ORIGIN,
16)
17from open_webui.env import (
18 ENABLE_WEBSOCKET_SUPPORT,
19 GLOBAL_LOG_LEVEL,
20 REDIS_KEY_PREFIX,
21 WEBSOCKET_EVENT_CALLER_TIMEOUT,
22 WEBSOCKET_HEARTBEAT_INTERVAL,
23 WEBSOCKET_MANAGER,
24 WEBSOCKET_REDIS_CLUSTER,
25 WEBSOCKET_REDIS_LOCK_TIMEOUT,
26 WEBSOCKET_REDIS_OPTIONS,
27 WEBSOCKET_REDIS_URL,
28 WEBSOCKET_SENTINEL_HOSTS,
29 WEBSOCKET_SENTINEL_PORT,
30 WEBSOCKET_SERVER_ENGINEIO_LOGGING,
31 WEBSOCKET_SERVER_LOGGING,
32 WEBSOCKET_SERVER_PING_INTERVAL,
33 WEBSOCKET_SERVER_PING_TIMEOUT,
34)
35from open_webui.models.access_grants import AccessGrants
36from open_webui.models.channels import Channels
37from open_webui.models.chats import Chats
38from open_webui.models.folders import Folders
39from open_webui.models.notes import Notes, NoteUpdateForm
40from open_webui.models.users import UserNameResponse, Users
41from open_webui.socket.utils import CachedRedisDict, RedisDict, RedisLock, YdocManager
42from open_webui.tasks import (
43 REDIS_PUBSUB_MAX_RECONNECT_INTERVAL,
44 REDIS_PUBSUB_RECONNECT_INTERVAL,
45 create_task,
46 stop_item_tasks,
47)
48from open_webui.utils.access_control import has_permission
49from open_webui.utils.auth import get_verified_user_by_token
50from open_webui.utils.chat_id import is_saved_chat_id
51from open_webui.utils.json_codec import SOCKETIO_JSON, JSONCodec, dumps_bytes
52from open_webui.utils.misc import get_output_text
53from open_webui.utils.redis import (
54 build_sentinel_url,
55 get_redis_connection,
56 get_sentinels_from_env,
57)
58from redis.exceptions import RedisError
59from socketio.packet import Packet
61logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
62log = logging.getLogger(__name__)
65# Let no connection opened in good faith be dropped without
66# cause, and let every message find the room it was meant for.
67REDIS = None
69# Configure CORS for Socket.IO
70SOCKETIO_CORS_ORIGINS = '*' if CORS_ALLOW_ORIGIN == ['*'] else CORS_ALLOW_ORIGIN
73def get_room_sid_map(manager, namespace: str, room: str):
74 """Return this process's Socket.IO sid map for a room, without copying it."""
75 return manager.rooms.get(namespace, {}).get(room)
78class JSONOnlyPacket(Packet):
79 """Packet class for JSON-serializable payloads only, skipping python-socketio's per-emit binary scan."""
81 uses_binary_events = False
83 @classmethod
84 def reconstruct_binary(cls, data: Any, attachments: list[bytes]):
85 """Normalize client attachments to int lists, the form the Yjs handlers store and apply."""
86 return super().reconstruct_binary(data, [list(attachment) for attachment in attachments])
89if WEBSOCKET_MANAGER == 'redis': 89 ↛ 90line 89 didn't jump to line 90 because the condition on line 89 was never true
90 sentinel_hosts = WEBSOCKET_SENTINEL_HOSTS or ''
91 ws_redis_url = (
92 build_sentinel_url(WEBSOCKET_REDIS_URL, sentinel_hosts, WEBSOCKET_SENTINEL_PORT)
93 if sentinel_hosts
94 else WEBSOCKET_REDIS_URL
95 )
96 redis_manager = socketio.AsyncRedisManager(ws_redis_url, redis_options=WEBSOCKET_REDIS_OPTIONS, json=SOCKETIO_JSON)
97 sio = socketio.AsyncServer(
98 cors_allowed_origins=SOCKETIO_CORS_ORIGINS,
99 async_mode='asgi',
100 json=SOCKETIO_JSON,
101 serializer=JSONOnlyPacket,
102 transports=(['websocket'] if ENABLE_WEBSOCKET_SUPPORT else ['polling']),
103 allow_upgrades=ENABLE_WEBSOCKET_SUPPORT,
104 always_connect=True,
105 client_manager=redis_manager,
106 logger=WEBSOCKET_SERVER_LOGGING,
107 ping_interval=WEBSOCKET_SERVER_PING_INTERVAL,
108 ping_timeout=WEBSOCKET_SERVER_PING_TIMEOUT,
109 engineio_logger=WEBSOCKET_SERVER_ENGINEIO_LOGGING,
110 )
111else:
112 sio = socketio.AsyncServer(
113 cors_allowed_origins=SOCKETIO_CORS_ORIGINS,
114 async_mode='asgi',
115 json=SOCKETIO_JSON,
116 serializer=JSONOnlyPacket,
117 transports=(['websocket'] if ENABLE_WEBSOCKET_SUPPORT else ['polling']),
118 allow_upgrades=ENABLE_WEBSOCKET_SUPPORT,
119 always_connect=True,
120 logger=WEBSOCKET_SERVER_LOGGING,
121 ping_interval=WEBSOCKET_SERVER_PING_INTERVAL,
122 ping_timeout=WEBSOCKET_SERVER_PING_TIMEOUT,
123 engineio_logger=WEBSOCKET_SERVER_ENGINEIO_LOGGING,
124 )
127# Timeout duration in seconds
128TIMEOUT_DURATION = 3
129SESSION_POOL_TIMEOUT = max(WEBSOCKET_HEARTBEAT_INTERVAL * 4, 120) if WEBSOCKET_HEARTBEAT_INTERVAL is not None else 120
131# Dictionary to maintain the user pool
133if WEBSOCKET_MANAGER == 'redis': 133 ↛ 134line 133 didn't jump to line 134 because the condition on line 133 was never true
134 log.debug('Using Redis to manage websockets.')
135 ws_sentinels = get_sentinels_from_env(WEBSOCKET_SENTINEL_HOSTS, WEBSOCKET_SENTINEL_PORT)
136 REDIS = get_redis_connection(
137 redis_url=WEBSOCKET_REDIS_URL,
138 redis_sentinels=ws_sentinels,
139 redis_cluster=WEBSOCKET_REDIS_CLUSTER,
140 async_mode=True,
141 )
143 MODELS = CachedRedisDict(
144 f'{REDIS_KEY_PREFIX}:models',
145 redis_url=WEBSOCKET_REDIS_URL,
146 redis_sentinels=ws_sentinels,
147 redis_cluster=WEBSOCKET_REDIS_CLUSTER,
148 )
150 SESSION_POOL = RedisDict(
151 f'{REDIS_KEY_PREFIX}:session_pool',
152 redis_url=WEBSOCKET_REDIS_URL,
153 redis_sentinels=ws_sentinels,
154 redis_cluster=WEBSOCKET_REDIS_CLUSTER,
155 )
156 USAGE_POOL = RedisDict(
157 f'{REDIS_KEY_PREFIX}:usage_pool',
158 redis_url=WEBSOCKET_REDIS_URL,
159 redis_sentinels=ws_sentinels,
160 redis_cluster=WEBSOCKET_REDIS_CLUSTER,
161 )
163 clean_up_lock = RedisLock(
164 redis_url=WEBSOCKET_REDIS_URL,
165 lock_name=f'{REDIS_KEY_PREFIX}:usage_cleanup_lock',
166 timeout_secs=WEBSOCKET_REDIS_LOCK_TIMEOUT,
167 redis_sentinels=ws_sentinels,
168 redis_cluster=WEBSOCKET_REDIS_CLUSTER,
169 )
170 aquire_func = clean_up_lock.aquire_lock
171 renew_func = clean_up_lock.renew_lock
172 release_func = clean_up_lock.release_lock
174 session_cleanup_lock = RedisLock(
175 redis_url=WEBSOCKET_REDIS_URL,
176 lock_name=f'{REDIS_KEY_PREFIX}:session_cleanup_lock',
177 timeout_secs=WEBSOCKET_REDIS_LOCK_TIMEOUT,
178 redis_sentinels=ws_sentinels,
179 redis_cluster=WEBSOCKET_REDIS_CLUSTER,
180 )
181 session_aquire_func = session_cleanup_lock.aquire_lock
182 session_renew_func = session_cleanup_lock.renew_lock
183 session_release_func = session_cleanup_lock.release_lock
184else:
185 MODELS = {}
187 SESSION_POOL = {}
188 USAGE_POOL = {}
190 aquire_func = release_func = renew_func = lambda: True
191 session_aquire_func = session_release_func = session_renew_func = lambda: True
194YDOC_MANAGER = YdocManager(
195 redis=REDIS,
196 redis_key_prefix=f'{REDIS_KEY_PREFIX}:ydoc:documents',
197)
199REDIS_EVENT_CHANNEL = f'{REDIS_KEY_PREFIX}:direct_completion'
201EVENT_QUEUES: dict[str, asyncio.Queue] = {}
202EVENT_PUBLISH_LOCK = asyncio.Lock()
205def get_session_pool_batches():
206 """All session pool entries, in bounded batches for the Redis backing."""
207 if WEBSOCKET_MANAGER == 'redis': 207 ↛ 208line 207 didn't jump to line 208 because the condition on line 207 was never true
208 return SESSION_POOL.scan_batches()
209 return [list(SESSION_POOL.items())]
212async def periodic_session_pool_cleanup():
213 """Reap orphaned SESSION_POOL entries that missed heartbeats (e.g. crashed instance)."""
214 retry_delay = random.uniform(WEBSOCKET_REDIS_LOCK_TIMEOUT / 2, WEBSOCKET_REDIS_LOCK_TIMEOUT)
215 renew_interval = max(WEBSOCKET_REDIS_LOCK_TIMEOUT / 2, 0.5)
216 while True:
217 try:
218 if not session_aquire_func(): 218 ↛ 219line 218 didn't jump to line 219 because the condition on line 218 was never true
219 log.debug('Session cleanup lock held by another node. Retrying.')
220 await asyncio.sleep(retry_delay)
221 continue
223 try:
224 while True:
225 if not session_renew_func(): 225 ↛ 226line 225 didn't jump to line 226 because the condition on line 225 was never true
226 log.warning('Unable to renew session cleanup lock. Retrying cleanup ownership.')
227 break
229 now = int(time.time())
230 for batch in get_session_pool_batches():
231 expired = [
232 sid
233 for sid, entry in batch
234 if entry and now - entry.get('last_seen_at', 0) > SESSION_POOL_TIMEOUT
235 ]
236 if expired: 236 ↛ 237line 236 didn't jump to line 237 because the condition on line 236 was never true
237 log.warning('Reaping %d orphaned session(s) from the session pool', len(expired))
238 if WEBSOCKET_MANAGER == 'redis':
239 SESSION_POOL.delete_many(*expired)
240 else:
241 for sid in expired:
242 SESSION_POOL.pop(sid, None)
243 await asyncio.sleep(0) # don't hold the loop for the whole sweep
245 next_cleanup_at = time.monotonic() + SESSION_POOL_TIMEOUT
246 lock_lost = False
247 while True:
248 sleep_for = min(renew_interval, next_cleanup_at - time.monotonic())
249 if sleep_for <= 0:
250 break
251 await asyncio.sleep(sleep_for)
252 if not session_renew_func(): 252 ↛ 253line 252 didn't jump to line 253 because the condition on line 252 was never true
253 log.warning('Unable to renew session cleanup lock. Retrying cleanup ownership.')
254 lock_lost = True
255 break
257 if lock_lost: 257 ↛ 258line 257 didn't jump to line 258 because the condition on line 257 was never true
258 break
259 finally:
260 session_release_func()
261 except Exception:
262 log.exception('Session pool cleanup failed. Retrying.')
263 await asyncio.sleep(retry_delay)
266async def periodic_usage_pool_cleanup():
267 retry_delay = random.uniform(WEBSOCKET_REDIS_LOCK_TIMEOUT / 2, WEBSOCKET_REDIS_LOCK_TIMEOUT)
268 while True:
269 try:
270 if not aquire_func(): 270 ↛ 271line 270 didn't jump to line 271 because the condition on line 270 was never true
271 log.debug('Usage cleanup lock held by another node. Retrying.')
272 await asyncio.sleep(retry_delay)
273 continue
275 try:
276 while True:
277 if not renew_func(): 277 ↛ 278line 277 didn't jump to line 278 because the condition on line 277 was never true
278 log.warning('Unable to renew usage cleanup lock. Retrying cleanup ownership.')
279 break
281 now = int(time.time())
282 for model_id, connections in list(USAGE_POOL.items()): 282 ↛ 283line 282 didn't jump to line 283 because the loop on line 282 never started
283 expired_sids = [
284 sid
285 for sid, details in connections.items()
286 if now - details['updated_at'] > TIMEOUT_DURATION
287 ]
289 if connections and not expired_sids:
290 continue
292 for sid in expired_sids:
293 del connections[sid]
295 if not connections:
296 log.debug('Cleaning up model %s from usage pool', model_id)
297 try:
298 del USAGE_POOL[model_id]
299 except KeyError:
300 pass
301 else:
302 USAGE_POOL[model_id] = connections
303 await asyncio.sleep(TIMEOUT_DURATION)
304 finally:
305 release_func()
306 except Exception:
307 log.exception('Usage pool cleanup failed. Retrying.')
308 await asyncio.sleep(retry_delay)
311app = socketio.ASGIApp(
312 sio,
313 socketio_path='/ws/socket.io',
314)
317def get_models_in_use():
318 # List models that are currently in use
319 models_in_use = list(USAGE_POOL.keys())
320 return models_in_use
323def get_user_id_from_session_pool(sid):
324 user = SESSION_POOL.get(sid)
325 if user:
326 return user['id']
327 return None
330async def get_socket_session_user(sid: str) -> dict | None:
331 """Session user from this worker's local Socket.IO store; only locally connected sids are ever looked up."""
332 try:
333 return (await sio.get_session(sid)).get('user')
334 except KeyError:
335 return None
338def get_session_ids_from_room(room):
339 """Get all session IDs from a specific room."""
340 members = get_room_sid_map(sio.manager, '/', room)
341 return list(members) if members else []
344def get_session_ids_by_user_id(user_id: str) -> list[str]:
345 """Get known session IDs for a user across the local rooms and shared session pool."""
346 session_ids = set(get_session_ids_from_room(f'user:{user_id}'))
347 session_ids.update(sid for sid, entry in SESSION_POOL.items() if entry and entry.get('id') == user_id)
348 return list(session_ids)
351async def get_user_ids_from_room(room) -> set[str]:
352 users = [await get_socket_session_user(session_id) for session_id in get_session_ids_from_room(room)]
353 return {user['id'] for user in users if user}
356async def emit_to_users(event: str, data: dict, user_ids: list[str]):
357 """
358 Send a message to specific users using their user:{id} rooms.
360 Args:
361 event (str): The event name to emit.
362 data (dict): The payload/data to send.
363 user_ids (list[str]): The target users' IDs.
364 """
365 try:
366 for user_id in user_ids:
367 await sio.emit(event, data, room=f'user:{user_id}')
368 except Exception as e:
369 log.debug('Failed to emit event %s to users %s: %s', event, user_ids, e)
372async def enter_room_for_users(room: str, user_ids: list[str]):
373 """
374 Make all sessions of a user join a specific room.
375 Args:
376 room (str): The room to join.
377 user_ids (list[str]): The target user's IDs.
378 """
379 try:
380 for user_id in user_ids:
381 session_ids = get_session_ids_from_room(f'user:{user_id}')
382 for sid in session_ids:
383 await sio.enter_room(sid, room)
384 except Exception as e:
385 log.debug('Failed to make users %s join room %s: %s', user_ids, room, e)
388async def disconnect_user_sessions(user_id: str):
389 """Disconnect all Socket.IO sessions belonging to a user.
391 Call this when a user's role is changed or the user is deleted so that
392 stale role/permission data cached in SESSION_POOL is invalidated.
393 The client will automatically reconnect and re-authenticate with
394 fresh data from the database.
395 """
396 session_ids = get_session_ids_by_user_id(user_id)
397 for sid in session_ids: 397 ↛ 398line 397 didn't jump to line 398 because the loop on line 397 never started
398 try:
399 await sio.disconnect(sid)
400 except Exception:
401 log.exception('Failed to disconnect session %s for user %s', sid, user_id)
403 if session_ids: 403 ↛ 404line 403 didn't jump to line 404 because the condition on line 403 was never true
404 log.info('Requested disconnect of %s session(s) for user %s', len(session_ids), user_id)
407@sio.on('usage')
408async def usage(sid, data):
409 if await get_socket_session_user(sid):
410 model_id = data['model']
411 # Record the timestamp for the last update
412 current_time = int(time.time())
414 # Store the new usage data and task
415 USAGE_POOL[model_id] = {
416 **(USAGE_POOL.get(model_id) or {}),
417 sid: {'updated_at': current_time},
418 }
421@sio.event
422async def connect(sid, environ, auth):
423 user = None
424 if auth and 'token' in auth:
425 scope = (environ or {}).get('asgi.scope') or {}
426 fastapi_app = scope.get('app')
427 redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
428 user = await get_verified_user_by_token(auth['token'], redis)
430 if user:
431 socket_user = {
432 **user.model_dump(
433 exclude=[
434 'profile_image_url',
435 'profile_banner_image_url',
436 'date_of_birth',
437 'bio',
438 'gender',
439 ]
440 ),
441 'last_seen_at': int(time.time()),
442 }
443 SESSION_POOL[sid] = socket_user
444 await sio.save_session(sid, {'user': socket_user})
445 await sio.enter_room(sid, f'user:{user.id}')
448@sio.on('user-join')
449async def user_join(sid, data):
450 auth = data.get('auth')
451 if not auth or 'token' not in auth:
452 return
454 environ = sio.get_environ(sid) or {}
455 scope = environ.get('asgi.scope') or {}
456 fastapi_app = scope.get('app')
457 redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
458 user = await get_verified_user_by_token(auth['token'], redis)
459 if not user:
460 return
462 socket_user = {
463 **user.model_dump(
464 exclude=[
465 'profile_image_url',
466 'profile_banner_image_url',
467 'date_of_birth',
468 'bio',
469 'gender',
470 ]
471 ),
472 'last_seen_at': int(time.time()),
473 }
475 SESSION_POOL[sid] = socket_user
476 await sio.save_session(sid, {'user': socket_user})
477 await sio.enter_room(sid, f'user:{user.id}')
479 # Join all the channels only if user has channels permission
480 if user.role == 'admin' or await has_permission(user.id, 'features.channels'):
481 channels = await Channels.get_channels_by_user_id(user.id)
482 log.debug('channels=%r', channels)
483 for channel in channels:
484 await sio.enter_room(sid, f'channel:{channel.id}')
486 return {'id': user.id, 'name': user.name}
489@sio.on('heartbeat')
490async def heartbeat(sid, data):
491 user = await get_socket_session_user(sid)
492 if user:
493 SESSION_POOL[sid] = {**user, 'last_seen_at': int(time.time())}
494 await Users.update_last_active_by_id(user['id'])
497@sio.on('join-channels')
498async def join_channel(sid, data):
499 auth = data['auth'] if 'auth' in data else None
500 if not auth or 'token' not in auth:
501 return
503 environ = sio.get_environ(sid) or {}
504 scope = environ.get('asgi.scope') or {}
505 fastapi_app = scope.get('app')
506 redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
507 user = await get_verified_user_by_token(auth['token'], redis)
508 if not user:
509 return
511 # Join all the channels only if user has channels permission
512 if user.role == 'admin' or await has_permission(user.id, 'features.channels'):
513 channels = await Channels.get_channels_by_user_id(user.id)
514 log.debug('channels=%r', channels)
515 for channel in channels:
516 await sio.enter_room(sid, f'channel:{channel.id}')
519@sio.on('join-note')
520async def join_note(sid, data):
521 auth = data['auth'] if 'auth' in data else None
522 if not auth or 'token' not in auth:
523 return
525 environ = sio.get_environ(sid) or {}
526 scope = environ.get('asgi.scope') or {}
527 fastapi_app = scope.get('app')
528 redis = getattr(getattr(fastapi_app, 'state', None), 'redis', None) or REDIS
529 user = await get_verified_user_by_token(auth['token'], redis)
530 if not user:
531 return
533 note = await Notes.get_note_by_id(data['note_id'])
534 if not note:
535 log.error(f'Note {data["note_id"]} not found for user {user.id}')
536 return
538 if (
539 user.role != 'admin'
540 and user.id != note.user_id
541 and not await AccessGrants.has_access(
542 user_id=user.id,
543 resource_type='note',
544 resource_id=note.id,
545 permission='read',
546 )
547 ):
548 log.error(f'User {user.id} does not have access to note {data["note_id"]}')
549 return
551 log.debug('Joining note %s for user %s', note.id, user.id)
552 await sio.enter_room(sid, f'note:{note.id}')
555@sio.on('events:channel')
556async def channel_events(sid, data):
557 room = f'channel:{data["channel_id"]}'
558 if sid not in (get_room_sid_map(sio.manager, '/', room) or {}):
559 return
561 event_data = data['data']
562 event_type = event_data['type']
564 user = await get_socket_session_user(sid)
566 if not user:
567 return
569 if event_type == 'typing':
570 await sio.emit(
571 'events:channel',
572 {
573 'channel_id': data['channel_id'],
574 'message_id': data.get('message_id', None),
575 'data': event_data,
576 'user': UserNameResponse(**user).model_dump(),
577 },
578 room=room,
579 )
580 elif event_type == 'last_read_at':
581 await Channels.update_member_last_read_at(data['channel_id'], user['id'])
584async def get_folder_unread_counts(user_id: str) -> dict[str, int]:
585 folder_list = await Folders.get_folders_by_user_id(user_id)
586 parent_by_id = {folder.id: folder.parent_id for folder in folder_list}
587 unread_counts = dict.fromkeys(parent_by_id.keys(), 0)
589 direct_unread_counts = await Chats.count_unread_by_folder_ids(user_id, list(parent_by_id.keys()))
590 for unread_folder_id, unread_count in direct_unread_counts.items():
591 current_id = unread_folder_id
592 seen = set()
593 while current_id and current_id not in seen:
594 seen.add(current_id)
595 if current_id in unread_counts:
596 unread_counts[current_id] += unread_count
597 current_id = parent_by_id.get(current_id)
599 return unread_counts
602@sio.on('events:chat')
603async def chat_events(sid, data):
604 user = await get_socket_session_user(sid)
605 if not user:
606 return
608 event_data = data.get('data', {})
609 event_type = event_data.get('type')
611 if event_type == 'last_read_at':
612 read_update = await Chats.update_chat_last_read_at_by_id(data['chat_id'], user['id'])
613 if not read_update:
614 return
615 last_read_at, was_unread = read_update
616 response_data = {
617 'chat_id': data['chat_id'],
618 'last_read_at': last_read_at,
619 }
620 if was_unread:
621 response_data['folder_unread_counts'] = await get_folder_unread_counts(user['id'])
623 await sio.emit(
624 'events',
625 {
626 'chat_id': data['chat_id'],
627 'data': {
628 'type': 'chat:list',
629 'data': response_data,
630 },
631 },
632 room=f'user:{user["id"]}',
633 )
634 try:
635 from open_webui.utils.timers import cancel_timers_for_chat
637 await cancel_timers_for_chat(data['chat_id'], 'chat.read', user['id'])
638 except Exception:
639 log.exception('Failed to cancel chat.read timers for chat %s', data.get('chat_id'))
642def normalize_document_id(document_id: str) -> str:
643 """Canonicalize document IDs to prevent auth bypass via prefix variants.
645 YdocManager normalizes storage keys by replacing ":" with "_", so
646 "note_abc" and "note:abc" resolve to the same underlying document.
647 We must rewrite underscore-prefixed IDs back to the colon form so
648 that authorization checks (which key on "note:") always fire.
649 """
650 if document_id.startswith('note_'):
651 document_id = 'note:' + document_id[5:]
652 return document_id
655@sio.on('ydoc:document:join')
656async def ydoc_document_join(sid, data):
657 """Handle user joining a document"""
658 user = await get_socket_session_user(sid)
659 if not user:
660 return
662 try:
663 document_id = normalize_document_id(data['document_id'])
665 if document_id.startswith('note:'):
666 note_id = document_id.split(':')[1]
667 note = await Notes.get_note_by_id(note_id)
668 if not note:
669 log.error(f'Note {note_id} not found')
670 return
672 if (
673 user.get('role') != 'admin'
674 and user.get('id') != note.user_id
675 and not await AccessGrants.has_access(
676 user_id=user.get('id'),
677 resource_type='note',
678 resource_id=note.id,
679 permission='read',
680 )
681 ):
682 log.error(f'User {user.get("id")} does not have access to note {note_id}')
683 return
685 user_id = data.get('user_id', sid)
686 user_name = data.get('user_name', 'Anonymous')
687 user_color = data.get('user_color', '#000000')
689 log.info('User %s joining document %s', user_id, document_id)
690 await YDOC_MANAGER.add_user(document_id=document_id, user_id=sid)
692 # Join Socket.IO room
693 await sio.enter_room(sid, f'doc_{document_id}')
695 active_session_ids = get_session_ids_from_room(f'doc_{document_id}')
697 # Get the Yjs document state
698 ydoc = Y.Doc()
699 updates = await YDOC_MANAGER.get_updates(document_id)
700 for update in updates:
701 ydoc.apply_update(bytes(update))
703 # Encode the entire document state as an update
704 state_update = ydoc.get_update()
705 await sio.emit(
706 'ydoc:document:state',
707 {
708 'document_id': document_id,
709 'state': list(state_update), # Convert bytes to list for JSON
710 'sessions': active_session_ids,
711 },
712 room=sid,
713 )
715 # Notify other users about the new user
716 await sio.emit(
717 'ydoc:user:joined',
718 {
719 'document_id': document_id,
720 'user_id': user_id,
721 'user_name': user_name,
722 'user_color': user_color,
723 },
724 room=f'doc_{document_id}',
725 skip_sid=sid,
726 )
728 log.info('User %s successfully joined document %s', user_id, document_id)
730 except Exception as e:
731 log.error(f'Error in yjs_document_join: {e}')
732 await sio.emit('error', {'message': 'Failed to join document'}, room=sid)
735async def document_save_handler(document_id, data, user):
736 document_id = normalize_document_id(document_id)
738 if document_id.startswith('note:'):
739 note_id = document_id.split(':')[1]
740 note = await Notes.get_note_by_id(note_id)
741 if not note:
742 log.error(f'Note {note_id} not found')
743 return
745 if (
746 user.get('role') != 'admin'
747 and user.get('id') != note.user_id
748 and not await AccessGrants.has_access(
749 user_id=user.get('id'),
750 resource_type='note',
751 resource_id=note.id,
752 permission='write',
753 )
754 ):
755 log.error(f'User {user.get("id")} does not have write access to note {note_id}')
756 return
758 await Notes.update_note_by_id(note_id, NoteUpdateForm(data=data))
761@sio.on('ydoc:document:state')
762async def yjs_document_state(sid, data):
763 """Send the current state of the Yjs document to the user"""
764 try:
765 document_id = data['document_id']
767 document_id = normalize_document_id(document_id)
768 room = f'doc_{document_id}'
770 active_session_ids = get_session_ids_from_room(room)
772 if sid not in active_session_ids:
773 log.warning(f'Session {sid} not in room {room}. Cannot send state.')
774 return
776 if not await YDOC_MANAGER.document_exists(document_id):
777 log.warning(f'Document {document_id} not found')
778 return
780 # Get the Yjs document state
781 ydoc = Y.Doc()
782 updates = await YDOC_MANAGER.get_updates(document_id)
783 for update in updates:
784 ydoc.apply_update(bytes(update))
786 # Encode the entire document state as an update
787 state_update = ydoc.get_update()
789 await sio.emit(
790 'ydoc:document:state',
791 {
792 'document_id': document_id,
793 'state': list(state_update), # Convert bytes to list for JSON
794 'sessions': active_session_ids,
795 },
796 room=sid,
797 )
798 except Exception as e:
799 log.error(f'Error in yjs_document_state: {e}')
802@sio.on('ydoc:document:update')
803async def yjs_document_update(sid, data):
804 """Handle Yjs document updates"""
805 try:
806 document_id = data['document_id']
808 document_id = normalize_document_id(document_id)
810 # Verify the sender actually joined this document room
811 room = f'doc_{document_id}'
812 active_session_ids = get_session_ids_from_room(room)
813 if sid not in active_session_ids:
814 log.warning(f'Session {sid} not in room {room}. Rejecting update.')
815 return
817 # Verify write permission — room membership only proves read access
818 user = await get_socket_session_user(sid)
819 if not user:
820 return
822 if document_id.startswith('note:'):
823 note_id = document_id.split(':')[1]
824 note = await Notes.get_note_by_id(note_id)
825 if not note:
826 log.error(f'Note {note_id} not found')
827 return
829 if (
830 user.get('role') != 'admin'
831 and user.get('id') != note.user_id
832 and not await AccessGrants.has_access(
833 user_id=user.get('id'),
834 resource_type='note',
835 resource_id=note.id,
836 permission='write',
837 )
838 ):
839 log.warning(f'User {user.get("id")} does not have write access to note {note_id}. Rejecting update.')
840 return
842 update = data.get('update') # List of bytes from frontend
844 if update:
845 user_id = data.get('user_id', sid)
847 await YDOC_MANAGER.append_to_updates(
848 document_id=document_id,
849 update=update, # Convert list of bytes to bytes
850 )
852 # Broadcast update to all other users in the document
853 await sio.emit(
854 'ydoc:document:update',
855 {
856 'document_id': document_id,
857 'user_id': user_id,
858 'update': update,
859 'socket_id': sid, # Add socket_id to match frontend filtering
860 },
861 room=f'doc_{document_id}',
862 skip_sid=sid,
863 )
865 async def debounced_save():
866 await asyncio.sleep(0.5)
867 await document_save_handler(document_id, data.get('data', {}), user)
869 if data.get('data'):
870 # Only drop the pending save when a new one takes its place.
871 # Updates without a content snapshot (the resync a client sends
872 # after rejoining a document) would otherwise cancel the pending
873 # save without scheduling a replacement, so the edits made just
874 # before the resync never reach the database.
875 try:
876 await stop_item_tasks(REDIS, document_id)
877 except Exception:
878 pass
880 await create_task(REDIS, debounced_save(), document_id)
882 except Exception as e:
883 log.error(f'Error in yjs_document_update: {e}')
886@sio.on('ydoc:document:leave')
887async def yjs_document_leave(sid, data):
888 """Handle user leaving a document"""
889 user = await get_socket_session_user(sid)
890 if not user: # authenticated session required (parity with sibling handlers)
891 return
892 try:
893 document_id = normalize_document_id(data['document_id'])
895 log.info('User %s leaving document %s', user['id'], document_id)
897 # Remove user from the document
898 await YDOC_MANAGER.remove_user(document_id=document_id, user_id=sid)
900 # Leave Socket.IO room
901 await sio.leave_room(sid, f'doc_{document_id}')
903 # Notify other users; user_id is the authenticated identity, not client-supplied
904 await sio.emit(
905 'ydoc:user:left',
906 {'document_id': document_id, 'user_id': user['id']},
907 room=f'doc_{document_id}',
908 )
910 if await YDOC_MANAGER.document_exists(document_id) and len(await YDOC_MANAGER.get_users(document_id)) == 0:
911 log.info('Cleaning up document %s as no users are left', document_id)
912 await YDOC_MANAGER.clear_document(document_id)
914 except Exception as e:
915 log.error(f'Error in yjs_document_leave: {e}')
918@sio.on('ydoc:awareness:update')
919async def yjs_awareness_update(sid, data):
920 """Handle awareness updates (cursors, selections, etc.)"""
921 user = await get_socket_session_user(sid)
922 if not user: # authenticated session required (parity with sibling handlers)
923 return
924 try:
925 document_id = normalize_document_id(data['document_id'])
926 room = f'doc_{document_id}'
927 if room not in sio.rooms(sid): # must have joined the document first
928 return
929 update = data['update']
931 # Broadcast to the room; user_id is the authenticated identity, not client-supplied
932 await sio.emit(
933 'ydoc:awareness:update',
934 {'document_id': document_id, 'user_id': user['id'], 'update': update},
935 room=room,
936 skip_sid=sid,
937 )
939 except Exception as e:
940 log.error(f'Error in yjs_awareness_update: {e}')
943@sio.event
944async def disconnect(sid, reason=None):
945 if sid in SESSION_POOL:
946 del SESSION_POOL[sid]
948 # Clean up USAGE_POOL entries for this session
949 for model_id, connections in list(USAGE_POOL.items()):
950 if sid in connections:
951 del connections[sid]
952 if not connections:
953 del USAGE_POOL[model_id]
954 else:
955 USAGE_POOL[model_id] = connections
957 await YDOC_MANAGER.remove_user_from_all_documents(sid)
958 else:
959 pass
960 # print(f"Unknown session ID {sid} disconnected")
963async def redis_event_listener() -> None:
964 """Route events received over Redis to their local queues."""
965 reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL
967 while True:
968 pubsub = None
969 try:
970 # RedisCluster can't route a pubsub subscribe until initialize() fills its slot cache.
971 await REDIS.initialize()
973 pubsub = REDIS.pubsub()
974 await pubsub.subscribe(REDIS_EVENT_CHANNEL)
975 reconnect_interval = REDIS_PUBSUB_RECONNECT_INTERVAL
977 async for message in pubsub.listen():
978 if message['type'] != 'message':
979 continue
980 event = JSONCodec.loads(message['data'])
981 queue = EVENT_QUEUES.get(event['channel'])
982 if queue is not None:
983 await queue.put(event['data'])
984 log.warning('Redis event listener stopped. Retrying.')
985 except asyncio.CancelledError:
986 raise
987 except Exception:
988 log.exception('Redis event listener failed. Retrying.')
989 finally:
990 if pubsub:
991 with suppress(Exception):
992 await pubsub.aclose()
994 await asyncio.sleep(reconnect_interval)
995 reconnect_interval = min(reconnect_interval * 2, REDIS_PUBSUB_MAX_RECONNECT_INTERVAL)
998@sio.on('*')
999async def socket_event_handler(event: Any, sid: str, *args: Any) -> None:
1000 """Route user-owned stream events to a local queue or another worker."""
1001 if not isinstance(event, str) or event.count(':') != 2 or not args:
1002 return
1004 user = await get_socket_session_user(sid)
1005 if not user or user.get('id') != event.split(':', 1)[0]:
1006 return
1008 queue = EVENT_QUEUES.get(event)
1009 if queue is not None:
1010 await queue.put(args[0])
1011 elif WEBSOCKET_MANAGER == 'redis':
1012 try:
1013 async with EVENT_PUBLISH_LOCK:
1014 await REDIS.publish(REDIS_EVENT_CHANNEL, dumps_bytes({'channel': event, 'data': args[0]}))
1015 except RedisError as e:
1016 log.debug('Failed to relay socket event %s: %s', event, e)
1019async def _make_channel_emitter(request_info):
1020 """Event emitter that routes pipeline output to a channel message.
1022 Translates chat:completion events into channel message:update socket
1023 emissions, throttled to avoid flooding with per-token updates.
1024 """
1025 channel_id = request_info['chat_id'].removeprefix('channel:')
1026 message_id = request_info['message_id']
1028 state = {'last_emit_at': 0.0, 'output': []}
1029 THROTTLE_INTERVAL = 0.15 # ~6 updates/sec
1031 async def _emit_channel_update(
1032 content: str,
1033 done: bool = False,
1034 output: list | None = None,
1035 data: dict | None = None,
1036 ):
1037 from open_webui.models.messages import MessageForm, Messages
1039 msg = await Messages.get_message_by_id(message_id)
1040 if not msg or msg.channel_id != channel_id:
1041 return
1043 update_data = data or ({'output': output} if output else None)
1044 update_form = MessageForm(content=content, data=update_data)
1045 if done:
1046 # Merge done flag into existing meta (preserve model_id etc.)
1047 existing_meta = msg.meta or {}
1048 update_form = MessageForm(
1049 content=content,
1050 data=update_data,
1051 meta={**existing_meta, 'done': True},
1052 )
1054 await Messages.update_message_by_id(message_id, update_form)
1055 message = await Messages.get_message_by_id(message_id)
1056 if message:
1057 await sio.emit(
1058 'events:channel',
1059 {
1060 'channel_id': channel_id,
1061 'message_id': message_id,
1062 'data': {
1063 'type': 'message:update',
1064 'data': message.model_dump(),
1065 },
1066 },
1067 to=f'channel:{channel_id}',
1068 )
1070 async def __channel_emitter__(event_data):
1071 event_type = event_data.get('type')
1073 if event_type == 'chat:completion':
1074 data = event_data.get('data', {})
1075 output = data.get('output')
1076 content = data.get('content') or get_output_text(output)
1077 done = data.get('done', False)
1079 if not content and not output and not done:
1080 return
1082 if isinstance(output, list):
1083 state['output'] = copy.deepcopy(output)
1085 now = time.time()
1086 # Tool boundaries must publish all results before waiting on the next model response.
1087 if done or data.get('flush') or (now - state['last_emit_at']) >= THROTTLE_INTERVAL:
1088 state['last_emit_at'] = now
1089 await _emit_channel_update(content, done, output if isinstance(output, list) else None)
1091 elif event_type == 'response:completion':
1092 from open_webui.utils.middleware import handle_responses_streaming_event
1094 data = event_data.get('data', {})
1095 state['output'], _ = handle_responses_streaming_event(data, state['output'])
1096 content = get_output_text(state['output'])
1098 now = time.time()
1099 if content and (now - state['last_emit_at']) >= THROTTLE_INTERVAL:
1100 state['last_emit_at'] = now
1101 await _emit_channel_update(content, False, state['output'])
1103 elif event_type in ('files', 'chat:message:files'):
1104 from open_webui.models.messages import Messages
1106 files = event_data.get('data', {}).get('files', [])
1107 if not files:
1108 return
1110 msg = await Messages.get_message_by_id(message_id)
1111 if not msg or msg.channel_id != channel_id:
1112 return
1114 existing_files = (msg.data or {}).get('files')
1115 for file in files:
1116 if isinstance(file, dict) and file.get('id'):
1117 file['url'] = file['id']
1118 await Channels.add_file_to_channel_by_id(channel_id, file['id'], msg.user_id)
1119 await Channels.set_file_message_id_in_channel_by_id(channel_id, file['id'], message_id)
1121 if isinstance(existing_files, list):
1122 files.extend(existing_files)
1124 await _emit_channel_update(msg.content, data={'files': files})
1126 elif event_type == 'chat:message:error':
1127 error = event_data.get('data', {}).get('error', {})
1128 error_content = error.get('content', 'An error occurred') if isinstance(error, dict) else str(error)
1129 await _emit_channel_update(f'Error: {error_content}', done=True)
1131 return __channel_emitter__
1134async def get_event_emitter(request_info, update_db=True):
1135 # Channel mode: route pipeline output to channel message updates
1136 if (request_info.get('chat_id') or '').startswith('channel:'): 1136 ↛ 1137line 1136 didn't jump to line 1137 because the condition on line 1136 was never true
1137 return await _make_channel_emitter(request_info)
1139 async def __event_emitter__(event_data):
1140 user_id = request_info['user_id']
1141 chat_id = request_info['chat_id']
1142 message_id = request_info['message_id']
1143 internal = request_info.get('internal') is True
1144 save_to_chat = update_db and message_id and is_saved_chat_id(chat_id)
1146 if internal and event_data.get('type') == 'notification': 1146 ↛ 1147line 1146 didn't jump to line 1147 because the condition on line 1146 was never true
1147 return
1149 room = f'user:{user_id}'
1150 # Local rooms are authoritative; Redis may have listeners on another instance.
1151 if WEBSOCKET_MANAGER == 'redis' or room in sio.manager.rooms.get('/', {}): 1151 ↛ 1152line 1151 didn't jump to line 1152 because the condition on line 1151 was never true
1152 await sio.emit(
1153 'events',
1154 {
1155 'chat_id': chat_id,
1156 'message_id': message_id,
1157 **({'internal': True} if internal else {}),
1158 'data': event_data,
1159 },
1160 room=room,
1161 )
1163 if save_to_chat:
1164 event_type = event_data.get('type')
1166 if event_type == 'status': 1166 ↛ 1167line 1166 didn't jump to line 1167 because the condition on line 1166 was never true
1167 await Chats.add_message_status_to_chat_by_id_and_message_id(
1168 request_info['chat_id'],
1169 request_info['message_id'],
1170 event_data.get('data', {}),
1171 )
1173 elif event_type == 'message': 1173 ↛ 1174line 1173 didn't jump to line 1174 because the condition on line 1173 was never true
1174 message = await Chats.get_message_by_id_and_message_id(
1175 request_info['chat_id'],
1176 request_info['message_id'],
1177 )
1179 if message:
1180 content = message.get('content', '')
1181 content += event_data.get('data', {}).get('content', '')
1183 await Chats.upsert_message_to_chat_by_id_and_message_id(
1184 request_info['chat_id'],
1185 request_info['message_id'],
1186 {
1187 'content': content,
1188 },
1189 )
1191 elif event_type == 'replace': 1191 ↛ 1192line 1191 didn't jump to line 1192 because the condition on line 1191 was never true
1192 content = event_data.get('data', {}).get('content', '')
1194 await Chats.upsert_message_to_chat_by_id_and_message_id(
1195 request_info['chat_id'],
1196 request_info['message_id'],
1197 {
1198 'content': content,
1199 },
1200 )
1202 elif event_type == 'embeds': 1202 ↛ 1203line 1202 didn't jump to line 1203 because the condition on line 1202 was never true
1203 event_payload = event_data.get('data', {})
1204 embeds = event_payload.get('embeds', [])
1206 if not event_payload.get('replace', False):
1207 existing_embeds = await Chats.get_message_metadata(chat_id, message_id, 'embeds')
1208 if isinstance(existing_embeds, list):
1209 embeds.extend(existing_embeds)
1211 await Chats.upsert_message_to_chat_by_id_and_message_id(
1212 chat_id,
1213 message_id,
1214 {
1215 'embeds': embeds,
1216 },
1217 touch=False,
1218 )
1220 elif event_type == 'files': 1220 ↛ 1221line 1220 didn't jump to line 1221 because the condition on line 1220 was never true
1221 files = event_data.get('data', {}).get('files', [])
1222 existing_files = await Chats.get_message_metadata(chat_id, message_id, 'files')
1223 if isinstance(existing_files, list):
1224 files.extend(existing_files)
1226 await Chats.upsert_message_to_chat_by_id_and_message_id(
1227 chat_id,
1228 message_id,
1229 {
1230 'files': files,
1231 },
1232 touch=False,
1233 )
1235 elif event_type in ('source', 'citation'): 1235 ↛ 1236line 1235 didn't jump to line 1236 because the condition on line 1235 was never true
1236 data = event_data.get('data', {})
1237 if data.get('type') is None:
1238 sources = await Chats.get_message_metadata(chat_id, message_id, 'sources')
1239 if not isinstance(sources, list):
1240 sources = []
1241 sources.append(data)
1243 await Chats.upsert_message_to_chat_by_id_and_message_id(
1244 chat_id,
1245 message_id,
1246 {
1247 'sources': sources,
1248 },
1249 touch=False,
1250 )
1252 if 'user_id' in request_info and 'chat_id' in request_info and 'message_id' in request_info: 1252 ↛ 1255line 1252 didn't jump to line 1255 because the condition on line 1252 was always true
1253 return __event_emitter__
1254 else:
1255 return None
1258async def get_event_call(request_info):
1259 async def __event_caller__(event_data):
1260 session_id = request_info['session_id']
1262 # session_id is client-supplied; only the requesting user's own live session may be targeted.
1263 session = SESSION_POOL.get(session_id)
1264 if session is None or session.get('id') != request_info.get('user_id'):
1265 log.warning(f'Event caller: session {session_id} not owned by requesting user or disconnected')
1266 return {'error': 'Client session disconnected.'}
1268 try:
1269 return await sio.call(
1270 'events',
1271 {
1272 'chat_id': request_info.get('chat_id', None),
1273 'message_id': request_info.get('message_id', None),
1274 'data': event_data,
1275 },
1276 to=session_id,
1277 timeout=WEBSOCKET_EVENT_CALLER_TIMEOUT,
1278 )
1279 except (TimeoutError, socketio.exceptions.TimeoutError):
1280 log.warning(f'Event caller timed out for session {session_id}')
1281 return {'error': 'Event call timed out. The browser tab may be inactive or closed.'}
1283 if 'session_id' in request_info and 'chat_id' in request_info and 'message_id' in request_info:
1284 return __event_caller__
1285 else:
1286 return None
1289get_event_caller = get_event_call