Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/redis/utils.py: 39%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

279 statements  

1import datetime 

2import inspect 

3import logging 

4import textwrap 

5import warnings 

6from collections.abc import Callable 

7from contextlib import contextmanager 

8from functools import cache, wraps 

9from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, TypeVar, Union 

10 

11from redis.exceptions import DataError 

12from redis.typing import AbsExpiryT, EncodableT, ExpiryT 

13 

14if TYPE_CHECKING: 

15 from redis.client import Redis 

16 

17try: 

18 import hiredis # noqa 

19 

20 # Only support Hiredis >= 3.0: 

21 hiredis_version = hiredis.__version__.split(".") 

22 HIREDIS_AVAILABLE = int(hiredis_version[0]) > 3 or ( 

23 int(hiredis_version[0]) == 3 and int(hiredis_version[1]) >= 2 

24 ) 

25 if not HIREDIS_AVAILABLE: 

26 raise ImportError("hiredis package should be >= 3.2.0") 

27except ImportError: 

28 HIREDIS_AVAILABLE = False 

29 

30try: 

31 import ssl # noqa 

32 

33 SSL_AVAILABLE = True 

34except ImportError: 

35 SSL_AVAILABLE = False 

36 

37try: 

38 import cryptography # noqa 

39 

40 CRYPTOGRAPHY_AVAILABLE = True 

41except ImportError: 

42 CRYPTOGRAPHY_AVAILABLE = False 

43 

44from importlib import metadata 

45 

46# Shared marker for omitted arguments, especially where None is a valid 

47# explicit value. Import this object from redis.utils instead of creating local 

48# sentinels, and compare it by identity only (`is` / `is not`). 

49SENTINEL = object() 

50 

51 

52def from_url(url: str, **kwargs: Any) -> "Redis": 

53 """ 

54 Returns an active Redis client generated from the given database URL. 

55 

56 Will attempt to extract the database id from the path url fragment, if 

57 none is provided. 

58 """ 

59 from redis.client import Redis 

60 

61 return Redis.from_url(url, **kwargs) 

62 

63 

64@contextmanager 

65def pipeline(redis_obj): 

66 p = redis_obj.pipeline() 

67 try: 

68 yield p 

69 p.execute() 

70 finally: 

71 p.reset() 

72 

73 

74def str_if_bytes(value: Union[str, bytes]) -> str: 

75 return ( 

76 value.decode("utf-8", errors="replace") if isinstance(value, bytes) else value 

77 ) 

78 

79 

80def safe_str(value): 

81 return str(str_if_bytes(value)) 

82 

83 

84def decode_field_value(value, key=None, field_encodings=None): 

85 """Decode a field value respecting optional per-field encoding overrides. 

86 

87 - If *field_encodings* is provided and *key* is in it, the corresponding 

88 encoding is used (``None`` means keep raw bytes). 

89 - Otherwise falls back to :func:`str_if_bytes`. 

90 """ 

91 if not isinstance(value, bytes): 

92 return value 

93 if field_encodings and key is not None and key in field_encodings: 

94 encoding = field_encodings[key] 

95 if encoding is None: 

96 return value 

97 return value.decode(encoding, "replace") 

98 return str_if_bytes(value) 

99 

100 

101def dict_merge(*dicts: Mapping[str, Any]) -> Dict[str, Any]: 

102 """ 

103 Merge all provided dicts into 1 dict. 

104 *dicts : `dict` 

105 dictionaries to merge 

106 """ 

107 merged = {} 

108 

109 for d in dicts: 

110 merged.update(d) 

111 

112 return merged 

113 

114 

115def list_keys_to_dict(key_list, callback): 

116 return dict.fromkeys(key_list, callback) 

117 

118 

119def merge_result(command, res): 

120 """ 

121 Merge all items in `res` into a list. 

122 

123 This command is used when sending a command to multiple nodes 

124 and the result from each node should be merged into a single list. 

125 

126 res : 'dict' 

127 """ 

128 result = set() 

129 

130 for v in res.values(): 

131 for value in v: 

132 result.add(value) 

133 

134 return list(result) 

135 

136 

137def warn_deprecated(name, reason="", version="", stacklevel=2): 

138 import warnings 

139 

140 msg = f"Call to deprecated {name}." 

141 if reason: 

142 msg += f" ({reason})" 

143 if version: 

144 msg += f" -- Deprecated since version {version}." 

145 warnings.warn(msg, category=DeprecationWarning, stacklevel=stacklevel) 

146 

147 

148def deprecated_function(reason="", version="", name=None): 

149 """ 

150 Decorator to mark a function as deprecated. 

151 """ 

152 

153 def decorator(func): 

154 if inspect.iscoroutinefunction(func): 

155 # Create async wrapper for async functions 

156 @wraps(func) 

157 async def async_wrapper(*args, **kwargs): 

158 warn_deprecated(name or func.__name__, reason, version, stacklevel=3) 

159 return await func(*args, **kwargs) 

160 

161 return async_wrapper 

162 else: 

163 # Create regular wrapper for sync functions 

164 @wraps(func) 

165 def wrapper(*args, **kwargs): 

166 warn_deprecated(name or func.__name__, reason, version, stacklevel=3) 

167 return func(*args, **kwargs) 

168 

169 return wrapper 

170 

171 return decorator 

172 

173 

174def warn_deprecated_arg_usage( 

175 arg_name: Union[list, str], 

176 function_name: str, 

177 reason: str = "", 

178 version: str = "", 

179 stacklevel: int = 2, 

180): 

181 import warnings 

182 

183 msg = ( 

184 f"Call to '{function_name}' function with deprecated" 

185 f" usage of input argument/s '{arg_name}'." 

186 ) 

187 if reason: 

188 msg += f" ({reason})" 

189 if version: 

190 msg += f" -- Deprecated since version {version}." 

191 warnings.warn(msg, category=DeprecationWarning, stacklevel=stacklevel) 

192 

193 

194C = TypeVar("C", bound=Callable) 

195 

196 

197def _get_filterable_args( 

198 func: Callable, args: tuple, kwargs: dict, allowed_args: Optional[List[str]] = None 

199) -> dict: 

200 """ 

201 Extract arguments from function call that should be checked for deprecation/experimental warnings. 

202 Excludes 'self' and any explicitly allowed args. 

203 """ 

204 arg_names = func.__code__.co_varnames[: func.__code__.co_argcount] 

205 filterable_args = dict(zip(arg_names, args)) 

206 filterable_args.update(kwargs) 

207 filterable_args.pop("self", None) 

208 if allowed_args: 

209 for allowed_arg in allowed_args: 

210 filterable_args.pop(allowed_arg, None) 

211 return filterable_args 

212 

213 

214def deprecated_args( 

215 args_to_warn: Optional[List[str]] = None, 

216 allowed_args: Optional[List[str]] = None, 

217 reason: str = "", 

218 version: str = "", 

219) -> Callable[[C], C]: 

220 """ 

221 Decorator to mark specified args of a function as deprecated. 

222 If '*' is in args_to_warn, all arguments will be marked as deprecated. 

223 """ 

224 if args_to_warn is None: 

225 args_to_warn = ["*"] 

226 if allowed_args is None: 

227 allowed_args = [] 

228 

229 def _check_deprecated_args(func, filterable_args): 

230 """Check and warn about deprecated arguments.""" 

231 for arg in args_to_warn: 

232 if arg == "*" and len(filterable_args) > 0: 

233 warn_deprecated_arg_usage( 

234 list(filterable_args.keys()), 

235 func.__name__, 

236 reason, 

237 version, 

238 stacklevel=5, 

239 ) 

240 elif arg in filterable_args: 

241 warn_deprecated_arg_usage( 

242 arg, func.__name__, reason, version, stacklevel=5 

243 ) 

244 

245 def decorator(func: C) -> C: 

246 if inspect.iscoroutinefunction(func): 

247 

248 @wraps(func) 

249 async def async_wrapper(*args, **kwargs): 

250 filterable_args = _get_filterable_args(func, args, kwargs, allowed_args) 

251 _check_deprecated_args(func, filterable_args) 

252 return await func(*args, **kwargs) 

253 

254 return async_wrapper 

255 else: 

256 

257 @wraps(func) 

258 def wrapper(*args, **kwargs): 

259 filterable_args = _get_filterable_args(func, args, kwargs, allowed_args) 

260 _check_deprecated_args(func, filterable_args) 

261 return func(*args, **kwargs) 

262 

263 return wrapper 

264 

265 return decorator 

266 

267 

268def _set_info_logger(): 

269 """ 

270 Set up a logger that log info logs to stdout. 

271 (This is used by the default push response handler) 

272 """ 

273 if "push_response" not in logging.root.manager.loggerDict.keys(): 

