From a4be6a9a6fdbd85b5dbc369abb0f2247be89cfc2 Mon Sep 17 00:00:00 2001 From: eugene-yao-zocdoc Date: Mon, 24 Aug 2026 16:36:23 -0400 Subject: [PATCH] fix(redis): address credential provider review findings Generated with AI Co-Authored-By: Claude Code --- litellm/_redis.py | 35 ++++++----- litellm/caching/redis_cache.py | 26 +++----- pyproject.toml | 4 +- tests/test_litellm/test_redis.py | 103 ++++++++++++++++++++++++++----- uv.lock | 4 +- 5 files changed, 121 insertions(+), 51 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index 0f9716a5396..6ff4c292c47 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -14,6 +14,7 @@ import json import os from collections.abc import Callable from typing import Final +from urllib.parse import urlsplit, urlunsplit import redis import redis.asyncio as async_redis @@ -156,7 +157,7 @@ def _get_redis_cluster_kwargs(client=None): def _get_redis_env_kwarg_mapping(): PREFIX: Final = "REDIS_" - exclude_from_environment: Final = {"credential_provider"} + exclude_from_environment: Final = frozenset({"credential_provider"}) return {f"{PREFIX}{x.upper()}": x for x in _get_redis_kwargs() if x not in exclude_from_environment} @@ -355,6 +356,14 @@ def get_redis_url_from_environment(): return f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" +def _url_without_userinfo(url: str) -> str: + """redis-py rejects a url that carries its own username or password next to a credential + provider, so the provider's credentials replace whatever userinfo the url was configured with.""" + parts: Final = urlsplit(url) + netloc: Final = parts.netloc.rsplit("@", 1)[-1] + return urlunsplit((parts.scheme, netloc, parts.path, parts.query, parts.fragment)) + + def _get_redis_client_logic(**env_overrides): """ Common functionality across sync + async redis client implementations @@ -476,10 +485,7 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("username", None) redis_kwargs.pop("password", None) if redis_kwargs.get("url") is not None: - from urllib.parse import urlsplit, urlunsplit - - parsed_url = urlsplit(redis_kwargs["url"]) - redis_kwargs["url"] = urlunsplit(parsed_url._replace(netloc=parsed_url.netloc.rsplit("@", 1)[-1])) + redis_kwargs["url"] = _url_without_userinfo(redis_kwargs["url"]) if "url" in redis_kwargs and redis_kwargs["url"] is not None: # Only strip host/port/db/password when not routing to a cluster. @@ -490,8 +496,6 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs.pop("port", None) redis_kwargs.pop("db", None) redis_kwargs.pop("password", None) - if redis_kwargs.get("credential_provider") is not None: - redis_kwargs.pop("username", None) elif ( "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None @@ -623,18 +627,17 @@ def _async_auth_kwargs(redis_kwargs: dict) -> dict: """Swaps a connect func an async path cannot run for the equivalent credential provider, which supersedes any static username or password redis-py would otherwise reject it with.""" explicit_provider: Final = redis_kwargs.get("credential_provider") - if explicit_provider is not None: - superseded: Final = frozenset({"redis_connect_func", "username", "password"}) - kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in superseded) - return dict(kept) - - credential_provider: Final = _async_credential_provider(redis_kwargs.get("redis_connect_func")) + credential_provider: Final = ( + explicit_provider + if explicit_provider is not None + else _async_credential_provider(redis_kwargs.get("redis_connect_func")) + ) if credential_provider is None: return redis_kwargs - automatic_superseded: Final = frozenset({"redis_connect_func", "username", "password"}) - automatic_kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in automatic_superseded) - return dict(automatic_kept, credential_provider=credential_provider) + superseded: Final = frozenset({"redis_connect_func", "username", "password"}) + kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in superseded) + return dict(kept, credential_provider=credential_provider) # mutable-ok: the branches below mutate these kwargs def get_redis_client(**env_overrides): diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 991b6c8c6c5..0207a571dd6 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -175,6 +175,10 @@ _RedisCallResult = TypeVar("_RedisCallResult") _swallowed_redis_failures: Final[ContextVar[int]] = ContextVar("litellm_swallowed_redis_failures", default=0) +def _opaque_kwarg_key(value: object) -> str: + return f"{type(value).__name__}-{id(value)}" + + @functools.lru_cache(maxsize=1) def _redis_health_error_types() -> tuple[type, ...]: """Exception types that mean the Redis backend itself is unhealthy. @@ -398,24 +402,14 @@ class RedisCache(BaseCache): """ Generate a cache key for the async Redis client based on connection parameters. This ensures different Redis configurations use different cached clients. + + Kwargs the caller hands over as live objects (a credential provider, a connect func) are not + JSON-serializable and carry no stable value identity, so they key on instance identity. """ - # Create a stable representation of redis_kwargs for hashing # Sort keys to ensure consistent hash regardless of parameter order - redis_kwargs: Final[dict[str, object]] = self.redis_kwargs - provider: Final = redis_kwargs.get("credential_provider") - redis_connect_func: Final = redis_kwargs.get("redis_connect_func") - sorted_kwargs: Final = sorted( - item for item in redis_kwargs.items() if item[0] not in {"credential_provider", "redis_connect_func"} - ) - kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True) - identity_suffix: Final = ( - "" - if provider is None and redis_connect_func is None - else f":provider-{id(provider)}" - if provider is not None - else f":connect-func-{id(redis_connect_func)}" - ) - kwargs_hash: Final = hashlib.sha256(f"{kwargs_str}{identity_suffix}".encode()).hexdigest()[:16] + sorted_kwargs: Final = sorted(self.redis_kwargs.items()) + kwargs_str: Final = json.dumps(sorted_kwargs, sort_keys=True, default=_opaque_kwarg_key) + kwargs_hash: Final = hashlib.sha256(kwargs_str.encode()).hexdigest()[:16] return f"async-redis-client-{kwargs_hash}" def init_async_client( diff --git a/pyproject.toml b/pyproject.toml index c80ba143512..b9c514dd88b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,8 +54,8 @@ proxy = [ "rq>=2.7.0,<3.0", "redis>=5.3.1,<6.0", "orjson>=3.11.6,<4.0", - # redis-py's C response parser. It arrives with redis (via rq) either way; naming - # it here is what makes redis-py select _HiredisParser instead of the Python one. + # redis-py's C response parser. Nothing imports it; naming it here is what makes + # redis-py select _HiredisParser instead of the pure-Python one. "hiredis>=3.0.0,<4.0", "apscheduler>=3.11.2,<4.0", "fastapi-sso>=0.19.0,<1.0", diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 38b2bd5296a..1476650ac3f 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -95,21 +95,6 @@ def clear_gcp_iam_token_cache(): _token_cache.clear() -def test_redis_uses_the_hiredis_response_parser(): - """The proxy extra must keep redis-py's C response parser available.""" - from redis._parsers import _HiredisParser - from redis.connection import HIREDIS_AVAILABLE, DefaultParser - - if not HIREDIS_AVAILABLE: - pytest.skip("hiredis is not installed in this test environment") - - assert DefaultParser is _HiredisParser - - client = get_redis_client(host="redis-host", port=6379) - connection = client.connection_pool.make_connection() - assert isinstance(connection._parser, _HiredisParser) - - def test_redis_allowlists_include_credential_provider(): assert "credential_provider" in _get_redis_kwargs() assert "credential_provider" in _get_redis_url_kwargs() @@ -230,6 +215,19 @@ def test_async_url_pool_preserves_credential_provider_identity(clean_redis_envir assert pool.connection_kwargs["credential_provider"] is provider +def test_async_url_pool_strips_userinfo_for_the_provider(clean_redis_environment): + """The url allowlist has to carry the provider through, and redis-py rejects it next to userinfo.""" + provider = _StubCredentialProvider() + + pool = get_redis_connection_pool(url="rediss://url-user:url-pass@redis-host:6379/3", credential_provider=provider) + + connection = pool.make_connection() + assert connection.credential_provider is provider + assert connection.username is None + assert connection.password is None + assert connection.db == 3 + + def test_sync_cluster_preserves_credential_provider_identity(clean_redis_environment): provider = _StubCredentialProvider() startup_nodes = [{"host": "cluster-node", "port": 6379}] @@ -274,6 +272,47 @@ def test_explicit_provider_skips_automatic_auth_and_callback(clean_redis_environ assert "redis_connect_func" not in redis_kwargs +@pytest.mark.parametrize( + "overrides", + [ + {"gcp_ssl_ca_certs": "/tmp/ca.pem"}, + {"gcp_service_account": "sa@example.com", "gcp_ssl_ca_certs": "/tmp/ca.pem"}, + ], + ids=["certs-without-service-account", "both-alongside-a-provider"], +) +def test_gcp_kwargs_never_survive_client_logic(clean_redis_environment, overrides): + """redis.Redis has no gcp_* parameters, so anything left behind raises TypeError on connect.""" + redis_kwargs = _get_redis_client_logic( + host="redis-host", + port=6379, + credential_provider=_StubCredentialProvider() if "gcp_service_account" in overrides else None, + **overrides, + ) + + assert "gcp_service_account" not in redis_kwargs + assert "gcp_ssl_ca_certs" not in redis_kwargs + + +def test_provider_keeps_the_rest_of_the_url_intact(clean_redis_environment): + """Stripping the userinfo must not take the database path, query, or scheme with it.""" + provider = _StubCredentialProvider() + + redis_kwargs = _get_redis_client_logic( + url="rediss://url-user:url-pass@redis-host:6379/3?protocol=3", + credential_provider=provider, + ) + + assert redis_kwargs["url"] == "rediss://redis-host:6379/3?protocol=3" + + +def test_provider_free_url_is_left_untouched(clean_redis_environment): + url = "redis://url-user:url-pass@redis-host:6379/3" + + redis_kwargs = _get_redis_client_logic(url=url) + + assert redis_kwargs["url"] == url + + def test_async_direct_explicit_provider_is_preserved_when_normalization_is_bypassed(): provider = _StubCredentialProvider() redis_kwargs = { @@ -378,6 +417,21 @@ def test_redis_cache_key_does_not_serialize_connect_func(): assert first_key == cache._get_async_client_cache_key() +def test_redis_cache_key_keys_opaque_kwargs_by_identity(): + """Any object a caller passes through must key by identity rather than crash the JSON dump.""" + + class _Opaque: + pass + + first = RedisCache.__new__(RedisCache) + first.redis_kwargs = {"host": "redis-host", "retry": _Opaque()} + second = RedisCache.__new__(RedisCache) + second.redis_kwargs = {"host": "redis-host", "retry": _Opaque()} + + assert first._get_async_client_cache_key() == first._get_async_client_cache_key() + assert first._get_async_client_cache_key() != second._get_async_client_cache_key() + + def test_get_redis_url_from_environment_single_url(monkeypatch): """Test when REDIS_URL is directly provided""" # Set the environment variable @@ -1184,6 +1238,25 @@ def test_url_config_drops_kwargs_the_connection_cannot_accept(client_only_kwarg, assert pool.connection_kwargs.get("socket_timeout") == 5.0 +def test_redis_uses_the_hiredis_response_parser(): + """The C parser must be the one redis-py actually picks. + + hiredis is declared in the `proxy` extra purely for speed; nothing imports it, so + dropping it from pyproject.toml would silently fall back to the pure-Python parser + with no other symptom. redis-py selects it at import time, so asserting on the + selection is what catches that. + """ + from redis._parsers import _HiredisParser + from redis.connection import HIREDIS_AVAILABLE, DefaultParser + + assert HIREDIS_AVAILABLE, "hiredis is not installed; redis-py fell back to the pure-Python parser" + assert DefaultParser is _HiredisParser, f"redis-py selected {DefaultParser.__name__}, expected _HiredisParser" + + client = get_redis_client(host="redis-host", port=6379) + connection = client.connection_pool.make_connection() + assert isinstance(connection._parser, _HiredisParser) + + def test_init_arg_names_sees_through_decorated_inits(): """redis-py >= 7.4 wraps AbstractConnection.__init__ with @deprecated_args, whose wrapper is declared (self, *args, **kwargs). Introspecting the wrapper directly diff --git a/uv.lock b/uv.lock index f8af2f60b11..fae22e3759e 100644 --- a/uv.lock +++ b/uv.lock @@ -4342,10 +4342,10 @@ proxy = [ { name = "pyroscope-io", marker = "sys_platform != 'win32'" }, { name = "python-multipart" }, { name = "pyyaml" }, + { name = "redis" }, { name = "restrictedpython" }, { name = "rich" }, { name = "rq" }, - { name = "redis" }, { name = "soundfile" }, { name = "starlette" }, { name = "uvicorn" }, @@ -4552,13 +4552,13 @@ requires-dist = [ { name = "python3-saml", marker = "extra == 'saml'", specifier = ">=1.16.0,<2.0" }, { name = "pyyaml", marker = "extra == 'cli'", specifier = ">=6.0.3,<7.0" }, { name = "pyyaml", marker = "extra == 'proxy'", specifier = ">=6.0.3,<7.0" }, + { name = "redis", marker = "extra == 'proxy'", specifier = ">=5.3.1,<6.0" }, { name = "redisvl", marker = "extra == 'extra-proxy'", specifier = ">=0.4.1,<1.0" }, { name = "requests", marker = "extra == 'cli'", specifier = ">=2.32.0,<3.0" }, { name = "resend", marker = "extra == 'extra-proxy'", specifier = ">=2.23.0,<3.0" }, { name = "restrictedpython", marker = "extra == 'proxy'", specifier = ">=8.1,<9.0" }, { name = "rich", marker = "extra == 'cli'", specifier = ">=13.9.4,<14.0" }, { name = "rich", marker = "extra == 'proxy'", specifier = ">=13.9.4,<14.0" }, - { name = "redis", marker = "extra == 'proxy'", specifier = ">=5.3.1,<6.0" }, { name = "rq", marker = "extra == 'proxy'", specifier = ">=2.7.0,<3.0" }, { name = "semantic-router", marker = "python_full_version < '3.14' and extra == 'semantic-router'", specifier = ">=0.1.15,<1.0" }, { name = "sentry-sdk", marker = "extra == 'proxy-runtime'", specifier = ">=2.21.0,<3.0" },