From 25ecff36d5b8432dd9ffc2cf5a0e4872f8a59666 Mon Sep 17 00:00:00 2001 From: apex-mochen <2756823972@qq.com> Date: Sat, 3 Oct 2026 15:44:51 +0800 Subject: [PATCH] fix(complexity_router): tolerate task pin deletion failures --- .../complexity_router/complexity_router.py | 7 +++- .../router_strategy/test_complexity_router.py | 40 ++++++++++++++++++- 2 files changed, 43 insertions(+), 4 deletions(-) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 3a8edf0c144..6e5ed29b277 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -4183,7 +4183,7 @@ class ComplexityRouter(CustomLogger): async def async_pre_routing_hook( self, model: str, - request_kwargs: dict, + request_kwargs: dict[str, object], messages: list[dict[str, Any]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, @@ -4395,7 +4395,10 @@ class ComplexityRouter(CustomLogger): and not pin_replay_allowed and self._matched_plan_mode_signal(request_kwargs, resolved_messages) is None ): - await self.litellm_router_instance.cache.async_delete_cache(key=cache_key) + try: + await self.litellm_router_instance.cache.async_delete_cache(key=cache_key) + except Exception: # noqa: BLE001 # optional cache failures must not prevent routing + verbose_router_logger.debug("ComplexityRouter: task pin invalidation failed; continuing classification") routed_response: Final = await self._classify_and_route( model=model, diff --git a/tests/unit/router_strategy/test_complexity_router.py b/tests/unit/router_strategy/test_complexity_router.py index f81e7cd609d..2375a53aed0 100644 --- a/tests/unit/router_strategy/test_complexity_router.py +++ b/tests/unit/router_strategy/test_complexity_router.py @@ -10,7 +10,7 @@ import logging import math import sys import time -from collections.abc import AsyncIterator, Mapping, Sequence +from collections.abc import AsyncIterator, Iterator, Mapping, Sequence from copy import deepcopy from functools import partial from types import MappingProxyType @@ -34,6 +34,7 @@ from litellm.router_utils.auto_router_model_naming import ( from litellm._logging import verbose_router_logger from litellm.caching.dual_cache import DualCache from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_cache import RedisCache from litellm.constants import ( OUTPUT_TOKEN_CEILING_PARAMS, RETURN_RAW_MODEL_NAME_METADATA_KEY, @@ -16963,9 +16964,37 @@ class TestNonReasoningTier: ) +@pytest.fixture +def unavailable_task_pin_redis() -> Iterator[RedisCache]: + import redis + import redis.asyncio + from litellm import in_memory_llm_clients_cache + + class UnavailableSyncConnection(redis.Connection): + def connect(self) -> None: + raise redis.ConnectionError("scripted Redis outage") + + class UnavailableAsyncConnection(redis.asyncio.Connection): + async def connect(self) -> None: + raise redis.ConnectionError("scripted Redis outage") + + sync_pool: Final = redis.ConnectionPool(connection_class=UnavailableSyncConnection) + async_pool: Final = redis.asyncio.BlockingConnectionPool(connection_class=UnavailableAsyncConnection) + client: Final = redis.asyncio.Redis(connection_pool=async_pool) + cache: Final = RedisCache(host="synthetic.invalid", connection_pool=sync_pool) + client_key: Final = cache._get_async_client_cache_key() + cache.redis_async_client = client + yield cache + in_memory_llm_clients_cache.delete_cache(client_key) + sync_pool.disconnect() + + @pytest.mark.asyncio @pytest.mark.parametrize("circuit_breaker_enabled", [True, False]) -async def test_new_user_ask_classifier_failure_clears_previous_turn_pin(circuit_breaker_enabled: bool) -> None: +@pytest.mark.parametrize("redis_unavailable", [False, True]) +async def test_new_user_ask_classifier_failure_clears_previous_turn_pin( + circuit_breaker_enabled: bool, redis_unavailable: bool, unavailable_task_pin_redis: RedisCache +) -> None: from openai import AsyncOpenAI outcomes: Final = iter(("SIMPLE", None, "COMPLEX")) @@ -17002,6 +17031,8 @@ async def test_new_user_ask_classifier_failure_clears_previous_turn_pin(circuit_ {"model_name": "default", "litellm_params": {"model": "openai/default"}}, ], ) + if redis_unavailable: + underlying.cache.redis_cache = unavailable_task_pin_redis router: Final = ComplexityRouter( "auto", underlying, @@ -17019,6 +17050,11 @@ async def test_new_user_ask_classifier_failure_clears_previous_turn_pin(circuit_ "tiers": {"SIMPLE": "cheap", "MEDIUM": "default", "COMPLEX": "default", "REASONING": "default"}, }, ) + if redis_unavailable: + litellm.in_memory_llm_clients_cache.set_cache( + unavailable_task_pin_redis._get_async_client_cache_key(), + unavailable_task_pin_redis.redis_async_client, + ) kwargs: Final = {"metadata": {"session_id": "stale-pin-regression"}} first: Final = [{"role": "user", "content": "Hello"}] second: Final = [