fix(complexity_router): invalidate stale pins on new user asks

Signed-off-by: apex-mochen <2756823972@qq.com>
This commit is contained in:
apex-mochen 2026-10-03 13:32:13 +08:00
parent 8efb4a21f6
commit 7424673d00
2 changed files with 93 additions and 0 deletions

View file

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

View file

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