mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(auto-router compression): tag-scoped markers now take precedence over untagged
An untagged marker (no tags key or empty tags list) was matching every request because requested.issuperset(frozenset()) is always true. When an alias carried multiple markers, the loop tried tag-matched markers first, but an untagged one could still match the tag-match query, and then the first one with a policy would be returned. Now only markers with a non-empty tags list can match via the tag-specific lookup; untagged markers are tried only after all tag-specific ones. Regression test added: test_tag_scoped_marker_takes_precedence_over_untagged fails with the old code. Also removed unused Any import per greptile's typing note.
This commit is contained in:
parent
9e286fe94b
commit
d0a8006737
2 changed files with 22 additions and 6 deletions
|
|
@ -15,7 +15,7 @@ each hop sees.
|
|||
import contextvars
|
||||
from collections.abc import Mapping, MutableMapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import AUTO_ROUTER_SUPPRESSED_COMPRESSION_GUARDRAILS_KEY
|
||||
|
|
@ -93,7 +93,9 @@ def policy_for_model(
|
|||
)
|
||||
requested: Final = frozenset(request_tags)
|
||||
tag_matched: Final = tuple(
|
||||
params for params in markers if requested.issuperset(frozenset(params.get("tags") or ()))
|
||||
params
|
||||
for params in markers
|
||||
if (tags := params.get("tags")) and requested.issuperset(frozenset(tags))
|
||||
)
|
||||
for params in (*tag_matched, *markers):
|
||||
policy = policy_from_litellm_params(params)
|
||||
|
|
@ -186,16 +188,16 @@ async def arm_pre_call(data: MutableMapping[str, object], llm_router: "Router |
|
|||
return data
|
||||
|
||||
|
||||
def _snapshot_messages() -> list[dict[str, Any]] | None:
|
||||
def _snapshot_messages() -> list[dict[str, object]] | None:
|
||||
snapshot: Final = _routing_messages_snapshot.get()
|
||||
return None if snapshot is None else [dict(message) for message in snapshot]
|
||||
|
||||
|
||||
async def messages_for_routing(
|
||||
policy: AutoRouterCompressionPolicy | None,
|
||||
messages: list[dict[str, Any]] | None,
|
||||
messages: list[dict[str, object]] | None,
|
||||
request_kwargs: Mapping[str, object],
|
||||
) -> list[dict[str, Any]] | None:
|
||||
) -> list[dict[str, object]] | None:
|
||||
"""Messages to use for a routing decision, per `policy.routing`.
|
||||
|
||||
Returns None when the caller should route on whatever messages it already has.
|
||||
|
|
@ -228,7 +230,9 @@ async def messages_for_routing(
|
|||
)
|
||||
return original
|
||||
|
||||
inputs: GenericGuardrailAPIInputs = {"structured_messages": [dict(m) for m in original]}
|
||||
inputs: GenericGuardrailAPIInputs = {
|
||||
"structured_messages": [dict(m) for m in original] # pyright: ignore[reportAssignmentType] # plain dicts, not AllMessageValues; see headroom.py's own use of this shape
|
||||
}
|
||||
# A throwaway request_data: apply_guardrail writes its stats onto this dict, not
|
||||
# the real request's metadata, so routing-side compression never double-counts
|
||||
# against extract_compression_saved_tokens's model-savings accounting.
|
||||
|
|
|
|||
|
|
@ -138,6 +138,18 @@ class TestPolicyForModel:
|
|||
)
|
||||
assert policy == AutoRouterCompressionPolicy(routing="headroom-a", model=None)
|
||||
|
||||
def test_tag_scoped_marker_takes_precedence_over_untagged(self):
|
||||
"""Regression: when multiple markers exist, the tag-scoped one the request
|
||||
actually matches should be used, not the first untagged one."""
|
||||
router = _FakeRouter(
|
||||
[
|
||||
_marker({"auto_router_routing_compression": "headroom-untagged"}),
|
||||
_marker({"auto_router_routing_compression": "headroom-eu"}, tags=["eu"]),
|
||||
]
|
||||
)
|
||||
policy = policy_for_model(llm_router=router, model_alias="smart-router", team_id=None, request_tags=("eu",))
|
||||
assert policy == AutoRouterCompressionPolicy(routing="headroom-eu", model=None)
|
||||
|
||||
|
||||
class _RecordingCompressionGuardrail(CustomGuardrail):
|
||||
"""A guardrail whose apply_guardrail marks every text message as compressed."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue