fix(complexity_router): sanitize budget reservation inside forwarded user_api_key_auth

This commit is contained in:
mateo-berri 2026-07-11 19:10:10 +00:00
parent 4a4262aa0d
commit 9388d03089
No known key found for this signature in database
2 changed files with 87 additions and 9 deletions

View file

@ -78,16 +78,30 @@ def _append_custom_keywords(base_keywords: list[str], custom_keywords: Optional[
#
# Note: user_api_key_auth itself is intentionally kept; it is required by
# _filter_deployments_by_model_access_groups to scope embedding/classifier model
# selection to the caller's authorized access groups. Only the budget reservation
# sub-field inside it is the problem -- stripping the whole object would allow
# an access-group-scoped caller to reach embedding deployments outside their group.
# selection to the caller's authorized access groups. It is forwarded as a sanitized
# copy with its budget_reservation sub-field removed, because the proxy cost callback
# (_get_budget_reservation_from_metadata) falls back to reading the reservation from
# inside the auth object when the top-level key is absent; forwarding it unsanitized
# would re-create the exact double-finalization this stripping exists to prevent.
_BUDGET_RESERVATION_METADATA_KEYS = frozenset({"user_api_key_budget_reservation"})
def _sanitize_user_api_key_auth(auth: Any) -> Any:
if isinstance(auth, dict):
return {k: v for k, v in auth.items() if k != "budget_reservation"}
if getattr(auth, "budget_reservation", None) is not None and hasattr(auth, "model_copy"):
return auth.model_copy(update={"budget_reservation": None})
return auth
def _classifier_call_metadata(metadata: Optional[dict[str, Any]]) -> Optional[dict[str, Any]]:
if not metadata:
return metadata
return {k: v for k, v in metadata.items() if k not in _BUDGET_RESERVATION_METADATA_KEYS}
return {
k: _sanitize_user_api_key_auth(v) if k == "user_api_key_auth" else v
for k, v in metadata.items()
if k not in _BUDGET_RESERVATION_METADATA_KEYS
}
class DimensionScore:

View file

@ -1260,13 +1260,20 @@ class TestLLMClassifier:
"hi", request_kwargs={"litellm_metadata": request_metadata}
)
call_kwargs = mock_router_instance.acompletion.call_args.kwargs
# user_api_key_budget_reservation is stripped (budget enforcement) but
# user_api_key_budget_reservation is stripped (budget enforcement) while
# user_api_key_auth is kept so _filter_deployments_by_model_access_groups
# can scope the classifier's model selection to the caller's access groups.
# can scope the classifier's model selection to the caller's access groups,
# but only as a sanitized copy without its budget_reservation sub-field:
# the cost callback falls back to reading the reservation from inside the
# auth object when the top-level key is absent.
assert call_kwargs["metadata"] == {
"user_api_key": "sk-abc",
"user_api_key_team_id": "team-1",
"user_api_key_auth": {"models": ["gpt-4o"], "budget_reservation": {"reserved_cost": 1.0}},
"user_api_key_auth": {"models": ["gpt-4o"]},
}
assert request_metadata["user_api_key_auth"] == {
"models": ["gpt-4o"],
"budget_reservation": {"reserved_cost": 1.0},
}
@pytest.mark.asyncio
@ -1652,14 +1659,20 @@ class TestSemanticKeywordTierRules:
assert fake_router.async_embedding_kwargs, "expected an embedding call for the prompt"
# user_api_key_budget_reservation is stripped to prevent budget-bypass.
# user_api_key_auth is kept so _filter_deployments_by_model_access_groups
# scopes the embedding model selection to the caller's authorized groups.
# scopes the embedding model selection to the caller's authorized groups,
# but its budget_reservation sub-field is removed because the cost callback
# falls back to reading the reservation from inside the auth object.
expected = {
"user_api_key_hash": "hash-abc",
"user_api_key_team_id": "team-1",
"user_api_key_auth": {"models": ["voyage-3-5"], "budget_reservation": {"reserved_cost": 1.0}},
"user_api_key_auth": {"models": ["voyage-3-5"]},
}
assert fake_router.async_embedding_kwargs[0]["metadata"] == expected
assert fake_router.async_embedding_kwargs[0]["litellm_metadata"] == expected
assert caller_metadata["user_api_key_auth"] == {
"models": ["voyage-3-5"],
"budget_reservation": {"reserved_cost": 1.0},
}
@pytest.mark.asyncio
async def test_semantic_routelayer_build_runs_off_event_loop(self, basic_config):
@ -1943,3 +1956,54 @@ class TestKeywordOverrideEdgeCases:
)
assert result is not None
assert result.model in {"gpt-4o-mini", "gpt-4o", "claude-sonnet-4-20250514", "o1-preview"}
class TestSubCallMetadataSanitization:
"""The proxy cost callback must not be able to recover the parent budget reservation
from sub-call metadata, in either of the shapes it knows how to read."""
def test_cost_callback_cannot_recover_reservation_from_sanitized_metadata(self):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.proxy_track_cost_callback import (
_get_budget_reservation_from_metadata,
)
from litellm.router_strategy.complexity_router.complexity_router import (
_classifier_call_metadata,
)
reservation = {"reserved_cost": 1.0}
auth_shapes = (
{"models": ["gpt-4o"], "budget_reservation": dict(reservation)},
UserAPIKeyAuth(api_key="sk-abc", budget_reservation=dict(reservation)),
)
for auth in auth_shapes:
metadata = {
"user_api_key_hash": "hash-abc",
"user_api_key_budget_reservation": dict(reservation),
"user_api_key_auth": auth,
}
assert _get_budget_reservation_from_metadata(metadata) == reservation
sanitized = _classifier_call_metadata(metadata)
assert sanitized is not None
assert sanitized["user_api_key_auth"] is not None
assert _get_budget_reservation_from_metadata(sanitized) is None
def test_sanitized_auth_keeps_access_group_fields_and_leaves_original_untouched(self):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.router_strategy.complexity_router.complexity_router import (
_classifier_call_metadata,
)
auth = UserAPIKeyAuth(
api_key="sk-abc",
team_id="team-1",
budget_reservation={"reserved_cost": 1.0},
)
sanitized = _classifier_call_metadata({"user_api_key_auth": auth})
assert sanitized is not None
sanitized_auth = sanitized["user_api_key_auth"]
assert sanitized_auth.budget_reservation is None
assert sanitized_auth.team_id == "team-1"
assert sanitized_auth.api_key == auth.api_key
assert auth.budget_reservation == {"reserved_cost": 1.0}