mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-11 03:38:02 +00:00
refac
This commit is contained in:
parent
1469f73b31
commit
8594b2f421
3 changed files with 83 additions and 42 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue