fix(passthrough): carry the budget reservation into request metadata (#36592)

A successful pass-through request left its pre-call budget reservation in
the shared Redis spend counter. `_init_kwargs_for_pass_through_endpoint`
built the request metadata from the sanitized key fields only, so
`_PROXY_track_cost_callback` resolved `budget_reservation = None` and
`increment_spend_counters` added the actual cost on top of a reservation
nobody released. The counter drifted above real spend on every request
until the key falsely tripped BudgetExceededError, while the Postgres
spend stayed far below the limit. The failure path was unaffected because
it releases `user_api_key_dict.budget_reservation` directly.

The reservation is now set alongside the other internal keys, after the
client-supplied metadata merge, so a request body cannot forge one that
names arbitrary counter keys.
This commit is contained in:
Yassin Kortam 2026-08-12 12:34:13 -07:00 committed by GitHub
parent 98a79ccf92
commit 258fe3e4ba
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 117 additions and 0 deletions

View file

@ -557,6 +557,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
# real parent span.
_metadata["user_api_key"] = user_api_key_dict.api_key
_metadata["litellm_parent_otel_span"] = user_api_key_dict.parent_otel_span
_metadata["user_api_key_budget_reservation"] = user_api_key_dict.budget_reservation
_metadata.update(
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
)

View file

@ -4877,3 +4877,119 @@ async def test_unusable_upstream_cost_records_zero_not_the_flat_estimate():
assert len(payloads) == 1
assert payloads[0]["response_cost"] == 0.0
assert payloads[0]["total_tokens"] == 1874
def _passthrough_kwargs_for_reservation(
user_api_key_dict: UserAPIKeyAuth, parsed_body: Optional[dict] = None
) -> dict:
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = (
"http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
)
mock_request.headers = Headers({})
return HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
request=mock_request,
user_api_key_dict=user_api_key_dict,
passthrough_logging_payload=MagicMock(),
logging_obj=MagicMock(),
_parsed_body=parsed_body if parsed_body is not None else {},
litellm_call_id="lit-5425-call-id",
)
async def _track_cost_for_passthrough_kwargs(kwargs: dict) -> AsyncMock:
from datetime import datetime
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
callback_kwargs = {
**kwargs,
"stream": False,
"standard_logging_object": {
"response_cost": 0.002,
"request_tags": None,
},
}
increment_spend_counters = AsyncMock()
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
patch(
"litellm.proxy.proxy_server.increment_spend_counters",
increment_spend_counters,
),
patch("litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock),
):
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock()
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
await _ProxyDBLogger()._PROXY_track_cost_callback(
kwargs=callback_kwargs,
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
return increment_spend_counters
@pytest.mark.asyncio
async def test_passthrough_success_reconciles_budget_reservation():
"""
A successful pass-through request must hand its pre-call budget reservation
to the spend-counter update so the reserved amount is reconciled down to the
actual cost. Without it the reservation stays in the shared Redis counter and
the actual cost is added on top, so the counter drifts above real spend until
the key falsely trips BudgetExceededError.
"""
budget_reservation = {
"reserved_cost": 0.5,
"entries": [{"counter_key": "spend:key:hashed-token", "reserved_cost": 0.5}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-token",
user_id="u1",
budget_reservation=budget_reservation,
)
reservation = user_api_key_dict.budget_reservation
kwargs = _passthrough_kwargs_for_reservation(user_api_key_dict)
assert (
kwargs["litellm_params"]["metadata"]["user_api_key_budget_reservation"]
is reservation
)
increment_spend_counters = await _track_cost_for_passthrough_kwargs(kwargs)
increment_spend_counters.assert_awaited_once()
assert increment_spend_counters.await_args.kwargs["budget_reservation"] is reservation
assert increment_spend_counters.await_args.kwargs["budget_reservation"] == budget_reservation
@pytest.mark.asyncio
async def test_passthrough_body_cannot_forge_budget_reservation():
"""
The reservation is an internal counter handle: a client-supplied metadata
field naming arbitrary counter keys must never reach the spend-counter
update, or a caller could decrement another entity's Redis counter.
"""
forged = {
"reserved_cost": 99.0,
"entries": [{"counter_key": "spend:team:victim", "reserved_cost": 99.0}],
}
user_api_key_dict = UserAPIKeyAuth(api_key="hashed-token", user_id="u1")
kwargs = _passthrough_kwargs_for_reservation(
user_api_key_dict,
parsed_body={"litellm_metadata": {"user_api_key_budget_reservation": forged}},
)
assert (
kwargs["litellm_params"]["metadata"]["user_api_key_budget_reservation"] is None
)
increment_spend_counters = await _track_cost_for_passthrough_kwargs(kwargs)
increment_spend_counters.assert_awaited_once()
assert increment_spend_counters.await_args.kwargs["budget_reservation"] is None