fix(redis): address credential provider review findings

Generated with AI

Co-Authored-By: Claude Code
This commit is contained in:
eugene-yao-zocdoc 2026-08-24 16:36:23 -04:00
parent 11061d13c9
commit a4be6a9a6f
5 changed files with 121 additions and 51 deletions

View file

@ -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):

View file

@ -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(

View file

@ -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",

View file

@ -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

4
uv.lock generated
View file

@ -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" },