Coverage for open_webui/utils/redis.py: 18%

164 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1"""Redis connection utilities. 

2 

3Provides connection factory functions for standalone, Sentinel, and Cluster 

4Redis deployments, with optional async support and automatic connection caching. 

5""" 

6 

7from __future__ import annotations 

8 

9import asyncio 

10import inspect 

11import logging 

12import time 

13from typing import Any 

14from urllib.parse import ParseResult, urlparse 

15 

16import redis as _redis_sync 

17from open_webui.env import ( 

18 REDIS_CLUSTER, 

19 REDIS_HEALTH_CHECK_INTERVAL, 

20 REDIS_RECONNECT_DELAY, 

21 REDIS_SENTINEL_HOSTS, 

22 REDIS_SENTINEL_MAX_RETRY_COUNT, 

23 REDIS_SENTINEL_PORT, 

24 REDIS_SOCKET_CONNECT_TIMEOUT, 

25 REDIS_SOCKET_KEEPALIVE, 

26 REDIS_SOCKET_TIMEOUT, 

27 REDIS_URL, 

28) 

29 

30log = logging.getLogger(__name__) 

31 

32_ACCEPTED_SCHEMES = frozenset({'redis', 'rediss'}) 

33_SENTINEL_RETRYABLE = ( 

34 _redis_sync.exceptions.ConnectionError, 

35 _redis_sync.exceptions.ReadOnlyError, 

36 _redis_sync.exceptions.TimeoutError, 

37) 

38_FACTORY_METHODS = frozenset({'pipeline', 'pubsub', 'monitor', 'client', 'transaction'}) 

39_CONNECTION_POOL: dict[tuple, Any] = {} 

40 

41 

42def parse_redis_url(url: str) -> dict[str, Any]: 

43 """Break a ``redis://`` URL into its parts: service, port, db, username, password.""" 

44 parts: ParseResult = urlparse(url) 

45 if parts.scheme not in _ACCEPTED_SCHEMES: 

46 raise ValueError(f"Invalid Redis URL scheme '{parts.scheme}'; expected 'redis' or 'rediss'.") 

47 return { 

48 'service': parts.hostname or 'mymaster', 

49 'port': parts.port or 6379, 

50 'db': int(parts.path.lstrip('/') or '0'), 

51 'username': parts.username or None, 

52 'password': parts.password or None, 

53 } 

54 

55 

56parse_redis_service_url = parse_redis_url 

57 

58 

59def get_sentinels_from_env( 

60 hosts_csv: str | None, 

61 port: str | int | None, 

62) -> list[tuple[str, int]]: 

63 """Turn a comma-separated host string into ``[(host, port), …]``.""" 

64 if not hosts_csv: 64 ↛ 66line 64 didn't jump to line 66 because the condition on line 64 was always true

65 return [] 

66 resolved_port = int(port) if port else 26379 

67 return [(host.strip(), resolved_port) for host in hosts_csv.split(',') if host.strip()] 

68 

69 

70def build_sentinel_url( 

71 base_url: str, 

72 hosts_csv: str, 

73 port: str | int, 

74) -> str: 

75 """Construct a ``redis+sentinel://`` connection string. 

76 

77 ``base_url`` supplies credentials, db index, and master service name. 

78 ``hosts_csv`` is a comma-separated list of sentinel hostnames. 

79 """ 

80 cfg = parse_redis_url(base_url) 

81 auth = '' 

82 if cfg['username'] or cfg['password']: 

83 auth = f'{cfg["username"] or ""}:{cfg["password"] or ""}@' 

84 nodes = ','.join(f'{host.strip()}:{port}' for host in hosts_csv.split(',') if host.strip()) 

85 return f'redis+sentinel://{auth}{nodes}/{cfg["db"]}/{cfg["service"]}' 

86 

87 

88def get_redis_client(async_mode: bool = False) -> Any | None: 

89 """Create a Redis connection using settings from environment variables. 

90 

91 Returns ``None`` when Redis is not configured or the connection fails. 

92 """ 

93 sentinel_list = get_sentinels_from_env(REDIS_SENTINEL_HOSTS, REDIS_SENTINEL_PORT) 

94 if not REDIS_URL and not sentinel_list: 94 ↛ 96line 94 didn't jump to line 96 because the condition on line 94 was always true

95 return None 

96 try: 

97 return get_redis_connection( 

98 REDIS_URL, 

99 redis_sentinels=sentinel_list, 

100 redis_cluster=REDIS_CLUSTER, 

101 async_mode=async_mode, 

102 ) 

103 except Exception: 

104 log.debug('Could not establish Redis connection', exc_info=True) 

