Coverage for open_webui/socket/utils.py: 18%
255 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
1"""Redis-backed distributed data structures for WebSocket state management."""
3from __future__ import annotations
5import hashlib
6import logging
7import uuid
9import pycrdt as Y
10from open_webui.env import REDIS_KEY_PREFIX
11from open_webui.utils.json_codec import JSONCodec
12from open_webui.utils.redis import get_redis_connection
13from redis.exceptions import RedisClusterException, RedisError
15log = logging.getLogger(__name__)
17YDOC_KEY_PREFIX = f'{REDIS_KEY_PREFIX}:ydoc:documents'
18SCAN_BATCH_SIZE = 200
21class RedisLock:
22 """Distributed lock backed by a Redis SET with NX/EX semantics."""
24 _RENEW_SCRIPT = """
25 if redis.call('get', KEYS[1]) == ARGV[1] then
26 return redis.call('expire', KEYS[1], ARGV[2])
27 end
28 return 0
29 """
30 _RELEASE_SCRIPT = """
31 if redis.call('get', KEYS[1]) == ARGV[1] then
32 return redis.call('del', KEYS[1])
33 end
34 return 0
35 """
37 def __init__(
38 self,
39 redis_url,
40 lock_name,
41 timeout_secs,
42 redis_sentinels=[],
43 redis_cluster=False,
44 ):
45 self.lock_name = lock_name
46 self.lock_id = str(uuid.uuid4())
47 self.timeout_secs = timeout_secs
48 self.lock_obtained = False
49 self.redis = get_redis_connection(
50 redis_url,
51 redis_sentinels,
52 redis_cluster=redis_cluster,
53 decode_responses=True,
54 )
56 def aquire_lock(self):
57 # nx=True will only set this key if it _hasn't_ already been set
58 self.lock_obtained = self.redis.set(self.lock_name, self.lock_id, nx=True, ex=self.timeout_secs)
59 return self.lock_obtained
61 def renew_lock(self):
62 return bool(self.redis.eval(self._RENEW_SCRIPT, 1, self.lock_name, self.lock_id, self.timeout_secs))
64 def release_lock(self):
65 try:
66 self.redis.eval(self._RELEASE_SCRIPT, 1, self.lock_name, self.lock_id)
67 except (RedisClusterException, RedisError) as e:
68 log.warning('Failed to release lock %s; it expires on its own: %s', self.lock_name, e)
71class RedisDict:
72 def __init__(
73 self,
74 name,
75 redis_url,
76 redis_sentinels=[],
77 redis_cluster=False,
78 cache_set_signature=False,
79 ):
80 self.name = name
81 self._signature_name = f'{name}:signature' if cache_set_signature else None
82 self.redis = get_redis_connection(
83 redis_url,
84 redis_sentinels,
85 redis_cluster=redis_cluster,
86 decode_responses=True,
87 )
89 def __setitem__(self, key, value):
90 serialized_value = JSONCodec.dumps(value)
91 self.redis.hset(self.name, key, serialized_value)
92 if self._signature_name:
93 self.redis.delete(self._signature_name)
95 def __getitem__(self, key):
96 value = self.redis.hget(self.name, key)
97 if value is None:
98 raise KeyError(key)
99 return JSONCodec.loads(value)
101 def __delitem__(self, key):
102 result = self.redis.hdel(self.name, key)
103 if result == 0:
104 raise KeyError(key)
105 if self._signature_name:
106 self.redis.delete(self._signature_name)
108 def __contains__(self, key):
109 return self.redis.hexists(self.name, key)
111 def __len__(self):
112 return self.redis.hlen(self.name)
114 def keys(self):
115 return self.redis.hkeys(self.name)
117 def values(self):
118 return [JSONCodec.loads(v) for v in self.redis.hvals(self.name)]
120 def items(self):
121 return [(k, JSONCodec.loads(v)) for k, v in self.redis.hgetall(self.name).items()]
123 def scan_batches(self):
124 """Yield lists of (key, value) pairs via incremental HSCAN; a field may repeat across batches."""
125 cursor = 0
126 while True:
127 cursor, batch = self.redis.hscan(self.name, cursor, count=SCAN_BATCH_SIZE)
128 if batch:
129 yield [(k, JSONCodec.loads(v)) for k, v in batch.items()]
130 if cursor == 0:
131 break
133 def delete_many(self, *keys):
134 """Delete fields in one HDEL; no keys is a no-op (HDEL rejects an empty field list)."""
135 if keys:
136 self.redis.hdel(self.name, *keys)
137 if self._signature_name:
138 self.redis.delete(self._signature_name)
140 def set(self, mapping: dict):
141 if not mapping:
142 self.clear()
143 return
145 # Serialize values once — reused for both the fingerprint and the write.
146 serialized = {k: JSONCodec.dumps(v) for k, v in mapping.items()}
147 digest = hashlib.sha256()
148 for key in sorted(serialized):
149 digest.update(key.encode())
150 digest.update(b'\0')
151 digest.update(serialized[key].encode())
152 digest.update(b'\0')
153 content_digest = digest.hexdigest()
155 if self._signature_name:
156 stored_signature = self.redis.get(self._signature_name)
157 if stored_signature and stored_signature.startswith(f'{content_digest}:'):
158 return
159 # Cleared first so readers refetch while the hash is being rewritten.
160 self.redis.delete(self._signature_name)
162 # Fetch existing keys before writing so we know which ones to remove.
163 # HKEYS is cheap — it transfers only short key strings, not large JSON values.
164 existing_keys = set(self.redis.hkeys(self.name))
165 new_keys = set(mapping.keys())
166 keys_to_remove = existing_keys - new_keys
168 # HSET first (add/update all new values), then HDEL (remove stale keys).
169 # We never DELETE the whole hash — this eliminates the race window
170 # where concurrent readers would see an empty models dict.
171 self.redis.hset(self.name, mapping=serialized)
172 if keys_to_remove:
173 self.redis.hdel(self.name, *keys_to_remove)
175 if self._signature_name:
176 self.redis.set(self._signature_name, f'{content_digest}:{uuid.uuid4().hex}')
178 def get(self, key, default=None):
179 try:
180 return self[key]
181 except KeyError:
182 return default
184 def clear(self):
185 if self._signature_name:
186 self.redis.delete(self.name)
187 self.redis.delete(self._signature_name)
188 else:
189 self.redis.delete(self.name)
191 def update(self, other=None, **kwargs):
192 if other is not None:
193 for k, v in other.items() if hasattr(other, 'items') else other:
194 self[k] = v
195 for k, v in kwargs.items():
196 self[k] = v
198 def setdefault(self, key, default=None):
199 if key not in self:
200 self[key] = default
201 return self[key]
204class CachedRedisDict(RedisDict):
205 """Answers reads from a per-worker cache of the hash, refetched whenever its signature changes."""
207 def __init__(self, name: str, redis_url: str, redis_sentinels: list = [], redis_cluster: bool = False):
208 super().__init__(name, redis_url, redis_sentinels, redis_cluster, cache_set_signature=True)
209 self._cache: dict = {}
210 self._cached_signature: str | None = None
212 def _refresh_cache(self) -> dict:
213 stored_signature = self.redis.get(self._signature_name)
214 if stored_signature is None or stored_signature != self._cached_signature:
215 self._cache = self.redis.hgetall(self.name)
216 self._cached_signature = stored_signature
217 return self._cache
219 def __getitem__(self, key):
220 value = self._refresh_cache().get(key)
221 if value is None:
222 raise KeyError(key)
223 return JSONCodec.loads(value)
225 def __contains__(self, key):
226 return key in self._refresh_cache()
228 def __len__(self):
229 return len(self._refresh_cache())
231 def keys(self):
232 return list(self._refresh_cache().keys())
234 def values(self):
235 return [JSONCodec.loads(v) for v in self._refresh_cache().values()]
237 def items(self):
238 return [(k, JSONCodec.loads(v)) for k, v in self._refresh_cache().items()]
241class YdocManager:
242 COMPACTION_THRESHOLD = 500
244 def __init__(
245 self,
246 redis=None,
247 redis_key_prefix: str = YDOC_KEY_PREFIX,
248 ):
249 self._updates = {}
250 self._users = {}
251 self._redis = redis
252 self._redis_key_prefix = redis_key_prefix
254 async def append_to_updates(self, document_id: str, update: bytes):
255 document_id = document_id.replace(':', '_')
256 if self._redis:
257 redis_key = f'{self._redis_key_prefix}:{document_id}:updates'
258 await self._redis.rpush(redis_key, JSONCodec.dumps(list(update)))
259 list_len = await self._redis.llen(redis_key)
260 if list_len >= self.COMPACTION_THRESHOLD:
261 await self._compact_updates_redis(document_id)
262 else:
263 if document_id not in self._updates:
264 self._updates[document_id] = []
265 self._updates[document_id].append(update)
266 if len(self._updates[document_id]) >= self.COMPACTION_THRESHOLD:
267 self._compact_updates_memory(document_id)
269 async def _compact_updates_redis(self, document_id: str):
270 """Rolling compaction: squash oldest half into one snapshot."""
271 redis_key = f'{self._redis_key_prefix}:{document_id}:updates'
272 all_updates = await self._redis.lrange(redis_key, 0, -1)
273 if len(all_updates) <= 1:
274 return
275 mid = len(all_updates) // 2
276 ydoc = Y.Doc()
277 for raw in all_updates[:mid]:
278 ydoc.apply_update(bytes(JSONCodec.loads(raw)))
279 snapshot = JSONCodec.dumps(list(ydoc.get_update()))
280 pipe = self._redis.pipeline()
281 pipe.delete(redis_key)
282 pipe.rpush(redis_key, snapshot, *all_updates[mid:])
283 await pipe.execute()
285 def _compact_updates_memory(self, document_id: str):
286 """Rolling compaction: squash oldest half into one snapshot."""
287 updates = self._updates.get(document_id, [])
288 if len(updates) <= 1:
289 return
290 mid = len(updates) // 2
291 ydoc = Y.Doc()
292 for update in updates[:mid]:
293 ydoc.apply_update(bytes(update))
294 self._updates[document_id] = [ydoc.get_update()] + updates[mid:]
296 async def get_updates(self, document_id: str) -> list[bytes]:
297 document_id = document_id.replace(':', '_')
299 if self._redis:
300 redis_key = f'{self._redis_key_prefix}:{document_id}:updates'
301 updates = await self._redis.lrange(redis_key, 0, -1)
302 return [bytes(JSONCodec.loads(update)) for update in updates]
303 else:
304 return self._updates.get(document_id, [])
306 async def document_exists(self, document_id: str) -> bool:
307 document_id = document_id.replace(':', '_')
309 if self._redis:
310 redis_key = f'{self._redis_key_prefix}:{document_id}:updates'
311 return await self._redis.exists(redis_key) > 0
312 else:
313 return document_id in self._updates
315 async def get_users(self, document_id: str) -> list[str]:
316 document_id = document_id.replace(':', '_')
318 if self._redis:
319 redis_key = f'{self._redis_key_prefix}:{document_id}:users'
320 users = await self._redis.smembers(redis_key)
321 return list(users)
322 else:
323 return self._users.get(document_id, [])
325 async def add_user(self, document_id: str, user_id: str):
326 document_id = document_id.replace(':', '_')
328 if self._redis:
329 redis_key = f'{self._redis_key_prefix}:{document_id}:users'
330 await self._redis.sadd(redis_key, user_id)
331 # Maintain a per-session reverse index so disconnect cleanup
332 # can look up only the documents this session joined, instead
333 # of issuing a cluster-wide SCAN over the entire keyspace.
334 session_key = f'{self._redis_key_prefix}:session:{user_id}:documents'
335 await self._redis.sadd(session_key, document_id)
336 else:
337 if document_id not in self._users:
338 self._users[document_id] = set()
339 self._users[document_id].add(user_id)
341 async def remove_user(self, document_id: str, user_id: str):
342 document_id = document_id.replace(':', '_')
344 if self._redis:
345 redis_key = f'{self._redis_key_prefix}:{document_id}:users'
346 await self._redis.srem(redis_key, user_id)
347 # Keep the reverse index in sync.
348 session_key = f'{self._redis_key_prefix}:session:{user_id}:documents'
349 await self._redis.srem(session_key, document_id)
350 else:
351 if document_id in self._users and user_id in self._users[document_id]:
352 self._users[document_id].remove(user_id)
354 async def remove_user_from_all_documents(self, user_id: str):
355 if self._redis:
356 # Use the per-session reverse index instead of a cluster-wide
357 # SCAN. This set contains only the document IDs that this
358 # session actually joined, so the cost is proportional to
359 # the session's footprint — not the total number of documents.
360 session_key = f'{self._redis_key_prefix}:session:{user_id}:documents'
361 document_ids = await self._redis.smembers(session_key)
363 for document_id in document_ids:
364 users_key = f'{self._redis_key_prefix}:{document_id}:users'
365 await self._redis.srem(users_key, user_id)
367 if len(await self.get_users(document_id)) == 0:
368 await self.clear_document(document_id)
370 # Clean up the reverse index itself.
371 await self._redis.delete(session_key)
373 else:
374 for document_id in list(self._users.keys()):
375 if user_id in self._users[document_id]:
376 self._users[document_id].remove(user_id)
377 if not self._users[document_id]:
378 del self._users[document_id]
380 await self.clear_document(document_id)
382 async def clear_document(self, document_id: str):
383 document_id = document_id.replace(':', '_')
385 if self._redis:
386 redis_key = f'{self._redis_key_prefix}:{document_id}:updates'
387 await self._redis.delete(redis_key)
388 redis_users_key = f'{self._redis_key_prefix}:{document_id}:users'
389 await self._redis.delete(redis_users_key)
390 else:
391 if document_id in self._updates:
392 del self._updates[document_id]
393 if document_id in self._users:
394 del self._users[document_id]