feat(complexity-router): add user_turn classification mode that holds one target per tool loop

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
tin 2026-08-29 18:11:50 +00:00
parent 352789257d
commit 4d7da425da
7 changed files with 270 additions and 4 deletions

View file

@ -386,6 +386,29 @@ def _iter_human_asks_newest_first(
)
def _turn_carries_human_ask(
messages: Sequence[Mapping[str, object]] | None,
marker_pairs: tuple[tuple[str, str], ...] = _DEFAULT_REMINDER_MARKERS,
) -> bool:
"""Whether this request's newest turn is a human speaking, rather than a tool loop continuing.
Only the newest turn, because the human ask that started a twenty-turn tool loop is still the
newest *user* turn on every one of those requests: reading anything but the last message reads
every continuation as a fresh ask. Tool output arrives as a `tool` role on chat completions and
as non-text `tool_result` blocks on a user turn on the Messages surface, and the latter flattens
to empty human text, so both read as a continuation here.
An unreadable request is treated as a human turn: that classifies it, which is the mode's own
fallback rather than holding a target on no evidence.
"""
newest: Final = messages[-1] if messages else None
if newest is None:
return True
if newest.get("role") != "user":
return False
return bool(_human_text(newest.get("content"), marker_pairs))
def _conversation_is_continuing(messages: Sequence[Mapping[str, object]] | None) -> bool:
"""Whether this request continues a conversation that was already underway.
@ -2247,7 +2270,10 @@ class ComplexityRouter(CustomLogger):
@property
def _uses_tier_pin(self) -> bool:
return bool(self.config.session_affinity and not self.config.plugins)
return bool(
(self.config.session_affinity or self.config.classification_mode == "user_turn")
and not self.config.plugins
)
@property
def _uses_deployment_pin(self) -> bool:
@ -2305,7 +2331,13 @@ class ComplexityRouter(CustomLogger):
session_id: Final = self._get_session_id_from_request_kwargs(request_kwargs) if use_session_affinity else None
cache_key = self._get_session_affinity_cache_key(session_id, request_kwargs) if session_id is not None else None
if cache_key is not None:
# session_affinity holds its pin on every turn; user_turn holds it only while the loop the
# human set off is still running, and lets the next human turn reclassify over it.
holds_pin_this_turn: Final = self.config.session_affinity or not _turn_carries_human_ask(
resolved_messages, self._reminder_markers
)
if cache_key is not None and holds_pin_this_turn:
pinned_value: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
pinned_pin: Final = _parse_session_affinity_pin(pinned_value)
if pinned_pin is not None:
@ -2354,10 +2386,13 @@ class ComplexityRouter(CustomLogger):
kwargs_metadata: Final = request_kwargs.setdefault("metadata", {})
if isinstance(kwargs_metadata, dict):
kwargs_metadata[ADAPTIVE_ROUTER_CHOSEN_MODEL_KEY] = routed_model
held_pin_cause: Final[RoutingDecisionCause] = (
"session_affinity_pin" if self.config.session_affinity else "user_turn_pin"
)
cause: RoutingDecisionCause = (
"plan_mode"
if plan_floored
else ("session_affinity_escalation" if escalated else "session_affinity_pin")
else ("session_affinity_escalation" if escalated else held_pin_cause)
)
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=%s, routed_model=%s", cause, routed_model

View file

@ -839,6 +839,23 @@ class ComplexityRouterConfig(BaseModel):
description="Minimum cosine similarity for a semantic keyword match",
)
classification_mode: Literal["every_request", "user_turn"] = Field(
default="every_request",
description=(
"Which requests get classified in a multi-turn agent session. 'every_request' (the "
"default) classifies every inference request, so a tool continuation can land on a "
"different tier than the ask that started it. 'user_turn' classifies only the requests "
"that carry a new human ask and holds that target for the tool loop the ask sets off, "
"which keeps one tool loop on one model, preserves its provider prompt cache, and skips "
"the classifier call on continuation turns. A continuation is a request whose newest "
"turn is tool output or an assistant continuation rather than human text. Needs a "
"resolvable session_id to key the held target on, and falls back to classifying every "
"request without one; suppressed when plugins are configured, for the same reason "
"session_affinity is. Inert when session_affinity is on, which holds one target for the "
"whole session and so never reclassifies at all."
),
)
# Session affinity: pin the first turn's routed model for the rest of the session
session_affinity: bool = Field(
default=False,

View file

@ -2842,6 +2842,11 @@ RoutingDecisionCause = Literal[
"housekeeping",
"session_affinity_pin",
"session_affinity_escalation",
# classification_mode 'user_turn': this request continued a tool loop rather than carrying a
# new human ask, so it reused the target chosen on the turn the human last spoke and the
# classifier was never called. Distinct from "session_affinity_pin", which holds one target for
# a whole session: this pin is replaced the next time a human speaks.
"user_turn_pin",
"default_fallback",
"keyword",
"quality_tier",

View file

@ -4420,6 +4420,201 @@ class TestSessionAffinity:
assert request_kwargs_2["metadata"]["adaptive_router_chosen_model"] == "cheap"
class TestUserTurnClassificationMode:
"""classification_mode 'user_turn': classify the requests carrying a human ask and hold that
target across the tool loop the ask sets off, instead of classifying every continuation."""
HUMAN_ASK = {"role": "user", "content": "Let's think step by step and reason through this refactor."}
FOLLOW_UP_ASK = {"role": "user", "content": "Hello!"}
ASSISTANT_TOOL_CALL = {
"role": "assistant",
"content": None,
"tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "bash", "arguments": "{}"}}],
}
TOOL_RESULT = {"role": "tool", "tool_call_id": "call_1", "content": "exit 0"}
MESSAGES_SURFACE_TOOL_RESULT = {
"role": "user",
"content": [{"type": "tool_result", "tool_use_id": "call_1", "content": "exit 0"}],
}
@staticmethod
def _router(mock_router_instance, **overrides) -> ComplexityRouter:
mock_router_instance.cache = DualCache()
return ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": "gpt-4o-mini",
"MEDIUM": "gpt-4o",
"COMPLEX": "claude-sonnet-4-20250514",
"REASONING": "o1-preview",
},
**overrides,
},
)
@staticmethod
def _request_kwargs(session_id: str | None = "loop-session") -> Dict:
return {"metadata": {"session_id": session_id}} if session_id is not None else {}
@pytest.mark.parametrize(
"continuation",
[TOOL_RESULT, MESSAGES_SURFACE_TOOL_RESULT, ASSISTANT_TOOL_CALL],
ids=["chat_completions_tool_role", "messages_surface_tool_result", "assistant_continuation"],
)
@pytest.mark.asyncio
async def test_continuation_holds_the_target_without_classifying(self, mock_router_instance, continuation):
router = self._router(mock_router_instance, classification_mode="user_turn")
request_kwargs = self._request_kwargs()
first = await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[self.HUMAN_ASK]
)
assert first.model == "o1-preview"
with patch.object(router, "aclassify", wraps=router.aclassify) as spy:
second = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, continuation],
)
spy.assert_not_called()
assert second.model == "o1-preview"
assert second.routing_decision["cause"] == "user_turn_pin"
@pytest.mark.asyncio
async def test_next_human_turn_reclassifies_over_the_held_target(self, mock_router_instance):
"""The held target is the tool loop's, not the session's: when the human speaks again the
request is classified on its own merits, which is what separates this from session_affinity."""
router = self._router(mock_router_instance, classification_mode="user_turn")
request_kwargs = self._request_kwargs()
await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[self.HUMAN_ASK]
)
held = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT],
)
reclassified = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT, self.FOLLOW_UP_ASK],
)
assert held.model == "o1-preview"
assert reclassified.model == "gpt-4o-mini"
assert reclassified.routing_decision["cause"] in ("heuristic_scorer", "reasoning_override")
@pytest.mark.asyncio
async def test_next_loop_holds_the_target_the_new_ask_classified_to(self, mock_router_instance):
"""The pin the second loop holds is the second ask's own, not the first loop's leftover."""
router = self._router(mock_router_instance, classification_mode="user_turn")
request_kwargs = self._request_kwargs()
await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[self.HUMAN_ASK]
)
history = [self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT, self.FOLLOW_UP_ASK]
await router.async_pre_routing_hook(model="test-model", request_kwargs=request_kwargs, messages=history)
second_loop = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[*history, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT],
)
assert second_loop.model == "gpt-4o-mini"
assert second_loop.routing_decision["cause"] == "user_turn_pin"
@pytest.mark.asyncio
async def test_every_request_mode_classifies_the_continuation(self, mock_router_instance):
"""The behavior 'user_turn' is measured against: the default classifies continuations too,
so a classifier that answers differently mid-loop moves the model mid-loop."""
from litellm.router_strategy.complexity_router.complexity_router import ClassificationOutcome
router = self._router(mock_router_instance)
outcomes = [
ClassificationOutcome(tier=ComplexityTier.REASONING, score=None, signals=(), cause="llm_classifier"),
ClassificationOutcome(tier=ComplexityTier.SIMPLE, score=None, signals=(), cause="llm_classifier"),
]
request_kwargs = self._request_kwargs()
with patch.object(router, "aclassify", new=AsyncMock(side_effect=outcomes)) as spy:
first = await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[self.HUMAN_ASK]
)
second = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT],
)
assert spy.call_count == 2
assert (first.model, second.model) == ("o1-preview", "gpt-4o-mini")
@pytest.mark.asyncio
async def test_no_session_id_falls_back_to_classifying_every_request(self, mock_router_instance):
"""There is nothing to key the held target on, so the mode is inert rather than guessing."""
router = self._router(mock_router_instance, classification_mode="user_turn")
with patch.object(router, "aclassify", wraps=router.aclassify) as spy:
continuation = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=self._request_kwargs(None),
messages=[self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT],
)
spy.assert_called_once()
assert continuation.model == "o1-preview"
assert continuation.routing_decision["cause"] in ("heuristic_scorer", "reasoning_override")
@pytest.mark.asyncio
async def test_plugins_suppress_the_held_target(self, mock_router_instance):
"""Same reason session_affinity is suppressed: a held target was never re-checked against a
policy plugin whose decision can change between turns."""
router = self._router(mock_router_instance, classification_mode="user_turn", plugins=[_DummyPlugin()])
request_kwargs = self._request_kwargs()
await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[self.HUMAN_ASK]
)
with patch.object(router, "aclassify", wraps=router.aclassify) as spy:
continuation = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT],
)
spy.assert_called_once()
assert continuation.routing_decision["cause"] in ("heuristic_scorer", "reasoning_override")
@pytest.mark.asyncio
async def test_session_affinity_still_holds_across_human_turns(self, mock_router_instance):
"""session_affinity is the stronger promise, so it decides both turns and keeps its own
cause even with the mode set."""
router = self._router(mock_router_instance, classification_mode="user_turn", session_affinity=True)
request_kwargs = self._request_kwargs()
await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[self.HUMAN_ASK]
)
follow_up = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT, self.FOLLOW_UP_ASK],
)
assert follow_up.model == "o1-preview"
assert follow_up.routing_decision["cause"] == "session_affinity_pin"
@pytest.mark.asyncio
async def test_held_target_carries_the_deployment_affinity_marker(self, mock_router_instance):
"""A loop frozen onto one model group but load-balanced across its deployments would still
go cache-cold, which is the whole point of holding the target."""
router = self._router(
mock_router_instance, classification_mode="user_turn", session_affinity_ttl_seconds=321
)
request_kwargs = self._request_kwargs()
first = await router.async_pre_routing_hook(
model="test-model", request_kwargs=request_kwargs, messages=[self.HUMAN_ASK]
)
held = await router.async_pre_routing_hook(
model="test-model",
request_kwargs=request_kwargs,
messages=[self.HUMAN_ASK, self.ASSISTANT_TOOL_CALL, self.TOOL_RESULT],
)
assert first.session_affinity_ttl_seconds == 321
assert held.session_affinity_ttl_seconds == 321
class _DummyPlugin:
async def run(self, context):
return context

View file

@ -186,6 +186,12 @@ describe("RoutingDecisionCard", () => {
expect(screen.queryByText("housekeeping")).not.toBeInTheDocument();
});
it("says a tool-loop turn reused the target held from the last user turn", () => {
render(<RoutingDecisionCard decision={{ ...heuristic, cause: "user_turn_pin", score: undefined }} />);
expect(screen.getByText("Held from the last user turn, classifier skipped")).toBeInTheDocument();
expect(screen.queryByText("user_turn_pin")).not.toBeInTheDocument();
});
it("shows the escalation keyword", () => {
render(
<RoutingDecisionCard decision={{ ...heuristic, escalated: true, escalation_keyword: "LITELLM ESCALATE" }} />,

View file

@ -89,6 +89,7 @@ const CONSTANT_CAUSE_LABELS: Record<string, string> = {
semantic_keyword_match: "Semantic keyword match",
session_affinity_pin: "Pinned to session",
session_affinity_escalation: "Escalated from session pin",
user_turn_pin: "Held from the last user turn, classifier skipped",
quality_tier: "Quality tier mapping",
bandit: "Adaptive bandit",
default_fallback: "Default model, no route matched",

View file

@ -33827,6 +33827,13 @@ export interface components {
adaptive_eligible: "all" | "classified_tier";
/** @description Quality vs cost weights for adaptive selection (used when adaptive=True) */
adaptive_weights?: components["schemas"]["AdaptiveRouterWeights"];
/**
* Classification Mode
* @description Which requests get classified in a multi-turn agent session. 'every_request' (the default) classifies every inference request, so a tool continuation can land on a different tier than the ask that started it. 'user_turn' classifies only the requests that carry a new human ask and holds that target for the tool loop the ask sets off, which keeps one tool loop on one model, preserves its provider prompt cache, and skips the classifier call on continuation turns. A continuation is a request whose newest turn is tool output or an assistant continuation rather than human text. Needs a resolvable session_id to key the held target on, and falls back to classifying every request without one; suppressed when plugins are configured, for the same reason session_affinity is. Inert when session_affinity is on, which holds one target for the whole session and so never reclassifies at all.
* @default every_request
* @enum {string}
*/
classification_mode: "every_request" | "user_turn";
/**
* Classification Prompt
* @description Replaces the opening instructions of the LLM classifier rubric (the judging-criteria prose) for a custom tier set. The per-tier bullets and the trust-boundary paragraph telling the classifier to ignore tier requests embedded in quoted caller text are always appended after it and cannot be overridden. Requires tier_definitions; a built-in-tier router customizes its prompt via classifier_llm_config.system_prompt or classification_rubric instead.
@ -35099,7 +35106,7 @@ export interface components {
* Cause
* @enum {string}
*/
cause?: "heuristic_scorer" | "reasoning_override" | "llm_classifier" | "heuristic_first_short_circuit" | "classifier_plugin" | "classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "session_affinity_pin" | "session_affinity_escalation" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
cause?: "heuristic_scorer" | "reasoning_override" | "llm_classifier" | "heuristic_first_short_circuit" | "classifier_plugin" | "classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_pin" | "default_fallback" | "keyword" | "quality_tier" | "bandit";
/** Classifier Cost */
classifier_cost?: number;
/** Classifier Model */