fix(caching): propagate auth cache invalidation over Redis Cluster via a node-level pub/sub client (#43110)

* fix(caching): give Redis Cluster clients a node-level pub/sub client so auth invalidation propagates

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(caching): shorten pub/sub client docstrings

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* style(caching): noqa BLE001 on best-effort pubsub client close

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(caching): satisfy LIT002/LIT006 in pubsub client derivation

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(auth): rename fake pubsub hook to init_pubsub_client

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): rename fake pubsub hook to init_pubsub_client

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(caching): drop class-level health ping patch from pub/sub client tests

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(caching): close cluster pubsub pools and cover failure paths

* fix(caching): clean up expired cluster pubsub clients safely

* test(mcp): arm cancellation deadline after requests start

---------

Co-authored-by: joshua <joshua@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-25 07:56:46 -07:00 • committed by GitHub
parent 61953318bf
commit d33d36ce86
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 319 additions and 59 deletions

View file

@ -8,9 +8,13 @@ on:
paths:
- "litellm/_redis.py"
- "litellm/_redis_credential_provider.py"
- "litellm/caching/redis_cache.py"
- "litellm/caching/evicted_client_closer.py"
- "tests/test_litellm/test_redis.py"
- "tests/local_testing/test_caching.py"
- "tests/test_litellm/caching/test_redis_connection_pool.py"
- "tests/test_litellm/caching/test_redis_cluster_cache.py"
- "tests/test_litellm/caching/test_evicted_client_closer.py"
- ".github/workflows/test-redis-compat.yml"
- "pyproject.toml"
- "uv.lock"
@ -82,6 +86,8 @@ jobs:
uv run --no-sync pytest \
tests/test_litellm/test_redis.py \
tests/test_litellm/caching/test_redis_connection_pool.py \
tests/test_litellm/caching/test_redis_cluster_cache.py \
tests/test_litellm/caching/test_evicted_client_closer.py \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_azure_credentials \
tests/local_testing/test_caching.py::test_sync_cluster_authenticates_with_gcp_credentials \
--tb=short -vv \

View file

@ -136,6 +136,8 @@ def _has_connection_in_flight(client: object) -> bool:
window as the only guard, exactly as it was before this check existed.
"""
try:
if getattr(getattr(client, "connection_pool", None), "_in_use_connections", None):
return True
transport: Final = _transport_of(client)
pooled_busy: Final = _pool_has_busy_connection(transport)
if pooled_busy is not None:

View file

@ -21,6 +21,7 @@ from collections.abc import Awaitable, Callable, Iterator, Sequence
from contextvars import ContextVar
from dataclasses import dataclass
from datetime import timedelta
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
from pydantic import TypeAdapter
@ -290,6 +291,29 @@ def _opaque_kwarg_key(value: object) -> str:
return f"{type(value).__name__}-{id(value)}"
_CLUSTER_ONLY_CONNECTION_KWARGS: Final[frozenset[str]] = frozenset({"response_callbacks"})
def _cluster_node_pubsub_client( # pyright: ignore[reportUnknownParameterType] # redis generics
cluster: async_redis_cluster_client, # pyright: ignore[reportUnknownParameterType] # redis generics
) -> async_redis_client:
"""Plain async client on one cluster node; classic PUBLISH/SUBSCRIBE is broadcast cluster-wide."""
from redis.asyncio import ConnectionPool, Redis
node: Final = cluster.get_default_node() or next(iter(cluster.nodes_manager.startup_nodes.values()), None)
if node is None: # pyright: ignore[reportUnnecessaryComparison] # get_default_node is None before cluster init
raise ValueError("cannot derive a pub/sub client: redis cluster has no default node and no startup nodes")
node_kwargs: Final = MappingProxyType(
{
key: value # pyright: ignore[reportAny] # connection_kwargs values are Any in redis stubs
for key, value in cluster.connection_kwargs.items() # pyright: ignore[reportAny] # connection_kwargs values are Any in redis stubs
if key not in _CLUSTER_ONLY_CONNECTION_KWARGS
}
)
pool: Final = ConnectionPool(host=node.host, port=node.port, **node_kwargs) # pyright: ignore[reportCallIssue, reportArgumentType] # cluster kwargs validated by redis-py at runtime
return Redis.from_pool(pool) # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # redis generics
@functools.lru_cache(maxsize=1)
def _redis_health_error_types() -> tuple[type, ...]:
"""Exception types that mean the Redis backend itself is unhealthy.
@ -738,6 +762,28 @@ class RedisCache(BaseCache):
self.redis_async_client = redis_async_client
return redis_async_client
def init_pubsub_client(self) -> async_redis_client: # pyright: ignore[reportUnknownParameterType] # redis generics
from redis.asyncio import RedisCluster
from litellm import in_memory_llm_clients_cache
client: Final = self.init_async_client() # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # redis generics
if not isinstance(client, RedisCluster):
return client # pyright: ignore[reportUnknownVariableType] # redis generics
cache_key: Final = f"{self._get_async_client_cache_key()}-pubsub"
cached_client: Final = in_memory_llm_clients_cache.get_cache( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # untyped in-memory client cache
key=cache_key
)
if cached_client is not None:
return cast( # cast-ok: per-loop pub/sub client stored by this method # pyright: ignore[reportUnknownVariableType] # redis generics
async_redis_client, cached_client
)
pubsub_client: Final = _cluster_node_pubsub_client( # pyright: ignore[reportUnknownVariableType] # redis generics
cluster=client
)
in_memory_llm_clients_cache.set_cache(key=cache_key, value=pubsub_client, litellm_owned_client=True) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # untyped in-memory client cache
return pubsub_client # pyright: ignore[reportUnknownVariableType] # redis generics
def _async_commands(self) -> _AsyncRedisCommands:
return self.init_async_client()
@ -1785,7 +1831,21 @@ class RedisCache(BaseCache):
self.redis_client.flushall()
async def disconnect(self):
await self.async_redis_conn_pool.disconnect(inuse_connections=True)
from litellm import in_memory_llm_clients_cache
if self.async_redis_conn_pool is not None:
await self.async_redis_conn_pool.disconnect(inuse_connections=True)
cached_pubsub_client: Final = cast( # cast-ok: only this module stores clients under this key # pyright: ignore[reportUnknownVariableType] # redis generics
async_redis_client | None,
in_memory_llm_clients_cache.get_cache( # pyright: ignore[reportUnknownMemberType] # untyped in-memory client cache
key=f"{self._get_async_client_cache_key()}-pubsub"
),
)
if cached_pubsub_client is not None:
try:
await cached_pubsub_client.aclose() # pyright: ignore[reportUnknownMemberType, reportAttributeAccessIssue] # redis stubs leave aclose unknown
except Exception as e: # noqa: BLE001 # best-effort close of a possibly-broken connection
verbose_logger.debug("Error closing cached pub/sub Redis client: %s", e)
try:
self.redis_client.close()
except Exception as e:

View file

@ -74,12 +74,6 @@ def _message_from_data(data: object) -> _CacheInvalidationMessage | None:
async def _publish_to_redis(redis_cache: "RedisCache", cache_key: str, message: str) -> None:
try:
client: Final = _pubsub_capable_client(redis_cache)
if client is None:
verbose_proxy_logger.debug(
"auth cache invalidation publish for %s skipped: cluster redis client has no pub/sub support",
cache_key,
)
return
async with _in_flight_publishes:
await client.publish(auth_cache_invalidation_channel(redis_cache), message)
except Exception as e: # noqa: BLE001 # best-effort publish; mutations must never fail on redis errors
@ -185,12 +179,6 @@ class AuthCacheInvalidationSubscriber:
while True:
try:
client = _pubsub_capable_client(self._redis_cache)
if client is None:
verbose_proxy_logger.warning(
"auth cache invalidation subscriber disabled: cluster redis client has no pub/sub support; "
"cross-worker eviction falls back to the local cache TTL"
)
return
pubsub = client.pubsub()
try:
await pubsub.subscribe(auth_cache_invalidation_channel(self._redis_cache))

View file