274 logger = logging.getLogger("push_response") 

275 logger.setLevel(logging.INFO) 

276 handler = logging.StreamHandler() 

277 handler.setLevel(logging.INFO) 

278 logger.addHandler(handler) 

279 

280 

281#: Default RESP protocol version used on the wire when the user does not 

282#: supply an explicit ``protocol`` to the client / connection / pool. Lives 

283#: in ``redis.utils`` so both ``redis.connection`` (for the HELLO handshake) 

284#: and ``check_protocol_version`` (for protocol-gated features) can read it 

285#: without a circular import. 

286DEFAULT_RESP_VERSION = 3 

287 

288 

289def check_protocol_version( 

290 protocol: str | int | object | None, expected_version: int = 3 

291) -> bool: 

292 if protocol is None or protocol is SENTINEL: 

293 protocol = DEFAULT_RESP_VERSION 

294 if isinstance(protocol, str): 

295 try: 

296 protocol = int(protocol) 

297 except ValueError: 

298 return False 

299 return protocol == expected_version 

300 

301 

302@cache 

303def get_lib_version(): 

304 try: 

305 libver = metadata.version("redis") 

306 except metadata.PackageNotFoundError: 

307 libver = "99.99.99" 

308 return libver 

309 

310 

311def format_error_message(host_error: str, exception: BaseException) -> str: 

312 if not exception.args: 

313 return f"Error connecting to {host_error}." 

314 elif len(exception.args) == 1: 

315 return f"Error {exception.args[0]} connecting to {host_error}." 

316 else: 

317 return ( 

318 f"Error {exception.args[0]} connecting to {host_error}. " 

319 f"{exception.args[1]}." 

320 ) 

321 

322 

323def compare_versions(version1: str, version2: str) -> int: 

324 """ 

325 Compare two versions. 

326 

327 :return: -1 if version1 > version2 

328 0 if both versions are equal 

329 1 if version1 < version2 

330 """ 

331 

332 num_versions1 = list(map(int, version1.split("."))) 

333 num_versions2 = list(map(int, version2.split("."))) 

334 

335 if len(num_versions1) > len(num_versions2): 

336 diff = len(num_versions1) - len(num_versions2) 

337 for _ in range(diff): 

338 num_versions2.append(0) 

339 elif len(num_versions1) < len(num_versions2): 

340 diff = len(num_versions2) - len(num_versions1) 

341 for _ in range(diff): 

342 num_versions1.append(0) 

343 

344 for i, ver in enumerate(num_versions1): 

345 if num_versions1[i] > num_versions2[i]: 

346 return -1 

347 elif num_versions1[i] < num_versions2[i]: 

348 return 1 

349 

350 return 0 

351 

352 

353def ensure_string(key): 

354 if isinstance(key, bytes): 

355 return key.decode("utf-8") 

356 elif isinstance(key, str): 

357 return key 

358 else: 

359 raise TypeError("Key must be either a string or bytes") 

360 

361 

362def extract_expire_flags( 

363 ex: ExpiryT | str | None = None, 

364 px: ExpiryT | str | None = None, 

365 exat: Optional[AbsExpiryT] = None, 

366 pxat: Optional[AbsExpiryT] = None, 

367) -> List[EncodableT]: 

368 exp_options: list[EncodableT] = [] 

369 if ex is not None: 

370 exp_options.append("EX") 

371 if isinstance(ex, datetime.timedelta): 

372 exp_options.append(int(ex.total_seconds())) 

373 elif isinstance(ex, int): 

374 exp_options.append(ex) 

375 elif isinstance(ex, str) and ex.isdecimal(): 

376 exp_options.append(int(ex)) 

377 else: 

378 raise DataError("ex must be datetime.timedelta or int") 

379 elif px is not None: 

380 exp_options.append("PX") 

381 if isinstance(px, datetime.timedelta): 

382 exp_options.append(int(px.total_seconds() * 1000)) 

383 elif isinstance(px, int): 

384 exp_options.append(px) 

385 elif isinstance(px, str) and px.isdecimal(): 

386 exp_options.append(int(px)) 

387 else: 

388 raise DataError("px must be datetime.timedelta or int") 

389 elif exat is not None: 

390 if isinstance(exat, datetime.datetime): 

391 exat = int(exat.timestamp()) 

392 exp_options.extend(["EXAT", exat]) 

393 elif pxat is not None: 

394 if isinstance(pxat, datetime.datetime): 