105 return None 

106 

107 

108# --------------------------------------------------------------------------- 

109# Sentinel proxy with automatic failover retry 

110# --------------------------------------------------------------------------- 

111 

112 

113class SentinelRedisProxy: 

114 """Transparent proxy that re-resolves the Sentinel master on connection errors. 

115 

116 Every call (sync or async) is wrapped with retry logic so that transient 

117 Sentinel failovers are handled without caller intervention. 

118 """ 

119 

120 def __init__( 

121 self, 

122 sentinel: Any, 

123 service_name: str, 

124 *, 

125 async_mode: bool = True, 

126 ) -> None: 

127 self._sentinel = sentinel 

128 self._service_name = service_name 

129 self._async_mode = async_mode 

130 self._master: Any | None = None 

131 

132 def __getattr__(self, name: str) -> Any: 

133 """Proxy attribute access with automatic Sentinel failover retry.""" 

134 current_master = self._resolve_master() 

135 original = getattr(current_master, name) 

136 

137 # Non-callable or factory attributes pass through without wrapping. 

138 if not callable(original) or name in _FACTORY_METHODS: 

139 return original 

140 

141 # Select the retry wrapper matching the execution mode. 

142 if not self._async_mode: 

143 return self._wrap_sync(name) 

144 return self._wrap_async(name, original) 

145 

146 def _resolve_master(self) -> Any: 

147 """Ask Sentinel for the current master connection.""" 

148 if self._master is None: 

149 self._master = self._sentinel.master_for(self._service_name) 

150 return self._master 

151 

152 def _clear_master(self) -> None: 

153 self._master = None 

154 

155 def _should_retry(self, attempt: int) -> bool: 

156 return attempt < REDIS_SENTINEL_MAX_RETRY_COUNT - 1 

157 

158 def _log_retry(self, exc: Exception, attempt: int) -> None: 

159 log.debug( 

160 'Sentinel failover (%s) — retry %d/%d', 

161 type(exc).__name__, 

162 attempt + 1, 

163 REDIS_SENTINEL_MAX_RETRY_COUNT, 

164 ) 

165 

166 def _log_exhausted(self, exc: Exception) -> None: 

167 log.error( 

168 'Redis operation failed after %d retries: %s', 

169 REDIS_SENTINEL_MAX_RETRY_COUNT, 

170 exc, 

171 ) 

172 

173 # -- async wrappers ----------------------------------------------------- 

174 

175 def _wrap_async(self, name: str, attr: Any) -> Any: 

176 if inspect.isasyncgenfunction(attr): 

177 return self._wrap_async_gen(name) 

178 return self._wrap_async_call(name) 

179 

180 def _wrap_async_gen(self, name: str) -> Any: 

181 proxy = self 

182 

183 def wrapper(*args: Any, **kwargs: Any) -> Any: 

184 async def _inner(): 

185 for attempt in range(REDIS_SENTINEL_MAX_RETRY_COUNT): 

186 try: 

187 method = getattr(proxy._resolve_master(), name) 

188 async for value in method(*args, **kwargs): 

189 yield value 

190 return 

191 except _SENTINEL_RETRYABLE as exc: 

192 if proxy._should_retry(attempt): 

193 proxy._log_retry(exc, attempt) 

194 proxy._clear_master() 

195 if REDIS_RECONNECT_DELAY: 

196 await asyncio.sleep(REDIS_RECONNECT_DELAY / 1000) 

197 continue 

198 proxy._log_exhausted(exc) 

199 raise 

200 

201 return _inner() 

202 

203 return wrapper 

204 

205 def _wrap_async_call(self, name: str) -> Any: 

206 proxy = self 

207 

208 async def wrapper(*args: Any, **kwargs: Any) -> Any: 

209 for attempt in range(REDIS_SENTINEL_MAX_RETRY_COUNT): 

210 try: 

211 method = getattr(proxy._resolve_master(), name) 

212 result = method(*args, **kwargs) 

213 if inspect.iscoroutine(result): 

214 return await result 

215 return result 

216 except _SENTINEL_RETRYABLE as exc: 

217 if proxy._should_retry(attempt): 

218 proxy._log_retry(exc, attempt) 

219 proxy._clear_master() 

220 if REDIS_RECONNECT_DELAY: 

221 await asyncio.sleep(REDIS_RECONNECT_DELAY / 1000) 

222 continue 

223 proxy._log_exhausted(exc) 

224 raise 

225 

226 return wrapper 

227 

228 # -- sync wrapper ------------------------------------------------------- 

229 

230 def _wrap_sync(self, name: str) -> Any: 

