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 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
302def get_lib_version():
303 try:
304 libver = metadata.version("redis")
305 except metadata.PackageNotFoundError:
306 libver = "99.99.99"
307 return libver
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 )
322def compare_versions(version1: str, version2: str) -> int:
323 """
324 Compare two versions.
326 :return: -1 if version1 > version2
327 0 if both versions are equal
328 1 if version1 < version2
329 """
331 num_versions1 = list(map(int, version1.split(".")))
332 num_versions2 = list(map(int, version2.split(".")))
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)
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
349 return 0
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")
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])
395 return exp_options
398def truncate_text(txt, max_length=100):
399 return textwrap.shorten(
400 text=txt, width=max_length, placeholder="...", break_long_words=True
401 )
404def dummy_fail():
405 """
406 Fake function for a Retry object if you don't need to handle each failure.
407 """
408 pass
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
418def experimental(cls):
419 """
420 Decorator to mark a class as experimental.
421 """
422 original_init = cls.__init__
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)
433 cls.__init__ = new_init
434 return cls
437def warn_experimental(name, stacklevel=2):
438 import warnings
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)
448def experimental_method() -> Callable[[C], C]:
449 """
450 Decorator to mark a function as experimental.
451 """
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)
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)
469 return wrapper
471 return decorator
474def warn_experimental_arg_usage(
475 arg_name: Union[list, str],
476 function_name: str,
477 stacklevel: int = 2,
478):
479 import warnings
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)
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 = ["*"]
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)
508 def decorator(func: C) -> C:
509 if inspect.iscoroutinefunction(func):
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)
518 return async_wrapper
519 else:
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)
528 return wrapper
530 return decorator