@ -82,22 +82,13 @@ def config_sync_channel(redis_cache: "RedisCache") -> str:
return f"{redis_cache.namespace}:{CONFIG_SYNC_CHANNEL}"
def _raw_async_client(redis_cache: "RedisCache") -> object:
return cast( # cast-ok: redis-py generics leave the client type partially unknown
object,
redis_cache.init_async_client(), # pyright: ignore[reportUnknownMemberType] # redis generics
def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient:
return cast( # cast-ok: protocol view of the pub/sub-capable async redis client
_ConfigSyncPubSubClient,
redis_cache.init_pubsub_client(), # pyright: ignore[reportUnknownMemberType] # redis generics
)
def _pubsub_capable_client(redis_cache: "RedisCache") -> _ConfigSyncPubSubClient | None:
from redis.asyncio import Redis
client: Final = _raw_async_client(redis_cache)
if isinstance(client, Redis):
return cast(_ConfigSyncPubSubClient, client) # cast-ok: protocol view of the standalone redis client
return None
@dataclass(frozen=True, slots=True)
class _ConfigChangeMessage:
object_type: str
@ -112,12 +103,6 @@ async def publish_config_change(redis_cache: "RedisCache | None", object_type: s
return
try:
client: Final = _pubsub_capable_client(redis_cache)
if client is None:
verbose_proxy_logger.debug(
"config sync publish for %s skipped: cluster redis client has no pub/sub support",
object_type,
)
return
await client.publish(config_sync_channel(redis_cache), _config_change_message_json(object_type))
except Exception as e: # noqa: BLE001 # best-effort publish; writes must never fail on redis errors
verbose_proxy_logger.warning("config sync publish for %s failed: %s", object_type, e)
@ -238,12 +223,6 @@ class ConfigSyncSubscriber:
while True:
try:
client = _pubsub_capable_client(self._redis_cache)
if client is None:
verbose_proxy_logger.warning(
"config sync subscriber disabled: cluster redis client has no pub/sub support; "
"interval polling remains the only sync mechanism"
)
return
pubsub = client.pubsub()
try:
await pubsub.subscribe(config_sync_channel(self._redis_cache))

View file

@ -10,9 +10,11 @@ never closed, because litellm does not own its lifecycle.
import asyncio
import gc
import weakref
from unittest.mock import AsyncMock
import httpx
import pytest
from redis.asyncio import ConnectionPool, Redis
from litellm.caching.evicted_client_closer import EvictedClientCloser
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@ -72,6 +74,33 @@ def make_closer(clock: FakeClock, grace_seconds: float = 60.0) -> EvictedClientC
return EvictedClientCloser(grace_seconds=grace_seconds, clock=clock)
@pytest.mark.asyncio
async def test_redis_client_is_closed_only_after_its_subscription_releases_the_connection():
closer = EvictedClientCloser(grace_seconds=0)
pool = ConnectionPool()
client = Redis.from_pool(pool)
closed = asyncio.Event()
connection = AsyncMock()
connection.disconnect.side_effect = closed.set
pool._available_connections.append(connection)
borrowed = pool.get_available_connection()
closer.mark_owned(client)
closer.schedule(client)
closer.reap()
await asyncio.sleep(0.05)
connection.disconnect.assert_not_awaited()
assert closer.pending_count == 1
await pool.release(borrowed)
closer.reap()
await asyncio.wait_for(closed.wait(), timeout=1)
connection.disconnect.assert_awaited_once()
assert closer.pending_count == 0
async def _trickling_upstream(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
"""Serves a chunked body slowly, so a request stays on the wire long enough to observe."""
await reader.read(4096)

View file

@ -1,13 +1,20 @@
import asyncio
from importlib import import_module
import json
from unittest.mock import MagicMock, patch
import ssl
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi.testclient import TestClient
from redis.asyncio import Redis, RedisCluster
from redis.asyncio.cluster import ClusterNode
from redis.asyncio.connection import SSLConnection
from litellm.caching.redis_cache import RedisCache
from litellm.caching.redis_cluster_cache import RedisClusterCache
from litellm.caching.evicted_client_closer import EvictedClientCloser
from litellm.caching.llm_caching_handler import LLMClientCache
@patch("litellm._redis.init_redis_cluster")
@ -175,3 +182,167 @@ def test_router_create_redis_cache_cluster_detection(
with patch.object(RedisCache, "__init__", _mock_redis_cache_init):
redis_cache = Router._create_redis_cache(cache_config)
assert isinstance(redis_cache, expected_cache_type)
def _isolated_redis_cache(host: str) -> RedisCache:
"""RedisCache whose sync client and pool are stubbed out."""
with (
patch("litellm._redis.get_redis_client", return_value=MagicMock()),
patch("litellm._redis.get_redis_connection_pool", return_value=MagicMock()),
):
return RedisCache(host=host, port=6379)
def _cluster_for_pubsub(startup_node_host: str = "10.9.9.9") -> RedisCluster:
"""Uninitialized RedisCluster carrying the connection kwargs a real one would."""
return RedisCluster(
startup_nodes=[ClusterNode(host=startup_node_host, port=7000)],
password="cluster-secret",
socket_timeout=7.0,
)
def test_init_pubsub_client_derives_a_node_client_for_cluster_backend() -> None:
"""LIT-8543: a cluster-backed cache must return a pub/sub-capable client.
The derived client pins a plain Redis connection pool to the cluster's
default node, inheriting the connection kwargs minus cluster-only keys.
"""
cache = _isolated_redis_cache("cluster-pubsub-default-node")
cluster = _cluster_for_pubsub()
node = ClusterNode(host="10.1.2.3", port=7001)
cluster.nodes_manager.default_node = node
cache.init_async_client = MagicMock(return_value=cluster)
client = cache.init_pubsub_client()
assert isinstance(client, Redis) and not isinstance(client, RedisCluster)
kwargs = client.connection_pool.connection_kwargs
assert kwargs["host"] == "10.1.2.3"
assert kwargs["port"] == 7001
assert kwargs["password"] == "cluster-secret"
assert kwargs["socket_timeout"] == 7.0
assert "response_callbacks" not in kwargs
def test_init_pubsub_client_falls_back_to_first_startup_node() -> None:
"""Before cluster initialization there is no default node; the first
startup node is a valid pub/sub target."""
cache = _isolated_redis_cache("cluster-pubsub-startup-fallback")
cluster = _cluster_for_pubsub(startup_node_host="10.8.8.8")
cache.init_async_client = MagicMock(return_value=cluster)
client = cache.init_pubsub_client()
assert isinstance(client, Redis)
assert client.connection_pool.connection_kwargs["host"] == "10.8.8.8"
def test_init_pubsub_client_returns_the_same_cached_client_on_repeat_calls() -> None:
cache = _isolated_redis_cache("cluster-pubsub-caching")
cluster = _cluster_for_pubsub()
cluster.nodes_manager.default_node = ClusterNode(host="10.1.2.3", port=7001)
cache.init_async_client = MagicMock(return_value=cluster)
first = cache.init_pubsub_client()
second = cache.init_pubsub_client()
assert first is second
def test_init_pubsub_client_returns_the_shared_async_client_for_standalone() -> None:
cache = _isolated_redis_cache("standalone-pubsub")
standalone = Redis()
cache.init_async_client = MagicMock(return_value=standalone)
assert cache.init_pubsub_client() is standalone
def test_init_pubsub_client_preserves_tls_and_authentication() -> None:
cache = _isolated_redis_cache("cluster-pubsub-tls")
cluster = RedisCluster(
startup_nodes=[ClusterNode(host="redis.example.test", port=7000)],
ssl=True,
ssl_cert_reqs="required",
ssl_check_hostname=True,
username="pubsub-user",
password="test-password",
socket_connect_timeout=3.0,
socket_keepalive=True,
)
cache.init_async_client = MagicMock(return_value=cluster)
client = cache.init_pubsub_client()
connection = client.connection_pool.make_connection()
assert isinstance(connection, SSLConnection)
assert connection.ssl_context.cert_reqs == ssl.CERT_REQUIRED
assert connection.ssl_context.check_hostname is True
assert connection.username == "pubsub-user"
assert connection.password == "test-password"
assert connection.socket_connect_timeout == 3.0
assert connection.socket_keepalive is True
def test_init_pubsub_client_rejects_missing_nodes_and_can_retry() -> None:
cache = _isolated_redis_cache("cluster-pubsub-no-nodes")
cluster = _cluster_for_pubsub()
cluster.nodes_manager.startup_nodes = {}
cache.init_async_client = MagicMock(return_value=cluster)
with pytest.raises(ValueError, match="no default node and no startup nodes"):
cache.init_pubsub_client()
cluster.nodes_manager.default_node = ClusterNode(host="recovered.example.test", port=7001)
client = cache.init_pubsub_client()
assert client.connection_pool.connection_kwargs["host"] == "recovered.example.test"
@pytest.mark.parametrize("close_fails", [False, True])
@pytest.mark.parametrize("has_shared_pool", [False, True])
def test_disconnect_closes_derived_pubsub_connections_even_when_pool_close_fails(
close_fails: bool, has_shared_pool: bool
) -> None:
cache = _isolated_redis_cache(f"cluster-pubsub-close-{close_fails}-{has_shared_pool}")
cache.async_redis_conn_pool = AsyncMock() if has_shared_pool else None
cache.init_async_client = MagicMock(return_value=_cluster_for_pubsub())
async def exercise() -> None:
client = cache.init_pubsub_client()
connection = AsyncMock()
connection.disconnect.side_effect = ConnectionError("connection close failed") if close_fails else None
client.connection_pool._available_connections.append(connection)
await cache.disconnect()
connection.disconnect.assert_awaited_once()
if has_shared_pool:
cache.async_redis_conn_pool.disconnect.assert_awaited_once_with(inuse_connections=True)
cache.redis_client.close.assert_called_once()
asyncio.run(exercise())
def test_expired_pubsub_client_closes_connections_after_eviction() -> None:
cache = _isolated_redis_cache("cluster-pubsub-expired")
cache.init_async_client = MagicMock(return_value=_cluster_for_pubsub())
clients = LLMClientCache(evicted_client_closer=EvictedClientCloser(grace_seconds=0))
async def exercise() -> None:
client = cache.init_pubsub_client()
closed = asyncio.Event()
connection = AsyncMock()
connection.disconnect.side_effect = closed.set
client.connection_pool._available_connections.append(connection)
cache_key = clients.update_cache_key_with_event_loop(f"{cache._get_async_client_cache_key()}-pubsub")
clients.ttl_dict[cache_key] = 0
replacement = cache.init_pubsub_client()
await asyncio.wait_for(closed.wait(), timeout=1)
assert replacement is not client
connection.disconnect.assert_awaited_once()
with patch("litellm.in_memory_llm_clients_cache", clients):
asyncio.run(exercise())

View file

@ -2762,6 +2762,7 @@ async def test_cancellation_delivers_termination_over_tcp(
cancel_mode: str, concurrency: int, termination: str, raise_on_error: bool
) -> None:
started: Final = asyncio.Event()
scope_ready: Final[asyncio.Future[anyio.CancelScope]] = asyncio.get_running_loop().create_future()
terminations: Final[list[bytes]] = []
starts: Final[list[bytes]] = []
stop: Final = asyncio.Event()
@ -2858,13 +2859,16 @@ async def test_cancellation_delivers_termination_over_tcp(
async def invoke():
if cancel_mode == "scope":
with anyio.fail_after(0.2):
with anyio.fail_after(None) as scope:
scope_ready.set_result(scope)
return await calls()
return await calls()
try:
task: Final = asyncio.create_task(invoke())
await asyncio.wait_for(started.wait(), 3)
if cancel_mode == "scope":
(await scope_ready).deadline = anyio.current_time() + 0.2
if cancel_mode == "task":
task.cancel()
expected_error: Final = (

View file

@ -16,7 +16,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
class _FakeRedisCache:
namespace = None
def init_async_client(self) -> object:
def init_pubsub_client(self) -> object:
return object()

View file

@ -8998,7 +8998,7 @@ async def test_invalidate_team_member_spend_state_broadcasts_the_spend_counter_t
def __init__(self) -> None:
self.namespace = None
def init_async_client(self) -> object:
def init_pubsub_client(self) -> object:
return _RecordingRedisClient()
local_spend_counter_cache = DualCache()
@ -9072,7 +9072,7 @@ async def test_invalidate_team_member_spend_state_self_delivered_broadcast_does_
def __init__(self) -> None:
self.namespace = None
def init_async_client(self) -> object:
def init_pubsub_client(self) -> object:
return _RecordingRedisClient()
local_spend_counter_cache = DualCache()

View file

@ -87,7 +87,7 @@ class _FakeRedisCache:
self._client = client
self.namespace = namespace
def init_async_client(self) -> object:
def init_pubsub_client(self) -> object:
return self._client

View file

@ -81,14 +81,25 @@ class _FailingPublishRedisClient(Redis):
raise ConnectionError("redis down")
class _NotRedisClient:
def __init__(self) -> None:
class _ScriptedPubSubClient:
"""Pub/sub-capable client that is not a redis.asyncio.Redis.
Mirrors what RedisCache.init_pubsub_client returns for a cluster backend:
a node-level client exposing publish/pubsub without being an instance of
the standalone Redis class.
"""
def __init__(self, pubsubs: Iterable["_QueuePubSub"]) -> None:
self._scripted_pubsubs = iter(pubsubs)
self.published: List[Tuple[str, str]] = []
async def publish(self, channel: str, message: str) -> int:
self.published.append((channel, message))
return 1
def pubsub(self) -> "_QueuePubSub":
return next(self._scripted_pubsubs)
class _QueuePubSub:
def __init__(self, initial_messages: Iterable[str] = ()) -> None:
@ -159,14 +170,14 @@ class _FakeRedisCache:
self._client = client
self.namespace = namespace
def init_async_client(self) -> object:
def init_pubsub_client(self) -> object:
return self._client
class _ExplodingRedisCache:
namespace: Optional[str] = None
def init_async_client(self) -> object:
def init_pubsub_client(self) -> object:
raise ConnectionError("cannot connect")
@ -215,13 +226,17 @@ async def test_publish_swallows_client_init_errors() -> None:
await publish_config_change(redis_cache=_ExplodingRedisCache(), object_type="litellm_proxymodeltable")
async def test_publish_skips_clients_without_pubsub_support() -> None:
client = _NotRedisClient()
async def test_publish_reaches_cluster_derived_pubsub_clients() -> None:
"""LIT-8543: a cluster-backed cache returns a node-level client from
init_pubsub_client; publishes must go out on it instead of being skipped."""
client = _ScriptedPubSubClient(pubsubs=[])
cache = _FakeRedisCache(client)
await publish_config_change(redis_cache=cache, object_type="litellm_proxymodeltable")
assert client.published == []
assert client.published == [
(CONFIG_SYNC_CHANNEL, json.dumps({"object_type": "litellm_proxymodeltable"}))
]
async def test_subscriber_runs_injected_callbacks_in_order_on_message() -> None:
@ -558,20 +573,26 @@ async def test_stop_before_start_is_a_noop() -> None:
await subscriber.stop()
async def test_subscriber_exits_without_callbacks_when_client_lacks_pubsub() -> None:
cache = _FakeRedisCache(_NotRedisClient())
async def test_subscriber_subscribes_on_cluster_derived_pubsub_client() -> None:
"""LIT-8543: the subscriber used to disable itself on cluster caches; now it
subscribes on the node-level client init_pubsub_client returns."""
pubsub = _QueuePubSub(initial_messages=[json.dumps({"object_type": "litellm_proxymodeltable"})])
cache = _FakeRedisCache(_ScriptedPubSubClient(pubsubs=[pubsub]))
resyncs: List[str] = []
fired = asyncio.Event()
subscriber = ConfigSyncSubscriber(
redis_cache=cache,
resync_callbacks=(_recording_callback(resyncs, "resync", asyncio.Event()),),
resync_callbacks=(_recording_callback(resyncs, "resync", fired),),
debounce_seconds=0.01,
jitter_max_seconds=0.0,
)
subscriber.start()
task = subscriber._task
assert task is not None
await asyncio.wait_for(task, timeout=5)
await asyncio.wait_for(fired.wait(), timeout=5)
await subscriber.stop()
assert resyncs == []
assert pubsub.subscribed_channels == [CONFIG_SYNC_CHANNEL]
assert resyncs == ["resync"]
class _FakeTableActions:

View file

@ -15127,7 +15127,7 @@ async def test_auth_cache_invalidation_subscriber_evicts_byok_credentials_cached
def __init__(self, client: object) -> None:
self._client = client
def init_async_client(self) -> object:
def init_pubsub_client(self) -> object:
return self._client
byok_credential_cache.flush_cache()