Skip to content

Commit 2417613

Browse files
committed
refactor(redis): build the async auth kwargs instead of mutating them twice
Both async entrypoints edited the kwargs dict in place with the same five lines. One shared transform returns the swapped copy instead.
1 parent 14faec9 commit 2417613

1 file changed

Lines changed: 14 additions & 12 deletions

File tree

litellm/_redis.py

Lines changed: 14 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -594,6 +594,18 @@ def _async_credential_provider(redis_connect_func: object | None) -> CredentialP
594594
return None
595595

596596

597+
def _async_auth_kwargs(redis_kwargs: dict) -> dict:
598+
"""Swaps a connect func an async path cannot run for the equivalent credential provider,
599+
which supersedes any static username or password redis-py would otherwise reject it with."""
600+
credential_provider: Final = _async_credential_provider(redis_kwargs.get("redis_connect_func"))
601+
if credential_provider is None:
602+
return redis_kwargs
603+
604+
superseded: Final = frozenset({"redis_connect_func", "username", "password"})
605+
kept: Final = ((k, v) for k, v in redis_kwargs.items() if k not in superseded)
606+
return dict(kept, credential_provider=credential_provider) # mutable-ok: the branches below mutate these kwargs
607+
608+
597609
def get_redis_client(**env_overrides):
598610
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
599611

@@ -620,12 +632,7 @@ def get_redis_async_client(
620632
connection_pool: async_redis.BlockingConnectionPool | None = None,
621633
**env_overrides,
622634
) -> async_redis.Redis | async_redis.RedisCluster:
623-
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
624-
credential_provider: Final = _async_credential_provider(redis_kwargs.get("redis_connect_func"))
625-
if credential_provider is not None:
626-
redis_kwargs["credential_provider"] = credential_provider
627-
for superseded in ("redis_connect_func", "username", "password"):
628-
redis_kwargs.pop(superseded, None)
635+
redis_kwargs: Final = _async_auth_kwargs(_get_redis_client_logic(**env_overrides))
629636

630637
if "startup_nodes" in redis_kwargs:
631638
from redis.cluster import ClusterNode
@@ -688,12 +695,7 @@ def get_redis_async_client(
688695
def get_redis_connection_pool(
689696
**env_overrides,
690697
) -> async_redis.BlockingConnectionPool | None:
691-
redis_kwargs: Final = _get_redis_client_logic(**env_overrides)
692-
credential_provider: Final = _async_credential_provider(redis_kwargs.get("redis_connect_func"))
693-
if credential_provider is not None:
694-
redis_kwargs["credential_provider"] = credential_provider
695-
for superseded in ("redis_connect_func", "username", "password"):
696-
redis_kwargs.pop(superseded, None)
698+
redis_kwargs: Final = _async_auth_kwargs(_get_redis_client_logic(**env_overrides))
697699
verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs)
698700

699701
if "startup_nodes" in redis_kwargs:

0 commit comments

Comments
 (0)