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

276 statements  

1import datetime 

2import inspect 

3import logging 

4import textwrap 

5import warnings 

6from collections.abc import Callable 

7from contextlib import contextmanager 

8from functools import 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 

302def get_lib_version(): 

303 try: 

304 libver = metadata.version("redis") 

305 except metadata.PackageNotFoundError: 

306 libver = "99.99.99" 

307 return libver 

308 

309 

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

311 if not exception.args: 

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

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

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

315 else: 

316 return ( 

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

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

319 ) 

320 

321 

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

323 """ 

324 Compare two versions. 

325 

326 :return: -1 if version1 > version2 

327 0 if both versions are equal 

328 1 if version1 < version2 

329 """ 

330 

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

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

333 

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

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

336 for _ in range(diff): 

337 num_versions2.append(0) 

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

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

340 for _ in range(diff): 

341 num_versions1.append(0) 

342 

343 for i, ver in enumerate(num_versions1): 

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

345 return -1 

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

347 return 1 

348 

349 return 0 

350 

351 

352def ensure_string(key): 

353 if isinstance(key, bytes): 

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

355 elif isinstance(key, str): 

356 return key 

357 else: 

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

359 

360 

361def extract_expire_flags( 

362 ex: Optional[ExpiryT] = None, 

363 px: Optional[ExpiryT] = None, 

364 exat: Optional[AbsExpiryT] = None, 

365 pxat: Optional[AbsExpiryT] = None, 

366) -> List[EncodableT]: 

367 exp_options: list[EncodableT] = [] 

368 if ex is not None: 

369 exp_options.append("EX") 

370 if isinstance(ex, datetime.timedelta): 

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

372 elif isinstance(ex, int): 

373 exp_options.append(ex) 

374 elif isinstance(ex, str) and ex.isdigit(): 

375 exp_options.append(int(ex)) 

376 else: 

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

378 elif px is not None: 

379 exp_options.append("PX") 

380 if isinstance(px, datetime.timedelta): 

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

382 elif isinstance(px, int): 

383 exp_options.append(px) 

384 else: 

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

386 elif exat is not None: 

387 if isinstance(exat, datetime.datetime): 

388 exat = int(exat.timestamp()) 

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

390 elif pxat is not None: 

391 if isinstance(pxat, datetime.datetime): 

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

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

394 

395 return exp_options 

396 

397 

398def truncate_text(txt, max_length=100): 

399 return textwrap.shorten( 

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

401 ) 

402 

403 

404def dummy_fail(): 

405 """ 

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

407 """ 

408 pass 

409 

410 

411async def dummy_fail_async(): 

412 """ 

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

414 """ 

415 pass 

416 

417 

418def experimental(cls): 

419 """ 

420 Decorator to mark a class as experimental. 

421 """ 

422 original_init = cls.__init__ 

423 

424 @wraps(original_init) 

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

426 warnings.warn( 

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

428 category=UserWarning, 

429 stacklevel=2, 

430 ) 

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

432 

433 cls.__init__ = new_init 

434 return cls 

435 

436 

437def warn_experimental(name, stacklevel=2): 

438 import warnings 

439 

440 msg = ( 

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

442 "Be aware that the function arguments can " 

443 "change or be removed in future versions." 

444 ) 

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

446 

447 

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

449 """ 

450 Decorator to mark a function as experimental. 

451 """ 

452 

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

454 if inspect.iscoroutinefunction(func): 

455 # Create async wrapper for async functions 

456 @wraps(func) 

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

458 warn_experimental(func.__name__, stacklevel=2) 

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

460 

461 return async_wrapper 

462 else: 

463 # Create regular wrapper for sync functions 

464 @wraps(func) 

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

466 warn_experimental(func.__name__, stacklevel=2) 

467 return func(*args, **kwargs) 

468 

469 return wrapper 

470 

471 return decorator 

472 

473 

474def warn_experimental_arg_usage( 

475 arg_name: Union[list, str], 

476 function_name: str, 

477 stacklevel: int = 2, 

478): 

479 import warnings 

480 

481 msg = ( 

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

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

484 ) 

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

486 

487 

488def experimental_args( 

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

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

491 """ 

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

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

494 """ 

495 if args_to_warn is None: 

496 args_to_warn = ["*"] 

497 

498 def _check_experimental_args(func, filterable_args): 

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

500 for arg in args_to_warn: 

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

502 warn_experimental_arg_usage( 

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

504 ) 

505 elif arg in filterable_args: 

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

507 

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

509 if inspect.iscoroutinefunction(func): 

510 

511 @wraps(func) 

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

513 filterable_args = _get_filterable_args(func, args, kwargs) 

514 if len(filterable_args) > 0: 

515 _check_experimental_args(func, filterable_args) 

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

517 

518 return async_wrapper 

519 else: 

520 

521 @wraps(func) 

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

523 filterable_args = _get_filterable_args(func, args, kwargs) 

524 if len(filterable_args) > 0: 

525 _check_experimental_args(func, filterable_args) 

526 return func(*args, **kwargs) 

527 

528 return wrapper 

529 

530 return decorator