From 7424673d00361f063021e204015ba48b8aefdf01 Mon Sep 17 00:00:00 2001 From: apex-mochen <2756823972@qq.com> Date: Sat, 3 Oct 2026 13:32:13 +0800 Subject: [PATCH] fix(complexity_router): invalidate stale pins on new user asks Signed-off-by: apex-mochen <2756823972@qq.com> --- .../complexity_router/complexity_router.py | 7 ++ .../router_strategy/test_complexity_router.py | 86 +++++++++++++++++++ 2 files changed, 93 insertions(+) diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 3610991a20d..3a8edf0c144 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -4390,6 +4390,13 @@ 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 + ): + await self.litellm_router_instance.cache.async_delete_cache(key=cache_key) + 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..f81e7cd609d 100644 --- a/tests/unit/router_strategy/test_complexity_router.py +++ b/tests/unit/router_strategy/test_complexity_router.py @@ -16961,3 +16961,89 @@ class TestNonReasoningTier: "complex", "reasoning", ) + + +@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: + 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"}}, + ], + ) + 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"}, + }, + ) + 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", + ]