"""Redis-backed distributed data structures for WebSocket state management.""" 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, 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 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.""" _RENEW_SCRIPT = """ if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('expire', KEYS[1], ARGV[2]) end return 0 """ _RELEASE_SCRIPT = """ if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) end return 0 """ def __init__( self, redis_url, lock_name, timeout_secs, redis_sentinels=[], redis_cluster=False, ): self.lock_name = lock_name self.lock_id = str(uuid.uuid4()) self.timeout_secs = timeout_secs self.lock_obtained = False self.redis = get_redis_connection( redis_url, redis_sentinels, redis_cluster=redis_cluster, decode_responses=True, ) def aquire_lock(self): # nx=True will only set this key if it _hasn't_ already been set self.lock_obtained = self.redis.set(self.lock_name, self.lock_id, nx=True, ex=self.timeout_secs) return self.lock_obtained def renew_lock(self): return bool(self.redis.eval(self._RENEW_SCRIPT, 1, self.lock_name, self.lock_id, self.timeout_secs)) def release_lock(self): try: self.redis.eval(self._RELEASE_SCRIPT, 1, self.lock_name, self.lock_id) except (RedisClusterException, RedisError) as e: log.warning('Failed to release lock %s; it expires on its own: %s', self.lock_name, e) class RedisDict: def __init__( self, name, redis_url, redis_sentinels=[], redis_cluster=False, cache_set_signature=False, ): self.name = name self._signature_name = f'{name}:signature' if cache_set_signature else None self.redis = get_redis_connection( redis_url, redis_sentinels, redis_cluster=redis_cluster, decode_responses=True, ) def __setitem__(self, key, value): serialized_value = JSONCodec.dumps(value) self.redis.hset(self.name, key, serialized_value) if self._signature_name: self.redis.delete(self._signature_name) def __getitem__(self, key): value = self.redis.hget(self.name, key) if value is None: raise KeyError(key) return JSONCodec.loads(value) def __delitem__(self, key): result = self.redis.hdel(self.name, key) if result == 0: raise KeyError(key) if self._signature_name: self.redis.delete(self._signature_name) def __contains__(self, key): return self.redis.hexists(self.name, key) def __len__(self): return self.redis.hlen(self.name) def keys(self): return self.redis.hkeys(self.name) def values(self): return [JSONCodec.loads(v) for v in self.redis.hvals(self.name)] def items(self): return [(k, JSONCodec.loads(v)) for k, v in self.redis.hgetall(self.name).items()] def scan_batches(self): """Yield lists of (key, value) pairs via incremental HSCAN; a field may repeat across batches.""" cursor = 0 while True: cursor, batch = self.redis.hscan(self.name, cursor, count=SCAN_BATCH_SIZE) if batch: yield [(k, JSONCodec.loads(v)) for k, v in batch.items()] if cursor == 0: break def delete_many(self, *keys): """Delete fields in one HDEL; no keys is a no-op (HDEL rejects an empty field list).""" if keys: self.redis.hdel(self.name, *keys) if self._signature_name: self.redis.delete(self._signature_name) def set(self, mapping: dict): if not mapping: self.clear() return # Serialize values once — reused for both the fingerprint and the write. serialized = {k: JSONCodec.dumps(v) for k, v in mapping.items()} digest = hashlib.sha256() for key in sorted(serialized): digest.update(key.encode()) digest.update(b'\0') digest.update(serialized[key].encode()) digest.update(b'\0') content_digest = digest.hexdigest() if self._signature_name: stored_signature = self.redis.get(self._signature_name) if stored_signature and stored_signature.startswith(f'{content_digest}:'): return # Cleared first so readers refetch while the hash is being rewritten. self.redis.delete(self._signature_name) # Fetch existing keys before writing so we know which ones to remove. # HKEYS is cheap — it transfers only short key strings, not large JSON values. existing_keys = set(self.redis.hkeys(self.name)) new_keys = set(mapping.keys()) keys_to_remove = existing_keys - new_keys # HSET first (add/update all new values), then HDEL (remove stale keys). # We never DELETE the whole hash — this eliminates the race window # where concurrent readers would see an empty models dict. self.redis.hset(self.name, mapping=serialized) if keys_to_remove: self.redis.hdel(self.name, *keys_to_remove) if self._signature_name: self.redis.set(self._signature_name, f'{content_digest}:{uuid.uuid4().hex}') def get(self, key, default=None): try: return self[key] except KeyError: return default def clear(self): if self._signature_name: self.redis.delete(self.name) self.redis.delete(self._signature_name) else: self.redis.delete(self.name) def update(self, other=None, **kwargs): if other is not None: for k, v in other.items() if hasattr(other, 'items') else other: self[k] = v for k, v in kwargs.items(): self[k] = v def setdefault(self, key, default=None): if key not in self: self[key] = default return self[key] class CachedRedisDict(RedisDict): """Answers reads from a per-worker cache of the hash, refetched whenever its signature changes.""" def __init__(self, name: str, redis_url: str, redis_sentinels: list = [], redis_cluster: bool = False): super().__init__(name, redis_url, redis_sentinels, redis_cluster, cache_set_signature=True) self._cache: dict = {} self._cached_signature: str | None = None def _refresh_cache(self) -> dict: stored_signature = self.redis.get(self._signature_name) if stored_signature is None or stored_signature != self._cached_signature: self._cache = self.redis.hgetall(self.name) self._cached_signature = stored_signature return self._cache def __getitem__(self, key): value = self._refresh_cache().get(key) if value is None: raise KeyError(key) return JSONCodec.loads(value) def __contains__(self, key): return key in self._refresh_cache() def __len__(self): return len(self._refresh_cache()) def keys(self): return list(self._refresh_cache().keys()) def values(self): return [JSONCodec.loads(v) for v in self._refresh_cache().values()] def items(self): return [(k, JSONCodec.loads(v)) for k, v in self._refresh_cache().items()] class YdocManager: COMPACTION_THRESHOLD = 500 MAX_DOCUMENTS_PER_SESSION = 20 MAX_DOCUMENT_SIZE = 2 * 1024 * 1024 def __init__( self, redis=None, redis_key_prefix: str = YDOC_KEY_PREFIX, ): self._updates = {} self._users = {} 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 try: update_bytes = bytes(update) Y.Doc().apply_update(update_bytes) # undecodable updates would stall compaction forever except (TypeError, ValueError): return False update_size = len(update_bytes) if update_size > self.MAX_DOCUMENT_SIZE: return False if self._redis: size_key = f'{self._redis_key_prefix}:{document_id}:size' if await self._redis.incrby(size_key, update_size) > self.MAX_DOCUMENT_SIZE: await self._redis.decrby(size_key, update_size) return False redis_key = f'{self._redis_key_prefix}:{document_id}:updates' await self._redis.rpush(redis_key, JSONCodec.dumps(list(update))) list_len = await self._redis.llen(redis_key) if list_len >= self.COMPACTION_THRESHOLD: await self._compact_updates_redis(document_id) else: if sum(len(u) for u in self._updates.get(document_id, [])) + update_size > self.MAX_DOCUMENT_SIZE: return False if document_id not in self._updates: self._updates[document_id] = [] self._updates[document_id].append(update_bytes) if len(self._updates[document_id]) >= self.COMPACTION_THRESHOLD: self._compact_updates_memory(document_id) return True async def count_documents_for_user(self, user_id: str) -> int: """Number of documents this session currently participates in.""" if self._redis: session_key = f'{self._redis_key_prefix}:session:{user_id}:documents' return await self._redis.scard(session_key) return sum(1 for members in self._users.values() if user_id in members) async def _compact_updates_redis(self, document_id: str): """Rolling compaction: squash oldest half into one snapshot.""" redis_key = f'{self._redis_key_prefix}:{document_id}:updates' all_updates = await self._redis.lrange(redis_key, 0, -1) if len(all_updates) <= 1: return mid = len(all_updates) // 2 ydoc = Y.Doc() for raw in all_updates[:mid]: ydoc.apply_update(bytes(JSONCodec.loads(raw))) snapshot_list = list(ydoc.get_update()) snapshot = JSONCodec.dumps(snapshot_list) pipe = self._redis.pipeline() pipe.delete(redis_key) pipe.rpush(redis_key, snapshot, *all_updates[mid:]) await pipe.execute() new_size = len(snapshot_list) + sum(len(JSONCodec.loads(raw)) for raw in all_updates[mid:]) await self._redis.set(f'{self._redis_key_prefix}:{document_id}:size', new_size) def _compact_updates_memory(self, document_id: str): """Rolling compaction: squash oldest half into one snapshot.""" updates = self._updates.get(document_id, []) if len(updates) <= 1: return mid = len(updates) // 2 ydoc = Y.Doc() for update in updates[:mid]: ydoc.apply_update(bytes(update)) self._updates[document_id] = [ydoc.get_update()] + updates[mid:] async def get_updates(self, document_id: str) -> list[bytes]: if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:updates' updates = await self._redis.lrange(redis_key, 0, -1) return [bytes(JSONCodec.loads(update)) for update in updates] else: return self._updates.get(document_id, []) async def document_exists(self, document_id: str) -> bool: if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:updates' return await self._redis.exists(redis_key) > 0 else: return document_id in self._updates async def get_users(self, document_id: str) -> list[str]: if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:users' users = await self._redis.smembers(redis_key) return list(users) else: return self._users.get(document_id, []) async def add_user(self, document_id: str, user_id: str): if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:users' await self._redis.sadd(redis_key, user_id) # Maintain a per-session reverse index so disconnect cleanup # can look up only the documents this session joined, instead # of issuing a cluster-wide SCAN over the entire keyspace. session_key = f'{self._redis_key_prefix}:session:{user_id}:documents' await self._redis.sadd(session_key, document_id) else: if document_id not in self._users: self._users[document_id] = set() self._users[document_id].add(user_id) async def remove_user(self, document_id: str, user_id: str): if self._redis: redis_key = f'{self._redis_key_prefix}:{document_id}:users' await self._redis.srem(redis_key, user_id) # Keep the reverse index in sync. session_key = f'{self._redis_key_prefix}:session:{user_id}:documents' await self._redis.srem(session_key, document_id) else: if document_id in self._users and 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] async def remove_user_from_all_documents(self, user_id: str): if self._redis: session_key = f'{self._redis_key_prefix}:session:{user_id}:documents' document_ids = await self._redis.smembers(session_key) else: document_ids = [document_id for document_id, users in self._users.items() if user_id in users] 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: redis_key = f'{self._redis_key_prefix}:{document_id}:updates' await self._redis.delete(redis_key) redis_users_key = f'{self._redis_key_prefix}:{document_id}:users' await self._redis.delete(redis_users_key) await self._redis.delete(f'{self._redis_key_prefix}:{document_id}:size') else: if document_id in self._updates: del self._updates[document_id] if document_id in self._users: del self._users[document_id]