This commit is contained in:
dingdangmao 2026-10-04 13:42:08 -07:00 • committed by GitHub
commit d8657fd207
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 134 additions and 2 deletions

View file

@ -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,

View file

@ -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",
]