231 proxy = self 

232 

233 def wrapper(*args: Any, **kwargs: Any) -> Any: 

234 for attempt in range(REDIS_SENTINEL_MAX_RETRY_COUNT): 

235 try: 

236 method = getattr(proxy._resolve_master(), name) 

237 return method(*args, **kwargs) 

238 except _SENTINEL_RETRYABLE as exc: 

239 if proxy._should_retry(attempt): 

240 proxy._log_retry(exc, attempt) 

241 proxy._clear_master() 

242 if REDIS_RECONNECT_DELAY: 

243 time.sleep(REDIS_RECONNECT_DELAY / 1000) 

244 continue 

245 proxy._log_exhausted(exc) 

246 raise 

247 

248 return wrapper 

249 

250 

251# --------------------------------------------------------------------------- 

252# Connection factory 

253# --------------------------------------------------------------------------- 

254 

255 

256def _socket_options() -> dict[str, Any]: 

257 """Collect optional socket-level kwargs once instead of repeating them.""" 

258 opts: dict[str, Any] = {} 

259 if REDIS_SOCKET_CONNECT_TIMEOUT is not None: 

260 opts['socket_connect_timeout'] = REDIS_SOCKET_CONNECT_TIMEOUT 

261 if REDIS_SOCKET_TIMEOUT: 

262 opts['socket_timeout'] = REDIS_SOCKET_TIMEOUT 

263 if REDIS_SOCKET_KEEPALIVE: 

264 opts['socket_keepalive'] = True 

265 if REDIS_HEALTH_CHECK_INTERVAL: 

266 opts['health_check_interval'] = REDIS_HEALTH_CHECK_INTERVAL 

267 return opts 

268 

269 

270def _build_sentinel( 

271 redis_module: Any, 

272 url: str, 

273 sentinels: list[tuple[str, int]], 

274 decode_responses: bool, 

275 async_mode: bool, 

276) -> SentinelRedisProxy: 

277 """Create a SentinelRedisProxy from a redis URL and sentinel list.""" 

278 cfg = parse_redis_url(url) 

279 sentinel = redis_module.sentinel.Sentinel( 

280 sentinels, 

281 port=cfg['port'], 

282 db=cfg['db'], 

283 username=cfg['username'], 

284 password=cfg['password'], 

285 decode_responses=decode_responses, 

286 socket_connect_timeout=REDIS_SOCKET_CONNECT_TIMEOUT, 

287 **{k: v for k, v in _socket_options().items() if k != 'socket_connect_timeout'}, 

288 ) 

289 return SentinelRedisProxy(sentinel, cfg['service'], async_mode=async_mode) 

290 

291 

292def get_redis_connection( 

293 redis_url: str | None, 

294 redis_sentinels: list[tuple[str, int]] | None = None, 

295 redis_cluster: bool = False, 

296 async_mode: bool = False, 

297 decode_responses: bool = True, 

298) -> Any | None: 

299 """Return a cached Redis connection (or create one). 

300 

301 Supports three topologies in order of precedence: 

302 1. **Sentinel** — when ``redis_sentinels`` is non-empty. 

303 2. **Cluster** — when ``redis_cluster`` is True. 

304 3. **Standalone** — plain ``redis://`` connection. 

305 """ 

306 cache_key = ( 

307 redis_url, 

308 tuple(redis_sentinels) if redis_sentinels else (), 

309 redis_cluster, 

310 async_mode, 

311 decode_responses, 

312 ) 

313 if cache_key in _CONNECTION_POOL: 

314 return _CONNECTION_POOL[cache_key] 

315 

316 extra = _socket_options() 

317 connection: Any = None 

318 

319 # Pick the right redis module for sync vs async. 

320 if async_mode: 

321 import redis.asyncio as redis_mod 

322 else: 

323 import redis as redis_mod # type: ignore[no-redef] 

324 

325 if redis_sentinels: 

326 connection = _build_sentinel(redis_mod, redis_url, redis_sentinels, decode_responses, async_mode) 

327 elif redis_cluster: 

328 if not redis_url: 

329 raise ValueError('Redis URL is required for cluster mode.') 

330 connection = redis_mod.cluster.RedisCluster.from_url( 

331 redis_url, 

332 decode_responses=decode_responses, 

333 **extra, 

334 ) 

335 elif redis_url: 

336 factory = getattr(redis_mod, 'from_url', None) or redis_mod.Redis.from_url 

337 connection = factory(redis_url, decode_responses=decode_responses, **extra) 

338 

339 _CONNECTION_POOL[cache_key] = connection 

340 return connection