This commit is contained in:
Timothy Jaeryang Baek 2026-10-10 22:05:25 +04:00
parent 1469f73b31
commit 8594b2f421
3 changed files with 83 additions and 42 deletions

View file

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

View file

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

View file

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