From 8594b2f421f69742b64297fcee85b4bb0125f238 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Sat, 10 Oct 2026 22:05:25 +0400 Subject: [PATCH] refac --- backend/open_webui/socket/main.py | 52 ++++++++++++++----- backend/open_webui/socket/utils.py | 51 ++++++++++-------- .../common/RichTextInput/Collaboration.ts | 22 +++++--- 3 files changed, 83 insertions(+), 42 deletions(-) diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index c85e6fea99..2e29d6d53c 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -6,8 +6,8 @@ import logging import random import sys import time -import weakref from contextlib import suppress +from functools import wraps from typing import Any from uuid import uuid4 @@ -43,11 +43,13 @@ 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.redis_room_channels import AsyncRedisRoomChannelManager -from open_webui.socket.utils import CachedRedisDict, RedisDict, RedisLock, YdocManager +from open_webui.socket.utils import SOCKET_EVENT_LOCKS, CachedRedisDict, RedisDict, RedisLock, YdocManager from open_webui.tasks import ( REDIS_PUBSUB_MAX_RECONNECT_INTERVAL, REDIS_PUBSUB_RECONNECT_INTERVAL, + cleanup_task, create_task, + has_active_tasks, stop_item_tasks, ) from open_webui.utils.access_control import has_permission @@ -211,7 +213,6 @@ 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(): @@ -349,7 +350,7 @@ async def periodic_socket_authentication(): await get_socket_session_user(sid) -async def get_socket_session_user(sid: str) -> dict | None: +async def get_socket_session_user(sid: str, *, wait_for_disconnect: bool = True) -> dict | None: """Session user from this worker's local Socket.IO store; only locally connected sids are ever looked up.""" try: session = await sio.get_session(sid) @@ -358,7 +359,11 @@ async def get_socket_session_user(sid: str) -> dict | None: except Exception: log.debug('Socket authentication expired for %s', sid) LOCAL_AUTHENTICATED_SIDS.discard(sid) - await sio.disconnect(sid) + if wait_for_disconnect: + await sio.disconnect(sid) + else: + # Document handlers hold a lock that disconnect cleanup also needs. + sio.start_background_task(sio.disconnect, sid) return None @@ -788,10 +793,23 @@ def normalize_document_id(document_id: str) -> str: return document_id +def with_document_lock(handler): + @wraps(handler) + async def wrapped(sid, data): + try: + async with YDOC_MANAGER.lock(normalize_document_id(data['document_id'])): + return await handler(sid, data) + except Exception: + log.exception('Error in %s', handler.__name__) + + return wrapped + + @sio.on('ydoc:document:join') +@with_document_lock async def ydoc_document_join(sid, data): """Handle user joining a document""" - user = await get_socket_session_user(sid) + user = await get_socket_session_user(sid, wait_for_disconnect=False) if not user: return @@ -855,6 +873,7 @@ async def ydoc_document_join(sid, data): { 'document_id': document_id, 'state': list(state_update), # Convert bytes to list for JSON + 'content': note.data.get('content') if document_id.startswith('note:') and note.data else None, 'sessions': active_session_ids, }, room=sid, @@ -903,10 +922,11 @@ async def document_save_handler(document_id, data, user): 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)) + return await Notes.update_note_by_id(note_id, NoteUpdateForm(data=data)) @sio.on('ydoc:document:state') +@with_document_lock async def yjs_document_state(sid, data): """Send the current state of the Yjs document to the user""" try: @@ -948,6 +968,7 @@ async def yjs_document_state(sid, data): @sio.on('ydoc:document:update') +@with_document_lock async def yjs_document_update(sid, data): """Handle Yjs document updates""" try: @@ -963,7 +984,7 @@ async def yjs_document_update(sid, data): return # Verify write permission — room membership only proves read access - user = await get_socket_session_user(sid) + user = await get_socket_session_user(sid, wait_for_disconnect=False) if not user: return @@ -1014,7 +1035,11 @@ async def yjs_document_update(sid, data): async def debounced_save(): await asyncio.sleep(0.5) - await document_save_handler(document_id, data.get('data', {}), user) + async with YDOC_MANAGER.lock(document_id): + if await document_save_handler(document_id, data.get('data', {}), user): + if not await YDOC_MANAGER.get_users(document_id): + await YDOC_MANAGER.clear_document(document_id) + await cleanup_task(REDIS, task_id, document_id) if document_id.startswith('note:') and data.get('data'): # Only drop the pending save when a new one takes its place. @@ -1027,16 +1052,17 @@ async def yjs_document_update(sid, data): except Exception: pass - await create_task(REDIS, debounced_save(), document_id) + task_id, _ = 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') +@with_document_lock async def yjs_document_leave(sid, data): """Handle user leaving a document""" - user = await get_socket_session_user(sid) + user = await get_socket_session_user(sid, wait_for_disconnect=False) if not user: # authenticated session required (parity with sibling handlers) return try: @@ -1057,7 +1083,7 @@ async def yjs_document_leave(sid, data): room=f'doc_{document_id}', ) - if await YDOC_MANAGER.document_exists(document_id) and len(await YDOC_MANAGER.get_users(document_id)) == 0: + if not await YDOC_MANAGER.get_users(document_id) and not await has_active_tasks(REDIS, document_id): log.info('Cleaning up document %s as no users are left', document_id) await YDOC_MANAGER.clear_document(document_id) @@ -1154,7 +1180,7 @@ async def socket_event_handler(event: Any, sid: str, *args: Any) -> None: # 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()): + async with SOCKET_EVENT_LOCKS.setdefault(('session', 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 diff --git a/backend/open_webui/socket/utils.py b/backend/open_webui/socket/utils.py index 7cdc5490c0..65234ac0e5 100644 --- a/backend/open_webui/socket/utils.py +++ b/backend/open_webui/socket/utils.py @@ -2,12 +2,16 @@ from __future__ import annotations +import asyncio import hashlib import logging import uuid +import weakref +from contextlib import asynccontextmanager import pycrdt as Y -from open_webui.env import REDIS_KEY_PREFIX +from open_webui.env import REDIS_KEY_PREFIX, WEBSOCKET_REDIS_LOCK_TIMEOUT +from open_webui.tasks import has_active_tasks from open_webui.utils.json_codec import JSONCodec from open_webui.utils.redis import get_redis_connection from redis.exceptions import RedisClusterException, RedisError @@ -17,6 +21,8 @@ log = logging.getLogger(__name__) YDOC_KEY_PREFIX = f'{REDIS_KEY_PREFIX}:ydoc:documents' SCAN_BATCH_SIZE = 200 +SOCKET_EVENT_LOCKS: weakref.WeakValueDictionary[tuple[str, str], asyncio.Lock] = weakref.WeakValueDictionary() + class RedisLock: """Distributed lock backed by a Redis SET with NX/EX semantics.""" @@ -253,6 +259,22 @@ class YdocManager: self._redis = redis self._redis_key_prefix = redis_key_prefix + @asynccontextmanager + async def lock(self, document_id: str): + # Local FIFO ordering also covers permission checks before an update is stored. + async with SOCKET_EVENT_LOCKS.setdefault(('document', document_id), asyncio.Lock()): + if self._redis: + async with self._redis.lock( + f'{self._redis_key_prefix}:{document_id}:lock', + timeout=WEBSOCKET_REDIS_LOCK_TIMEOUT, + blocking_timeout=WEBSOCKET_REDIS_LOCK_TIMEOUT, + ): + # Cancel stalled work before another worker can acquire the expired lease. + async with asyncio.timeout(WEBSOCKET_REDIS_LOCK_TIMEOUT / 2): + yield + else: + yield + async def append_to_updates(self, document_id: str, update: list[int]) -> bool: if not isinstance(update, list): return False @@ -373,31 +395,16 @@ class YdocManager: async def remove_user_from_all_documents(self, user_id: str): if self._redis: - # Use the per-session reverse index instead of a cluster-wide - # SCAN. This set contains only the document IDs that this - # session actually joined, so the cost is proportional to - # the session's footprint — not the total number of documents. session_key = f'{self._redis_key_prefix}:session:{user_id}:documents' document_ids = await self._redis.smembers(session_key) - - for document_id in document_ids: - users_key = f'{self._redis_key_prefix}:{document_id}:users' - await self._redis.srem(users_key, user_id) - - if len(await self.get_users(document_id)) == 0: - await self.clear_document(document_id) - - # Clean up the reverse index itself. - await self._redis.delete(session_key) - else: - for document_id in list(self._users.keys()): - if user_id in self._users[document_id]: - self._users[document_id].remove(user_id) - if not self._users[document_id]: - del self._users[document_id] + document_ids = [document_id for document_id, users in self._users.items() if user_id in users] - await self.clear_document(document_id) + for document_id in document_ids: + async with self.lock(document_id): + await self.remove_user(document_id, user_id) + if not await self.get_users(document_id) and not await has_active_tasks(self._redis, document_id): + await self.clear_document(document_id) async def clear_document(self, document_id: str): if self._redis: diff --git a/src/lib/components/common/RichTextInput/Collaboration.ts b/src/lib/components/common/RichTextInput/Collaboration.ts index 0496704342..d16b0a5913 100644 --- a/src/lib/components/common/RichTextInput/Collaboration.ts +++ b/src/lib/components/common/RichTextInput/Collaboration.ts @@ -1,4 +1,5 @@ import * as Y from 'yjs'; +import { marked } from 'marked'; import { ySyncPlugin, ySyncPluginKey, @@ -103,15 +104,15 @@ export class SocketIOCollaborationProvider { } } - private applyInitialContent() { - if (!this.editor || !this.initialContent) return; + private applyInitialContent(content = this.initialContent) { + if (!this.editor || !content) return; - if (typeof this.initialContent === 'string') { - this.editor.commands.setContent(this.initialContent); + if (typeof content === 'string') { + this.editor.commands.setContent(content); return; } - const doc = prosemirrorJSONToYDoc(this.editor.schema, this.initialContent); + const doc = prosemirrorJSONToYDoc(this.editor.schema, content); Y.applyUpdate(this.doc, Y.encodeStateAsUpdate(doc)); } @@ -181,10 +182,17 @@ export class SocketIOCollaborationProvider { this.doc.getXmlFragment('prosemirror').length === 0 ) { if ( - this.initialContent && + (data.content || this.initialContent) && [...(data.sessions ?? [])].sort()[0] === this.socket.id ) { - this.applyInitialContent(); + // The HTTP snapshot may predate the last socket save. + this.applyInitialContent( + data.content + ? data.content.json || + data.content.html || + marked.parse(data.content.md ?? '', { async: false }) + : this.initialContent + ); } } else { // If the editor already has content, we don't need to send an empty state