mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge 443468d530 into e768ad55ce
This commit is contained in:
commit
a887ffe299
16 changed files with 560 additions and 1 deletions
|
|
@ -519,6 +519,7 @@ model_list:
|
|||
complexity_router_config:
|
||||
classifier_type: heuristic_first
|
||||
heuristic_first_max_tier: SIMPLE
|
||||
heuristic_first_max_context_tokens: 8000
|
||||
classifier_llm_config:
|
||||
model: gpt-5-mini
|
||||
reasoning_effort: low
|
||||
|
|
@ -559,6 +560,18 @@ A request short-circuits, meaning it routes on the scorer's own tier with no cla
|
|||
two things hold: the scorer landed at or below `heuristic_first_max_tier`, and it produced at least
|
||||
one signal. Everything else goes to the classifier, which then decides as it normally would.
|
||||
|
||||
Set `heuristic_first_max_context_tokens` to veto that shortcut when the estimated whole conversation
|
||||
exceeds the limit. The estimate counts all message text at approximately four characters per token,
|
||||
so a short newest nudge in a long agentic session still reaches the classifier. Leave it unset to
|
||||
keep the scorer's tier in control at any conversation size
|
||||
|
||||
The context veto only applies on turns that are classified, so it never overrides a held pin. With
|
||||
`classification_mode: user_turn`, a continuation turn inside a session that already holds a pin
|
||||
replays that pin (`x-litellm-complexity-router-cause: user_turn_continuation`) with no classifier
|
||||
call, however large the conversation has grown. The threshold applies again on the next human ask,
|
||||
which falls through to classification and, when the conversation exceeds the limit, goes to the LLM
|
||||
classifier. With `session_affinity` on, the pin wins for new asks as well
|
||||
|
||||
The signal requirement is what keeps this from quietly routing everything to your cheapest model.
|
||||
A prompt where no dimension fires scores exactly 0.0, which is below `simple_medium`, so the score
|
||||
to tier mapping calls it SIMPLE by default rather than by evidence. Around half of general traffic
|
||||
|
|
|
|||
|
|
@ -497,6 +497,43 @@ def _message_text(content: object) -> str:
|
|||
return content if isinstance(content, str) else ""
|
||||
|
||||
|
||||
def _estimated_tool_result_characters(content: object) -> int:
|
||||
if isinstance(content, str):
|
||||
return len(content)
|
||||
if not isinstance(content, list):
|
||||
return 0
|
||||
return sum(
|
||||
len(part)
|
||||
if isinstance(part, str)
|
||||
else len(text)
|
||||
if isinstance(part, Mapping) and isinstance(text := part.get("text"), str)
|
||||
else 0
|
||||
for part in content
|
||||
)
|
||||
|
||||
|
||||
def _estimated_content_characters(content: object) -> int:
|
||||
if isinstance(content, str):
|
||||
return len(content)
|
||||
if not isinstance(content, list):
|
||||
return 0
|
||||
text_characters: Final = sum(
|
||||
len(text)
|
||||
for part in content
|
||||
if isinstance(part, Mapping) and part.get("type") == "text" and isinstance(text := part.get("text"), str)
|
||||
)
|
||||
tool_result_characters: Final = sum(
|
||||
_estimated_tool_result_characters(part.get("content"))
|
||||
for part in content
|
||||
if isinstance(part, Mapping) and part.get("type") == "tool_result"
|
||||
)
|
||||
return text_characters + tool_result_characters
|
||||
|
||||
|
||||
def _estimated_conversation_tokens(messages: Sequence[Mapping[str, object]] | None) -> int:
|
||||
return sum(_estimated_content_characters(message.get("content")) // 4 for message in messages or ())
|
||||
|
||||
|
||||
def _reminder_block_spans(lowered: str, open_marker: str, close_marker: str) -> Iterator[tuple[int, int]]:
|
||||
"""Span of each complete reminder block for one marker pair, left to right.
|
||||
|
||||
|
|
@ -1988,6 +2025,9 @@ class ComplexityRouter(CustomLogger):
|
|||
A turn carrying images the classifier would see is never decided cheaply: the scorer reads
|
||||
text alone, so its confidence describes a request it has only partly seen, and a trivial
|
||||
caption beside a screenshot is exactly the misrouting vision classification exists to stop.
|
||||
|
||||
A configured conversation-size limit also vetoes the cheap decision because a short newest
|
||||
turn can conceal a complex task in the preceding agentic context.
|
||||
"""
|
||||
tier, score, signals, cause = self._score_and_classify(prompt, system_prompt)
|
||||
scored: Final = ClassificationOutcome(tier=tier, score=score, signals=signals, cause=cause)
|
||||
|
|
@ -1996,12 +2036,17 @@ class ComplexityRouter(CustomLogger):
|
|||
threshold is not None
|
||||
and bool(signals)
|
||||
and not self._classifier_image_parts(messages)
|
||||
and not self._exceeds_heuristic_first_context(messages)
|
||||
and self._active_tier_severity(tier) <= self._active_tier_severity(threshold)
|
||||
)
|
||||
if decided_cheaply:
|
||||
return ClassificationOutcome(tier=tier, score=score, signals=signals, cause="heuristic_first_short_circuit")
|
||||
return await self._llm_classifier_outcome(prompt, system_prompt, request_kwargs, messages, scored=scored)
|
||||
|
||||
def _exceeds_heuristic_first_context(self, messages: Sequence[Mapping[str, object]] | None) -> bool:
|
||||
limit: Final = self.config.heuristic_first_max_context_tokens
|
||||
return limit is not None and _estimated_conversation_tokens(messages) > limit
|
||||
|
||||
async def _classify_hybrid(
|
||||
self,
|
||||
prompt: str,
|
||||
|
|
@ -2786,7 +2831,7 @@ class ComplexityRouter(CustomLogger):
|
|||
else ()
|
||||
)
|
||||
|
||||
cumulative_tokens: Final = sum(len(_message_text(msg.get("content"))) // 4 for msg in messages or ())
|
||||
cumulative_tokens: Final = _estimated_conversation_tokens(messages)
|
||||
trajectory_block: Final = (
|
||||
(f"\nConversation so far: ~{cumulative_tokens} tokens across the request",)
|
||||
if has_prior_conversation
|
||||
|
|
|
|||
|
|
@ -1207,6 +1207,18 @@ class ComplexityRouterConfig(BaseModel):
|
|||
"may not name the highest one, since that would make the LLM classifier unreachable."
|
||||
),
|
||||
)
|
||||
heuristic_first_max_context_tokens: int | None = Field(
|
||||
default=None,
|
||||
gt=0,
|
||||
description=(
|
||||
"The estimated size of the whole conversation, counting all message text at approximately four "
|
||||
"characters per token, above which the local scorer may not decide cheaply and the request goes to "
|
||||
"the LLM classifier even when the newest turn scores at or below heuristic_first_max_tier. The "
|
||||
"newest turn in a long agentic session is usually a short nudge such as 'run the tests' or 'why did "
|
||||
"that fail?' whose token-count signal says nothing about the task living in the conversation. None "
|
||||
"keeps the scorer's tier at any conversation size."
|
||||
),
|
||||
)
|
||||
hybrid_boundary_margin: float | None = Field(
|
||||
default=None,
|
||||
ge=0,
|
||||
|
|
@ -1918,6 +1930,15 @@ class ComplexityRouterConfig(BaseModel):
|
|||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_heuristic_first_max_context_tokens(self) -> "ComplexityRouterConfig":
|
||||
if self.classifier_type != "heuristic_first" and self.heuristic_first_max_context_tokens is not None:
|
||||
raise ValueError(
|
||||
f"heuristic_first_max_context_tokens is set but classifier_type is {self.classifier_type!r}; "
|
||||
"set classifier_type 'heuristic_first' or remove heuristic_first_max_context_tokens"
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_hybrid_boundary_margin(self) -> "ComplexityRouterConfig":
|
||||
if self.classifier_type != "hybrid":
|
||||
|
|
|
|||
204
tests/e2e/router/test_heuristic_first_long_context_e2e.py
Normal file
204
tests/e2e/router/test_heuristic_first_long_context_e2e.py
Normal file
|
|
@ -0,0 +1,204 @@
|
|||
"""Live e2e repros for heuristic-first classification of short turns with context."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from complexity_router_client import ComplexityRouterClient
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import StreamingResponse, require_successful_call
|
||||
from lifecycle import ResourceManager
|
||||
from models import (
|
||||
ChatAssistantTurn,
|
||||
ChatBody,
|
||||
ChatMessage,
|
||||
ChatToolResultTurn,
|
||||
KeyGenerateBody,
|
||||
LiteLLMParamsBody,
|
||||
ToolCall,
|
||||
ToolCallFunction,
|
||||
)
|
||||
from proxy_client import ProxyClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
ROUTER_BACKENDS: Final = ("gpt-5.5", "claude-haiku-4-5")
|
||||
SIMPLE_MODELS: Final = frozenset(("openai/gpt-5.5", "gpt-5.5"))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def heuristic_first_router(proxy: ProxyClient, request: pytest.FixtureRequest) -> str:
|
||||
router_name: Final = f"e2e-heuristic-first-router-{unique_marker()}"
|
||||
model_id: Final = proxy.create_model(
|
||||
router_name,
|
||||
LiteLLMParamsBody(
|
||||
model="auto_router/complexity_router",
|
||||
complexity_router_config={
|
||||
"classifier_type": "heuristic_first",
|
||||
"heuristic_first_max_tier": "MEDIUM",
|
||||
"heuristic_first_max_context_tokens": 8000,
|
||||
"classifier_fallback": "heuristic",
|
||||
"classifier_llm_config": {"model": "gpt-5.5", "timeout_ms": 30000},
|
||||
"tiers": {
|
||||
"SIMPLE": "gpt-5.5",
|
||||
"MEDIUM": "claude-haiku-4-5",
|
||||
"COMPLEX": "claude-haiku-4-5",
|
||||
"REASONING": "claude-haiku-4-5",
|
||||
},
|
||||
},
|
||||
),
|
||||
)
|
||||
request.addfinalizer(lambda: proxy.delete_model(model_id))
|
||||
return router_name
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def heuristic_first_key(
|
||||
resources: ResourceManager,
|
||||
client: ComplexityRouterClient,
|
||||
heuristic_first_router: str,
|
||||
) -> str:
|
||||
key: Final = client.proxy.generate_key(
|
||||
KeyGenerateBody(
|
||||
models=[heuristic_first_router, *ROUTER_BACKENDS],
|
||||
user_id=f"e2e-heuristic-first-{unique_marker()}",
|
||||
)
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_key(key))
|
||||
return key
|
||||
|
||||
|
||||
def _agentic_messages(marker: str) -> tuple[ChatMessage | ChatAssistantTurn | ChatToolResultTurn, ...]:
|
||||
system: Final = ChatMessage(
|
||||
role="system",
|
||||
content=(
|
||||
f"You are a coding agent operating on a repository. Preserve the marker {marker}. "
|
||||
"Use tools to inspect files, run tests, and diagnose failures. Keep track of prior "
|
||||
"commands and their outputs before proposing a fix. Never discard relevant logs or "
|
||||
"assume that a failed command was unrelated to the current change."
|
||||
),
|
||||
)
|
||||
rounds: Final = tuple(
|
||||
turn
|
||||
for round_index in range(5)
|
||||
for turn in (
|
||||
ChatMessage(
|
||||
role="user",
|
||||
content=(
|
||||
f"Inspect the repository state for debugging round {round_index} using marker {marker}. "
|
||||
"Run the relevant checks and report every warning, traceback, and changed file."
|
||||
),
|
||||
),
|
||||
ChatAssistantTurn(
|
||||
content=None,
|
||||
tool_calls=[
|
||||
ToolCall(
|
||||
id=f"{marker}-call-{round_index}",
|
||||
type="function",
|
||||
function=ToolCallFunction(
|
||||
name="run_tests",
|
||||
arguments=f'{{"round": {round_index}, "marker": "{marker}"}}',
|
||||
),
|
||||
)
|
||||
],
|
||||
),
|
||||
ChatToolResultTurn(
|
||||
tool_call_id=f"{marker}-call-{round_index}",
|
||||
content="\n".join(
|
||||
f"{marker} round={round_index} line={line_index} "
|
||||
"synthetic test output records a failing assertion, a retry, a provider "
|
||||
"response, a stack frame, and the captured repository state for diagnosis"
|
||||
for line_index in range(120)
|
||||
),
|
||||
),
|
||||
)
|
||||
)
|
||||
return (system, *rounds, ChatMessage(role="user", content=f"why did that fail? {marker}"))
|
||||
|
||||
|
||||
def _send(
|
||||
client: ComplexityRouterClient,
|
||||
key: str,
|
||||
body: ChatBody,
|
||||
) -> StreamingResponse:
|
||||
response: Final = client.proxy.transport.send(
|
||||
"/chat/completions",
|
||||
headers=client.proxy.transport.bearer(key),
|
||||
json=body,
|
||||
stream=False,
|
||||
)
|
||||
require_successful_call(response)
|
||||
return response
|
||||
|
||||
|
||||
def _assert_classifier_consulted(response: StreamingResponse, context: str) -> None:
|
||||
assert response.headers.get("x-litellm-complexity-router-cause") == "llm_classifier", (
|
||||
f"{context}: expected a successful classifier decision; observed headers={response.headers!r}"
|
||||
)
|
||||
classifier_cost: Final = response.headers.get("x-litellm-classifier-cost")
|
||||
assert classifier_cost is not None, (
|
||||
f"{context}: classifier header missing; observed headers={response.headers!r}; body={response.body[:300]!r}"
|
||||
)
|
||||
try:
|
||||
parsed_cost: Final = float(classifier_cost)
|
||||
except ValueError as exc:
|
||||
raise AssertionError(
|
||||
f"{context}: classifier header was not parseable as float: {classifier_cost!r}; "
|
||||
f"observed headers={response.headers!r}"
|
||||
) from exc
|
||||
assert parsed_cost >= 0, f"{context}: classifier cost was negative: {parsed_cost}; headers={response.headers!r}"
|
||||
|
||||
|
||||
@pytest.mark.covers("reliability.routing.complexity_heuristic.scores_current_ask_only")
|
||||
class TestHeuristicFirstLongContext:
|
||||
def test_short_turn_in_long_agentic_conversation_consults_classifier(
|
||||
self,
|
||||
client: ComplexityRouterClient,
|
||||
heuristic_first_key: str,
|
||||
heuristic_first_router: str,
|
||||
) -> None:
|
||||
marker: Final = unique_marker()
|
||||
response: Final = _send(
|
||||
client,
|
||||
heuristic_first_key,
|
||||
ChatBody(
|
||||
model=heuristic_first_router,
|
||||
messages=_agentic_messages(marker),
|
||||
max_tokens=16,
|
||||
),
|
||||
)
|
||||
_assert_classifier_consulted(
|
||||
response,
|
||||
f"long agentic context marker={marker}",
|
||||
)
|
||||
|
||||
def test_short_single_turn_stays_on_heuristic_path(
|
||||
self,
|
||||
client: ComplexityRouterClient,
|
||||
heuristic_first_key: str,
|
||||
heuristic_first_router: str,
|
||||
) -> None:
|
||||
marker: Final = unique_marker()
|
||||
response: Final = _send(
|
||||
client,
|
||||
heuristic_first_key,
|
||||
ChatBody(
|
||||
model=heuristic_first_router,
|
||||
messages=[ChatMessage(role="user", content=f"why did that fail? {marker}")],
|
||||
max_tokens=16,
|
||||
),
|
||||
)
|
||||
assert "x-litellm-classifier-cost" not in response.headers, (
|
||||
f"single-turn heuristic path unexpectedly consulted classifier; "
|
||||
f"observed headers={response.headers!r}; body={response.body[:300]!r}"
|
||||
)
|
||||
assert response.headers.get("x-litellm-complexity-router-cause") == "heuristic_first_short_circuit", (
|
||||
f"single-turn request should bypass the classifier; observed headers={response.headers!r}"
|
||||
)
|
||||
rows: Final = client.proxy.poll_logs_for_key(heuristic_first_key, min_rows=1)
|
||||
served: Final = tuple(row.model for row in rows if row.model is not None)
|
||||
assert len(served) == 1 and served[0] in SIMPLE_MODELS, (
|
||||
f"single-turn heuristic path should serve SIMPLE backend {sorted(SIMPLE_MODELS)!r}; "
|
||||
f"observed spend-log models={served!r}; headers={response.headers!r}"
|
||||
)
|
||||
|
|
@ -50,6 +50,7 @@ from litellm.router_strategy.complexity_router.complexity_router import (
|
|||
KeywordOverride,
|
||||
_built_in_prompt,
|
||||
_ClassifierCircuitBreaker,
|
||||
_estimated_conversation_tokens,
|
||||
_is_classifier_timeout,
|
||||
_matched_plan_mode_sentinel,
|
||||
classification_system_prompt,
|
||||
|
|
@ -13668,6 +13669,35 @@ class TestHeuristicFirstConfig:
|
|||
assert config.uses_llm_classifier is True
|
||||
assert ComplexityRouterConfig(tiers=dict(HEURISTIC_FIRST_TIERS)).uses_llm_classifier is False
|
||||
|
||||
def test_context_limit_is_accepted_on_heuristic_first(self):
|
||||
config = ComplexityRouterConfig(
|
||||
tiers=dict(HEURISTIC_FIRST_TIERS),
|
||||
classifier_type="heuristic_first",
|
||||
heuristic_first_max_tier="SIMPLE",
|
||||
heuristic_first_max_context_tokens=8000,
|
||||
classifier_llm_config={"model": "haiku-classifier"},
|
||||
)
|
||||
assert config.heuristic_first_max_context_tokens == 8000
|
||||
|
||||
def test_context_limit_is_rejected_on_llm(self):
|
||||
with pytest.raises(ValidationError, match="heuristic_first_max_context_tokens is set but classifier_type"):
|
||||
ComplexityRouterConfig(
|
||||
tiers=dict(HEURISTIC_FIRST_TIERS),
|
||||
classifier_type="llm",
|
||||
heuristic_first_max_context_tokens=8000,
|
||||
classifier_llm_config={"model": "haiku-classifier"},
|
||||
)
|
||||
|
||||
def test_context_limit_rejects_zero(self):
|
||||
with pytest.raises(ValidationError, match="greater than 0"):
|
||||
ComplexityRouterConfig(
|
||||
tiers=dict(HEURISTIC_FIRST_TIERS),
|
||||
classifier_type="heuristic_first",
|
||||
heuristic_first_max_tier="SIMPLE",
|
||||
heuristic_first_max_context_tokens=0,
|
||||
classifier_llm_config={"model": "haiku-classifier"},
|
||||
)
|
||||
|
||||
|
||||
class TestHeuristicFirst:
|
||||
"""Behavior of the heuristic-first chain: when the classifier call is skipped, and when it is not."""
|
||||
|
|
@ -13685,6 +13715,139 @@ class TestHeuristicFirst:
|
|||
assert outcome.signals
|
||||
assert outcome.classifier_cost is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"history_content",
|
||||
[
|
||||
"x" * 80,
|
||||
[{"type": "text", "text": "x" * 80}],
|
||||
[{"type": "tool_result", "content": "x" * 80}],
|
||||
[{"type": "tool_result", "content": [{"type": "text", "text": "x" * 80}]}],
|
||||
],
|
||||
)
|
||||
async def test_long_context_vetoes_cheap_short_turn(self, mock_router_instance, history_content: object):
|
||||
mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "MEDIUM"}'))
|
||||
router = _heuristic_first_router(
|
||||
mock_router_instance,
|
||||
heuristic_first_max_tier="MEDIUM",
|
||||
heuristic_first_max_context_tokens=10,
|
||||
)
|
||||
messages = [
|
||||
{"role": "user", "content": history_content},
|
||||
{"role": "user", "content": "why did that fail?"},
|
||||
]
|
||||
|
||||
outcome = await router.aclassify("why did that fail?", messages=messages)
|
||||
|
||||
mock_router_instance.acompletion.assert_awaited_once()
|
||||
assert outcome.cause != "heuristic_first_short_circuit"
|
||||
assert outcome.cause == "llm_classifier"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("context_limit,history_characters", [(None, 40000), (100, 0), (10, 24)])
|
||||
async def test_context_within_limit_or_unset_keeps_cheap_short_turn(
|
||||
self, mock_router_instance, context_limit: int | None, history_characters: int
|
||||
):
|
||||
mock_router_instance.acompletion = AsyncMock()
|
||||
router = _heuristic_first_router(
|
||||
mock_router_instance,
|
||||
heuristic_first_max_tier="MEDIUM",
|
||||
heuristic_first_max_context_tokens=context_limit,
|
||||
)
|
||||
messages = [
|
||||
{"role": "assistant", "content": "x" * history_characters},
|
||||
{"role": "user", "content": "why did that fail?"},
|
||||
]
|
||||
|
||||
outcome = await router.aclassify("why did that fail?", messages=messages)
|
||||
|
||||
mock_router_instance.acompletion.assert_not_called()
|
||||
assert outcome.cause == "heuristic_first_short_circuit"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_context_limit_preserves_user_turn_pin_until_next_human_ask(self, mock_router_instance):
|
||||
mock_router_instance.cache = DualCache()
|
||||
mock_router_instance.acompletion = AsyncMock(
|
||||
side_effect=(_llm_response('{"tier": "MEDIUM"}'), _llm_response('{"tier": "COMPLEX"}'))
|
||||
)
|
||||
router: Final = _heuristic_first_router(
|
||||
mock_router_instance,
|
||||
heuristic_first_max_tier="MEDIUM",
|
||||
heuristic_first_max_context_tokens=8000,
|
||||
classification_mode="user_turn",
|
||||
)
|
||||
ask: Final = [
|
||||
{"role": "user", "content": "Check whether this rollback is safe."},
|
||||
TestClassificationMode.TOOL_CALL_1,
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": "x" * 40000},
|
||||
{"role": "user", "content": "why did that fail?"},
|
||||
]
|
||||
first: Final = await router.async_pre_routing_hook(
|
||||
model="test-complexity-router",
|
||||
request_kwargs={"metadata": {"session_id": "context-threshold-pin"}},
|
||||
messages=ask,
|
||||
)
|
||||
assert first is not None and first.routing_decision is not None
|
||||
assert (first.routing_decision["cause"], first.routing_decision["tier"]) == ("llm_classifier", "MEDIUM")
|
||||
mock_router_instance.acompletion.assert_awaited_once()
|
||||
|
||||
continuation: Final = [*ask, TestClassificationMode.TOOL_CALL_2, TestClassificationMode.TOOL_RESULT_2]
|
||||
second: Final = await router.async_pre_routing_hook(
|
||||
model="test-complexity-router",
|
||||
request_kwargs={"metadata": {"session_id": "context-threshold-pin"}},
|
||||
messages=continuation,
|
||||
)
|
||||
assert second is not None and second.routing_decision is not None
|
||||
assert (second.model, second.routing_decision["cause"]) == (first.model, "user_turn_continuation")
|
||||
mock_router_instance.acompletion.assert_awaited_once()
|
||||
|
||||
third: Final = await router.async_pre_routing_hook(
|
||||
model="test-complexity-router",
|
||||
request_kwargs={"metadata": {"session_id": "context-threshold-pin"}},
|
||||
messages=[*continuation, {"role": "user", "content": "is that safe?"}],
|
||||
)
|
||||
assert third is not None and third.routing_decision is not None
|
||||
assert (third.routing_decision["cause"], third.routing_decision["tier"]) == ("llm_classifier", "COMPLEX")
|
||||
assert third.model == HEURISTIC_FIRST_TIERS["COMPLEX"]
|
||||
assert mock_router_instance.acompletion.await_count == 2
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"messages, expected",
|
||||
[
|
||||
(None, 0),
|
||||
([{"role": "assistant", "content": None}], 0),
|
||||
([{"role": "user", "content": [{"type": "tool_result", "content": None}]}], 0),
|
||||
(
|
||||
[
|
||||
{"role": "system", "content": "abcd"},
|
||||
{"role": "user", "content": [{"type": "text", "text": "efghij"}, {"type": "image_url"}]},
|
||||
{"role": "assistant", "content": "klmnopqr"},
|
||||
],
|
||||
4,
|
||||
),
|
||||
(
|
||||
[{"role": "user", "content": [{"type": "tool_result", "content": "abcdefghijklmnop"}]}],
|
||||
4,
|
||||
),
|
||||
(
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_result",
|
||||
"content": ["abcd", {"type": "text", "text": "efgh"}, {"type": "image"}],
|
||||
}
|
||||
],
|
||||
}
|
||||
],
|
||||
2,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_estimated_conversation_tokens_counts_text_and_tool_result_parts(self, messages, expected):
|
||||
assert _estimated_conversation_tokens(messages) == expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_signal_prompt_escalates_even_though_it_scores_simple(self, mock_router_instance):
|
||||
"""The core guard. This prompt scores 0.0 and the mapping calls it SIMPLE, which is at the
|
||||
|
|
|
|||
|
|
@ -60,6 +60,7 @@ const CLASSIFIER_TIMEOUT_ID = "classifier-timeout-ms";
|
|||
const CLASSIFIER_CONTEXT_WINDOW_SIZE_ID = "classifier-context-window-size";
|
||||
const CLASSIFIER_CONTEXT_BUDGET_CHARS_ID = "classifier-context-budget-chars";
|
||||
const HYBRID_BOUNDARY_MARGIN_ID = "hybrid-boundary-margin";
|
||||
const HEURISTIC_FIRST_MAX_CONTEXT_TOKENS_ID = "heuristic-first-max-context-tokens";
|
||||
const HEURISTIC_V2_SUCCESS_THRESHOLD_ID = "heuristic-v2-success-threshold";
|
||||
|
||||
const CUSTOM_PROMPT_WITH_HEURISTIC_FALLBACK =
|
||||
|
|
@ -241,6 +242,20 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
|
|||
onChange({ ...value, heuristic_first_max_tier: tier });
|
||||
};
|
||||
|
||||
const handleHeuristicFirstMaxContextTokensChange = (raw: string) => {
|
||||
setDraft({ id: HEURISTIC_FIRST_MAX_CONTEXT_TOKENS_ID, raw });
|
||||
if (raw.trim() === "") {
|
||||
onChange({ ...value, heuristic_first_max_context_tokens: undefined });
|
||||
return;
|
||||
}
|
||||
const parsed: number = Number(raw);
|
||||
if (Number.isFinite(parsed)) {
|
||||
onChange({ ...value, heuristic_first_max_context_tokens: Math.max(1, Math.round(parsed)) });
|
||||
return;
|
||||
}
|
||||
onChange({ ...value, heuristic_first_max_context_tokens: undefined });
|
||||
};
|
||||
|
||||
const handleHybridBoundaryMarginChange = (raw: string) => {
|
||||
setDraft({ id: HYBRID_BOUNDARY_MARGIN_ID, raw });
|
||||
const parsed = Number(raw);
|
||||
|
|
@ -466,6 +481,26 @@ const ClassificationMethodConfig: React.FC<ClassificationMethodConfigProps> = ({
|
|||
A request the scorer places at or below this tier routes there without a classifier call. Anything the
|
||||
scorer places higher, and anything it found no signal for at all, goes to the classifier instead
|
||||
</p>
|
||||
<Label htmlFor={HEURISTIC_FIRST_MAX_CONTEXT_TOKENS_ID}>Max conversation tokens before classifier</Label>
|
||||
<Input
|
||||
id={HEURISTIC_FIRST_MAX_CONTEXT_TOKENS_ID}
|
||||
type="text"
|
||||
inputMode="numeric"
|
||||
min={1}
|
||||
aria-describedby={`${HEURISTIC_FIRST_MAX_CONTEXT_TOKENS_ID}-help`}
|
||||
value={
|
||||
draft?.id === HEURISTIC_FIRST_MAX_CONTEXT_TOKENS_ID
|
||||
? draft.raw
|
||||
: String(value.heuristic_first_max_context_tokens ?? "")
|
||||
}
|
||||
onChange={(event) => handleHeuristicFirstMaxContextTokensChange(event.target.value)}
|
||||
onBlur={() => setDraft(null)}
|
||||
className="w-full"
|
||||
/>
|
||||
<p id={`${HEURISTIC_FIRST_MAX_CONTEXT_TOKENS_ID}-help`} className="text-sm text-muted-foreground">
|
||||
Above this estimated conversation size, consult the classifier even for a short ask. Leave blank to disable
|
||||
this limit. With user-turn classification, tool continuations keep their pinned model
|
||||
</p>
|
||||
</div>
|
||||
)}
|
||||
|
||||
|
|
|
|||
|
|
@ -52,6 +52,49 @@ const baseProps = {
|
|||
};
|
||||
|
||||
describe("ComplexityRouterConfig", () => {
|
||||
it("edits and clears the heuristic-first conversation limit", () => {
|
||||
const initialValue: ComplexityRouterConfigValue = {
|
||||
...defaultValue,
|
||||
classifier_type: "heuristic_first",
|
||||
heuristic_first_max_tier: "SIMPLE",
|
||||
heuristic_first_max_context_tokens: 8000,
|
||||
};
|
||||
const onChange = vi.fn();
|
||||
const StatefulConfig = () => {
|
||||
const [value, setValue] = React.useState(initialValue);
|
||||
return (
|
||||
<ComplexityRouterConfig
|
||||
{...baseProps}
|
||||
value={value}
|
||||
onChange={(nextValue) => {
|
||||
onChange(nextValue);
|
||||
setValue(nextValue);
|
||||
}}
|
||||
/>
|
||||
);
|
||||
};
|
||||
renderWithProviders(<StatefulConfig />);
|
||||
openAutoRouterAdvanced("Classification Method");
|
||||
const limit = screen.getByRole("textbox", { name: "Max conversation tokens before classifier" });
|
||||
expect(limit).toHaveValue("8000");
|
||||
|
||||
fireEvent.change(limit, { target: { value: "12000" } });
|
||||
fireEvent.blur(limit);
|
||||
expect(limit).toHaveValue("12000");
|
||||
expect(onChange).toHaveBeenLastCalledWith({ ...initialValue, heuristic_first_max_context_tokens: 12000 });
|
||||
|
||||
fireEvent.change(limit, { target: { value: "invalid" } });
|
||||
fireEvent.blur(limit);
|
||||
expect(limit).toHaveValue("");
|
||||
expect(onChange).toHaveBeenLastCalledWith({ ...initialValue, heuristic_first_max_context_tokens: undefined });
|
||||
|
||||
fireEvent.change(limit, { target: { value: "8000" } });
|
||||
fireEvent.change(limit, { target: { value: "" } });
|
||||
fireEvent.blur(limit);
|
||||
expect(limit).toHaveValue("");
|
||||
expect(onChange).toHaveBeenLastCalledWith({ ...initialValue, heuristic_first_max_context_tokens: undefined });
|
||||
});
|
||||
|
||||
it("should render", async () => {
|
||||
renderWithProviders(<ComplexityRouterConfig {...baseProps} />);
|
||||
expect(screen.getByText("Models by tier")).toBeInTheDocument();
|
||||
|
|
|
|||
|
|
@ -375,6 +375,8 @@ export interface ComplexityRouterConfigValue {
|
|||
classification_examples?: string;
|
||||
/** Highest tier the scorer may decide alone under heuristic_first. Required by that type, rejected by the others. */
|
||||
heuristic_first_max_tier?: string;
|
||||
/** Conversation token estimate above which heuristic_first defers to the classifier. */
|
||||
heuristic_first_max_context_tokens?: number;
|
||||
/** How near a tier boundary a score may land before hybrid defers to the classifier. Required by that type, rejected by the others. */
|
||||
hybrid_boundary_margin?: number;
|
||||
classification_mode?: ClassificationMode;
|
||||
|
|
|
|||
|
|
@ -1129,6 +1129,7 @@ describe("heuristic_first", () => {
|
|||
...baseParams,
|
||||
classifierType: "heuristic_first",
|
||||
heuristicFirstMaxTier: "SIMPLE",
|
||||
heuristicFirstMaxContextTokens: 8000,
|
||||
classifierLlmConfig: { model: "gpt-4o-mini", timeout_ms: 3000 },
|
||||
classifierContextWindowSize: 5,
|
||||
classifierContextBudgetChars: 4000,
|
||||
|
|
@ -1139,6 +1140,15 @@ describe("heuristic_first", () => {
|
|||
const config = buildComplexityRouterConfig(heuristicFirstParams);
|
||||
expect(config.classifier_type).toBe("heuristic_first");
|
||||
expect(config.heuristic_first_max_tier).toBe("SIMPLE");
|
||||
expect(config.heuristic_first_max_context_tokens).toBe(8000);
|
||||
});
|
||||
|
||||
it("omits heuristic_first_max_context_tokens when empty", () => {
|
||||
const config = buildComplexityRouterConfig({
|
||||
...heuristicFirstParams,
|
||||
heuristicFirstMaxContextTokens: undefined,
|
||||
});
|
||||
expect(config.heuristic_first_max_context_tokens).toBeUndefined();
|
||||
});
|
||||
|
||||
it("keeps every classifier key the operator set, since heuristic_first still calls the classifier", () => {
|
||||
|
|
|
|||
|
|
@ -151,6 +151,7 @@ export interface StoredComplexityRouterConfig {
|
|||
classification_prompt?: unknown;
|
||||
classification_examples?: unknown;
|
||||
heuristic_first_max_tier?: unknown;
|
||||
heuristic_first_max_context_tokens?: unknown;
|
||||
hybrid_boundary_margin?: unknown;
|
||||
tier_labels?: unknown;
|
||||
classifier_type?: ClassifierType | "oss_classifier";
|
||||
|
|
@ -219,6 +220,7 @@ export interface BuildComplexityRouterConfigParams {
|
|||
classificationPrompt: string | undefined;
|
||||
classificationExamples: string | undefined;
|
||||
heuristicFirstMaxTier: string | undefined;
|
||||
heuristicFirstMaxContextTokens?: number;
|
||||
hybridBoundaryMargin?: number;
|
||||
classificationMode: ClassificationMode | undefined;
|
||||
sessionAffinity: boolean;
|
||||
|
|
@ -299,6 +301,7 @@ export interface ComplexityRouterConfigPayload {
|
|||
classification_prompt?: string;
|
||||
classification_examples?: string;
|
||||
heuristic_first_max_tier?: string;
|
||||
heuristic_first_max_context_tokens?: number;
|
||||
hybrid_boundary_margin?: number;
|
||||
classification_mode: ClassificationMode;
|
||||
session_affinity: boolean;
|
||||
|
|
@ -599,6 +602,7 @@ const classifierWireFields = (
|
|||
classifierLlmConfig,
|
||||
classifierFallback,
|
||||
heuristicFirstMaxTier,
|
||||
heuristicFirstMaxContextTokens,
|
||||
hybridBoundaryMargin,
|
||||
classifierContextWindowSize,
|
||||
classifierContextBudgetChars,
|
||||
|
|
@ -609,6 +613,7 @@ const classifierWireFields = (
|
|||
| "classifierLlmConfig"
|
||||
| "classifierFallback"
|
||||
| "heuristicFirstMaxTier"
|
||||
| "heuristicFirstMaxContextTokens"
|
||||
| "hybridBoundaryMargin"
|
||||
| "classifierContextWindowSize"
|
||||
| "classifierContextBudgetChars"
|
||||
|
|
@ -627,6 +632,10 @@ const classifierWireFields = (
|
|||
...(supportsFallback && classifierFallback !== undefined && { classifier_fallback: classifierFallback }),
|
||||
...(effectiveType === "heuristic_first" &&
|
||||
heuristicFirstMaxTier?.trim() && { heuristic_first_max_tier: heuristicFirstMaxTier }),
|
||||
...(effectiveType === "heuristic_first" &&
|
||||
heuristicFirstMaxContextTokens !== undefined && {
|
||||
heuristic_first_max_context_tokens: heuristicFirstMaxContextTokens,
|
||||
}),
|
||||
...(effectiveType === "hybrid" &&
|
||||
hybridBoundaryMargin !== undefined && { hybrid_boundary_margin: hybridBoundaryMargin }),
|
||||
...(usesClassifierContext(effectiveType) &&
|
||||
|
|
@ -669,6 +678,7 @@ export const buildComplexityRouterConfig = ({
|
|||
classificationPrompt,
|
||||
classificationExamples,
|
||||
heuristicFirstMaxTier,
|
||||
heuristicFirstMaxContextTokens,
|
||||
hybridBoundaryMargin,
|
||||
classificationMode,
|
||||
sessionAffinity,
|
||||
|
|
@ -732,6 +742,7 @@ export const buildComplexityRouterConfig = ({
|
|||
classifierLlmConfig,
|
||||
classifierFallback,
|
||||
heuristicFirstMaxTier,
|
||||
heuristicFirstMaxContextTokens,
|
||||
hybridBoundaryMargin,
|
||||
classifierContextWindowSize,
|
||||
classifierContextBudgetChars,
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ const standard: ComplexityRouterConfigValue = {
|
|||
classifier_context_budget_chars: 16000,
|
||||
classifier_context_include_assistant_turns: true,
|
||||
classifier_fallback: "default_model",
|
||||
heuristic_first_max_context_tokens: 8000,
|
||||
tiers: { SIMPLE: ["efficient"], MEDIUM: ["middle"], COMPLEX: [], REASONING: ["capable"] },
|
||||
};
|
||||
|
||||
|
|
@ -62,6 +63,7 @@ describe("transitionClassifierType", () => {
|
|||
classifier_context_budget_chars: 16000,
|
||||
classifier_context_include_assistant_turns: true,
|
||||
classifier_fallback: "default_model",
|
||||
...(target === "heuristic_first" && { heuristic_first_max_context_tokens: 8000 }),
|
||||
};
|
||||
expect(result).toMatchObject(expectedSettings);
|
||||
});
|
||||
|
|
@ -70,6 +72,7 @@ describe("transitionClassifierType", () => {
|
|||
const result = transitionClassifierType(standard, target);
|
||||
expect(result.classifier_llm_config).toEqual({ model: "judge", timeout_ms: 20000 });
|
||||
expect(result.classifier_fallback).toBeUndefined();
|
||||
expect(result.heuristic_first_max_context_tokens).toBeUndefined();
|
||||
if (target === "capability") {
|
||||
expect(result.capability_classifier_config?.base_threshold).toBeNaN();
|
||||
} else {
|
||||
|
|
|
|||
|
|
@ -51,6 +51,8 @@ export const transitionClassifierType = (
|
|||
classifierType === "heuristic_first"
|
||||
? value.heuristic_first_max_tier ?? DEFAULT_HEURISTIC_FIRST_MAX_TIER
|
||||
: undefined,
|
||||
heuristic_first_max_context_tokens:
|
||||
classifierType === "heuristic_first" ? value.heuristic_first_max_context_tokens : undefined,
|
||||
hybrid_boundary_margin:
|
||||
classifierType === "hybrid" ? value.hybrid_boundary_margin ?? DEFAULT_HYBRID_BOUNDARY_MARGIN : undefined,
|
||||
...nonReasoningTierFields(classifierType, value),
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ export const builderParamsFromValue = (
|
|||
classificationPrompt: value.classification_prompt,
|
||||
classificationExamples: value.classification_examples,
|
||||
heuristicFirstMaxTier: value.heuristic_first_max_tier,
|
||||
heuristicFirstMaxContextTokens: value.heuristic_first_max_context_tokens,
|
||||
hybridBoundaryMargin: value.hybrid_boundary_margin,
|
||||
classificationMode: value.classification_mode,
|
||||
tierLabels: value.tier_labels,
|
||||
|
|
|
|||
|
|
@ -850,6 +850,7 @@ describe("managed keys survive an untouched open-and-save", () => {
|
|||
classifier_type: "heuristic_first",
|
||||
heuristic_v2_success_threshold: 0.89,
|
||||
heuristic_first_max_tier: "SIMPLE",
|
||||
heuristic_first_max_context_tokens: 8000,
|
||||
classifier_llm_config: { model: "gpt-4o-mini", timeout_ms: 3000, reasoning_effort: "low" },
|
||||
classifier_context_window_size: 5,
|
||||
classifier_context_budget_chars: 4000,
|
||||
|
|
|
|||
|
|
@ -109,6 +109,7 @@ export const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([
|
|||
"classification_prompt",
|
||||
"classification_examples",
|
||||
"heuristic_first_max_tier",
|
||||
"heuristic_first_max_context_tokens",
|
||||
"hybrid_boundary_margin",
|
||||
"heuristic_v2_success_threshold",
|
||||
"classification_mode",
|
||||
|
|
|
|||
|
|
@ -111,6 +111,10 @@ export const hydrateComplexityRouterConfig = (
|
|||
typeof parsedConfig.heuristic_first_max_tier === "string" && parsedConfig.heuristic_first_max_tier.trim() !== ""
|
||||
? parsedConfig.heuristic_first_max_tier
|
||||
: undefined,
|
||||
heuristic_first_max_context_tokens:
|
||||
typeof parsedConfig.heuristic_first_max_context_tokens === "number"
|
||||
? parsedConfig.heuristic_first_max_context_tokens
|
||||
: undefined,
|
||||
hybrid_boundary_margin:
|
||||
typeof parsedConfig.hybrid_boundary_margin === "number" ? parsedConfig.hybrid_boundary_margin : undefined,
|
||||
classification_mode:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue