diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index fc0865654c7..5d2889e2c9d 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -4185,7 +4185,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, @@ -4392,6 +4392,16 @@ class ComplexityRouter(CustomLogger): ) ) + if ( + cache_key is not None + and not pin_replay_allowed + and self._matched_plan_mode_signal(request_kwargs, resolved_messages) is None + ): + 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, request_kwargs=request_kwargs, diff --git a/tests/unit/router_strategy/test_complexity_router.py b/tests/unit/router_strategy/test_complexity_router.py index 333524ffffc..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, @@ -16961,3 +16962,124 @@ class TestNonReasoningTier: "complex", "reasoning", ) + + +@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]) +@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")) + + def classifier_response(request: httpx.Request) -> httpx.Response: + tier: Final = next(outcomes) + if tier is None: + raise httpx.ReadTimeout("classifier unavailable", request=request) + return httpx.Response( + 200, + json={ + "id": "classification", + "object": "chat.completion", + "created": 0, + "model": "classifier", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": json.dumps({"tier": tier})}, + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + ) + + async with httpx.AsyncClient(transport=httpx.MockTransport(classifier_response)) as transport: + async with AsyncOpenAI(api_key="synthetic", http_client=transport, max_retries=0) as client: + underlying: Final = Router( + default_litellm_params={"client": client}, + model_list=[ + {"model_name": "judge", "litellm_params": {"model": "openai/classifier", "api_key": "synthetic"}}, + {"model_name": "cheap", "litellm_params": {"model": "openai/cheap"}}, + {"model_name": "default", "litellm_params": {"model": "openai/default"}}, + ], + ) + if redis_unavailable: + underlying.cache.redis_cache = unavailable_task_pin_redis + router: Final = ComplexityRouter( + "auto", + underlying, + { + "classifier_type": "llm", + "classification_mode": "user_turn", + "classifier_context_include_assistant_turns": True, + "classifier_llm_config": { + "model": "judge", + "classification_rubric": "agentic", + "circuit_breaker_enabled": circuit_breaker_enabled, + }, + "classifier_fallback": "default_model", + "default_model": "default", + "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 = [ + *first, + {"role": "assistant", "content": "Hello!"}, + {"role": "user", "content": "Fix the queue worker"}, + ] + continuation: Final = [ + *second, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "t1", "type": "function", "function": {"name": "read_file", "arguments": "{}"}} + ], + }, + {"role": "tool", "tool_call_id": "t1", "content": "File contents"}, + ] + responses: Final = [ + await router.async_pre_routing_hook("auto", kwargs, messages=messages) + for messages in (first, second, continuation) + ] + assert [response.model for response in responses] == ["cheap", "default", "default"] + assert [response.routing_decision["cause"] for response in responses] == [ + "llm_classifier", + "default_model_fallback", + "default_model_fallback" if circuit_breaker_enabled else "llm_classifier", + ]