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
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
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
11from redis.exceptions import DataError
12from redis.typing import AbsExpiryT, EncodableT, ExpiryT
14if TYPE_CHECKING:
15 from redis.client import Redis
17try:
18 import hiredis # noqa
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
30try:
31 import ssl # noqa
33 SSL_AVAILABLE = True
34except ImportError:
35 SSL_AVAILABLE = False
37try:
38 import cryptography # noqa
40 CRYPTOGRAPHY_AVAILABLE = True
41except ImportError:
42 CRYPTOGRAPHY_AVAILABLE = False
44from importlib import metadata
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()
52def from_url(url: str, **kwargs: Any) -> "Redis":
53 """
54 Returns an active Redis client generated from the given database URL.
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
61 return Redis.from_url(url, **kwargs)
64@contextmanager
65def pipeline(redis_obj):
66 p = redis_obj.pipeline()
67 try:
68 yield p
69 p.execute()
70 finally:
71 p.reset()
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 )
80def safe_str(value):
81 return str(str_if_bytes(value))
84def decode_field_value(value, key=None, field_encodings=None):
85 """Decode a field value respecting optional per-field encoding overrides.
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)
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 = {}
109 for d in dicts:
110 merged.update(d)
112 return merged
115def list_keys_to_dict(key_list, callback):
116 return dict.fromkeys(key_list, callback)
119def merge_result(command, res):
120 """
121 Merge all items in `res` into a list.
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.
126 res : 'dict'
127 """
128 result = set()
130 for v in res.values():
131 for value in v:
132 result.add(value)
134 return list(result)
137def warn_deprecated(name, reason="", version="", stacklevel=2):
138 import warnings
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)
148def deprecated_function(reason="", version="", name=None):
149 """
150 Decorator to mark a function as deprecated.
151 """
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)
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)
169 return wrapper
171 return decorator
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
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)
194C = TypeVar("C", bound=Callable)
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
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 = []
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 )
245 def decorator(func: C) -> C:
246 if inspect.iscoroutinefunction(func):
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)
254 return async_wrapper
255 else:
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)
263 return wrapper
265 return decorator
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)
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
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
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
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 )
323def compare_versions(version1: str, version2: str) -> int:
324 """
325 Compare two versions.
327 :return: -1 if version1 > version2
328 0 if both versions are equal
329 1 if version1 < version2
330 """
332 num_versions1 = list(map(int, version1.split(".")))
333 num_versions2 = list(map(int, version2.split(".")))
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)
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
350 return 0
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")
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])
398 return exp_options
401def truncate_text(txt, max_length=100):
402 return textwrap.shorten(
403 text=txt, width=max_length, placeholder="...", break_long_words=True
404 )
407def dummy_fail():
408 """
409 Fake function for a Retry object if you don't need to handle each failure.
410 """
411 pass
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
421def experimental(cls):
422 """
423 Decorator to mark a class as experimental.
424 """
425 original_init = cls.__init__
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)
436 cls.__init__ = new_init
437 return cls
440def warn_experimental(name, stacklevel=2):
441 import warnings
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)
451def experimental_method() -> Callable[[C], C]:
452 """
453 Decorator to mark a function as experimental.
454 """
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)
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)
472 return wrapper
474 return decorator
477def warn_experimental_arg_usage(
478 arg_name: Union[list, str],
479 function_name: str,
480 stacklevel: int = 2,
481):
482 import warnings
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)
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 = ["*"]
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)
511 def decorator(func: C) -> C:
512 if inspect.iscoroutinefunction(func):
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)
521 return async_wrapper
522 else:
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)
531 return wrapper
533 return decorator