395 pxat = int(pxat.timestamp() * 1000) 

396 exp_options.extend(["PXAT", pxat]) 

397 

398 return exp_options 

399 

400 

401def truncate_text(txt, max_length=100): 

402 return textwrap.shorten( 

403 text=txt, width=max_length, placeholder="...", break_long_words=True 

404 ) 

405 

406 

407def dummy_fail(): 

408 """ 

409 Fake function for a Retry object if you don't need to handle each failure. 

410 """ 

411 pass 

412 

413 

414async def dummy_fail_async(): 

415 """ 

416 Async fake function for a Retry object if you don't need to handle each failure. 

417 """ 

418 pass 

419 

420 

421def experimental(cls): 

422 """ 

423 Decorator to mark a class as experimental. 

424 """ 

425 original_init = cls.__init__ 

426 

427 @wraps(original_init) 

428 def new_init(self, *args, **kwargs): 

429 warnings.warn( 

430 f"{cls.__name__} is an experimental and may change or be removed in future versions.", 

431 category=UserWarning, 

432 stacklevel=2, 

433 ) 

434 original_init(self, *args, **kwargs) 

435 

436 cls.__init__ = new_init 

437 return cls 

438 

439 

440def warn_experimental(name, stacklevel=2): 

441 import warnings 

442 

443 msg = ( 

444 f"Call to experimental method {name}. " 

445 "Be aware that the function arguments can " 

446 "change or be removed in future versions." 

447 ) 

448 warnings.warn(msg, category=UserWarning, stacklevel=stacklevel) 

449 

450 

451def experimental_method() -> Callable[[C], C]: 

452 """ 

453 Decorator to mark a function as experimental. 

454 """ 

455 

456 def decorator(func: C) -> C: 

457 if inspect.iscoroutinefunction(func): 

458 # Create async wrapper for async functions 

459 @wraps(func) 

460 async def async_wrapper(*args, **kwargs): 

461 warn_experimental(func.__name__, stacklevel=2) 

462 return await func(*args, **kwargs) 

463 

464 return async_wrapper 

465 else: 

466 # Create regular wrapper for sync functions 

467 @wraps(func) 

468 def wrapper(*args, **kwargs): 

469 warn_experimental(func.__name__, stacklevel=2) 

470 return func(*args, **kwargs) 

471 

472 return wrapper 

473 

474 return decorator 

475 

476 

477def warn_experimental_arg_usage( 

478 arg_name: Union[list, str], 

479 function_name: str, 

480 stacklevel: int = 2, 

481): 

482 import warnings 

483 

484 msg = ( 

485 f"Call to '{function_name}' method with experimental" 

486 f" usage of input argument/s '{arg_name}'." 

487 ) 

488 warnings.warn(msg, category=UserWarning, stacklevel=stacklevel) 

489 

490 

491def experimental_args( 

492 args_to_warn: Optional[List[str]] = None, 

493) -> Callable[[C], C]: 

494 """ 

495 Decorator to mark specified args of a function as experimental. 

496 If '*' is in args_to_warn, all arguments will be marked as experimental. 

497 """ 

498 if args_to_warn is None: 

499 args_to_warn = ["*"] 

500 

501 def _check_experimental_args(func, filterable_args): 

502 """Check and warn about experimental arguments.""" 

503 for arg in args_to_warn: 

504 if arg == "*" and len(filterable_args) > 0: 

505 warn_experimental_arg_usage( 

506 list(filterable_args.keys()), func.__name__, stacklevel=4 

507 ) 

508 elif arg in filterable_args: 

509 warn_experimental_arg_usage(arg, func.__name__, stacklevel=4) 

510 

511 def decorator(func: C) -> C: 

512 if inspect.iscoroutinefunction(func): 

513 

514 @wraps(func) 

515 async def async_wrapper(*args, **kwargs): 

516 filterable_args = _get_filterable_args(func, args, kwargs) 

517 if len(filterable_args) > 0: 

518 _check_experimental_args(func, filterable_args) 

519 return await func(*args, **kwargs) 

520 

521 return async_wrapper 

522 else: 

523 

524 @wraps(func) 

525 def wrapper(*args, **kwargs): 

526 filterable_args = _get_filterable_args(func, args, kwargs) 

527 if len(filterable_args) > 0: 

528 _check_experimental_args(func, filterable_args) 

529 return func(*args, **kwargs) 

530 

531 return wrapper 

532 

533 return decorator