mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(redis): address credential provider review findings
Generated with AI Co-Authored-By: Claude Code
This commit is contained in:
parent
11061d13c9
commit
a4be6a9a6f
5 changed files with 121 additions and 51 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
4
uv.lock
generated
|
|
@ -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" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue