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

1"""Redis-backed distributed data structures for WebSocket state management.""" 

2 

3from __future__ import annotations 

4 

5import hashlib 

6import logging 

7import uuid 

8 

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 

14 

15log = logging.getLogger(__name__) 

16 

17YDOC_KEY_PREFIX = f'{REDIS_KEY_PREFIX}:ydoc:documents' 

18SCAN_BATCH_SIZE = 200 

19 

20 

21class RedisLock: 

22 """Distributed lock backed by a Redis SET with NX/EX semantics.""" 

23 

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

36 

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 ) 

55 

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 

60 

61 def renew_lock(self): 

62 return bool(self.redis.eval(self._RENEW_SCRIPT, 1, self.lock_name, self.lock_id, self.timeout_secs)) 

63 

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) 

69 

70 

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 ) 

88 

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) 

94 

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) 

100 

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) 

107 

108 def __contains__(self, key): 

109 return self.redis.hexists(self.name, key) 

110 

111 def __len__(self): 

112 return self.redis.hlen(self.name) 

113 

114 def keys(self): 

115 return self.redis.hkeys(self.name) 

116 

117 def values(self): 

118 return [JSONCodec.loads(v) for v in self.redis.hvals(self.name)] 

119 

120 def items(self): 

121 return [(k, JSONCodec.loads(v)) for k, v in self.redis.hgetall(self.name).items()] 

122 

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 

132 

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) 

139 

140 def set(self, mapping: dict): 

141 if not mapping: 

142 self.clear() 

143 return 

144 

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() 

154 

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) 

161 

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 

167 

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) 

174 

175 if self._signature_name: 

176 self.redis.set(self._signature_name, f'{content_digest}:{uuid.uuid4().hex}') 

177 

178 def get(self, key, default=None): 

179 try: 

180 return self[key] 

181 except KeyError: 

182 return default 

183 

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) 

190 

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 

197 

198 def setdefault(self, key, default=None): 

199 if key not in self: 

200 self[key] = default 

201 return self[key] 

202 

203 

204class CachedRedisDict(RedisDict): 

205 """Answers reads from a per-worker cache of the hash, refetched whenever its signature changes.""" 

206 

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 

211 

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 

218 

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) 

224 

225 def __contains__(self, key): 

226 return key in self._refresh_cache() 

227 

228 def __len__(self): 

229 return len(self._refresh_cache()) 

230 

231 def keys(self): 

232 return list(self._refresh_cache().keys()) 

233 

234 def values(self): 

235 return [JSONCodec.loads(v) for v in self._refresh_cache().values()] 

236 

237 def items(self): 

238 return [(k, JSONCodec.loads(v)) for k, v in self._refresh_cache().items()] 

239 

240 

241class YdocManager: 

242 COMPACTION_THRESHOLD = 500 

243 

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 

253 

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) 

268 

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() 

284 

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

295 

296 async def get_updates(self, document_id: str) -> list[bytes]: 

297 document_id = document_id.replace(':', '_') 

298 

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, []) 

305 

306 async def document_exists(self, document_id: str) -> bool: 

307 document_id = document_id.replace(':', '_') 

308 

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 

314 

315 async def get_users(self, document_id: str) -> list[str]: 

316 document_id = document_id.replace(':', '_') 

317 

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, []) 

324 

325 async def add_user(self, document_id: str, user_id: str): 

326 document_id = document_id.replace(':', '_') 

327 

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) 

340 

341 async def remove_user(self, document_id: str, user_id: str): 

342 document_id = document_id.replace(':', '_') 

343 

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) 

353 

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) 

362 

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) 

366 

367 if len(await self.get_users(document_id)) == 0: 

368 await self.clear_document(document_id) 

369 

370 # Clean up the reverse index itself. 

371 await self._redis.delete(session_key) 

372 

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] 

379 

380 await self.clear_document(document_id) 

381 

382 async def clear_document(self, document_id: str): 

383 document_id = document_id.replace(':', '_') 

384 

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]