mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(proxy): settle folded tag reservations to their own input cost and roll back only acquired counters
A hook-added tag is priced after the hook rewrote the request, so its entries keep that input cost when they join the auth-time reservation and a cancellation settles each entry to the cost it was priced with. An entry counts as applied only once its counter increment completed, so a cancellation during the increment no longer refunds spend that was never reserved.
This commit is contained in:
parent
f85cb1dac5
commit
a690b5a543
5 changed files with 82 additions and 12 deletions
|
|
@ -77,7 +77,7 @@ from litellm.proxy.common_utils.sse_keepalive import (
|
|||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.spend_tracking.budget_reservation import reserve_budget_for_added_tags
|
||||
from litellm.proxy.spend_tracking.budget_reservation import merge_budget_reservation, reserve_budget_for_added_tags
|
||||
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
|
||||
from litellm.router import Router
|
||||
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
|
||||
|
|
@ -2156,7 +2156,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
return
|
||||
existing: Final = user_api_key_dict.budget_reservation
|
||||
if existing is not None:
|
||||
existing["entries"].extend(reservation["entries"])
|
||||
merge_budget_reservation(existing=existing, added=reservation)
|
||||
return
|
||||
user_api_key_dict.budget_reservation = reservation # rebind-ok: the failure and cancel paths read it here
|
||||
_, metadata_bucket = get_or_create_metadata_bucket(self.data)
|
||||
|
|
|
|||
|
|
@ -325,7 +325,6 @@ async def _reserve_counters(
|
|||
counter=counter,
|
||||
reserved_cost=reservation_cost,
|
||||
)
|
||||
applied_entries.append(entry)
|
||||
try:
|
||||
reserved_value = await _reserve_counter(
|
||||
counter=counter,
|
||||
|
|
@ -337,11 +336,11 @@ async def _reserve_counters(
|
|||
entries=[entry],
|
||||
default_reserved_cost=reservation_cost,
|
||||
)
|
||||
applied_entries.remove(entry)
|
||||
if fail_closed_budget_enforcement:
|
||||
_raise_reservation_unavailable(counter_key=counter.counter_key)
|
||||
continue
|
||||
|
||||
applied_entries.append(entry)
|
||||
if reserved_value is not None:
|
||||
current_spend = reserved_value
|
||||
else:
|
||||
|
|
@ -434,15 +433,47 @@ async def release_budget_reservation_on_cancel(
|
|||
"""
|
||||
if not budget_reservation or budget_reservation.get("finalized") is True:
|
||||
return
|
||||
incurred_cost: Final = float(budget_reservation.get("input_cost") or 0.0)
|
||||
try:
|
||||
await asyncio.shield(
|
||||
reconcile_budget_reservation(budget_reservation=budget_reservation, actual_cost=incurred_cost)
|
||||
)
|
||||
await asyncio.shield(_reconcile_entries_to_their_input_cost(budget_reservation=budget_reservation))
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
|
||||
|
||||
async def _reconcile_entries_to_their_input_cost(
|
||||
budget_reservation: dict[str, object], # mutable-ok: the reservation is stamped finalized in place
|
||||
) -> None:
|
||||
"""Entries folded in by ``merge_budget_reservation`` carry the input cost of the request that
|
||||
went upstream; the rest were priced at auth and settle to the reservation's own input cost."""
|
||||
shared_input_cost: Final = float(cast(SupportsFloat, budget_reservation.get("input_cost") or 0.0))
|
||||
reserved_cost: Final = float(cast(SupportsFloat, budget_reservation.get("reserved_cost") or 0.0))
|
||||
entries: Final = cast(list[dict[str, float | str]], budget_reservation.get("entries") or [])
|
||||
for input_cost in dict.fromkeys(_entry_input_cost(entry, shared_input_cost) for entry in entries):
|
||||
await _set_reserved_entries_actual_cost(
|
||||
entries=[entry for entry in entries if _entry_input_cost(entry, shared_input_cost) == input_cost],
|
||||
actual_cost=input_cost,
|
||||
default_reserved_cost=reserved_cost,
|
||||
)
|
||||
budget_reservation["finalized"] = True # rebind-ok: the settlement paths share this one dict
|
||||
|
||||
|
||||
def _entry_input_cost(entry: Mapping[str, float | str], shared_input_cost: float) -> float:
|
||||
return float(entry.get("input_cost", shared_input_cost))
|
||||
|
||||
|
||||
def merge_budget_reservation(
|
||||
existing: dict[str, object], # mutable-ok: the request's reservation, extended in place
|
||||
added: Mapping[str, object],
|
||||
) -> None:
|
||||
"""Fold a later reservation's entries into the request's reservation so every settlement path sees one.
|
||||
Each added entry keeps the input cost it was priced with, since a pre-call hook may have rewritten the
|
||||
request between the two estimates and a cancellation settles each entry to that cost."""
|
||||
added_entries: Final = cast(list[dict[str, float | str]], added.get("entries") or [])
|
||||
added_input_cost: Final = float(cast(SupportsFloat, added.get("input_cost") or 0.0))
|
||||
for entry in added_entries:
|
||||
entry["input_cost"] = added_input_cost # rebind-ok: the entry settles to the cost it was priced with
|
||||
cast(list[dict[str, float | str]], existing["entries"]).extend(added_entries) # rebind-ok: one dict for all paths
|
||||
|
||||
|
||||
async def invalidate_budget_reservation_counters(
|
||||
budget_reservation: dict | None,
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -2517,6 +2517,8 @@ def test_route_skips_budget_checks_matches_auth_scope(route, expected):
|
|||
False,
|
||||
False,
|
||||
),
|
||||
("/bria", {"pass_through_endpoints": "not-a-list"}, "sk-master", False, False),
|
||||
("/bria", {"pass_through_endpoints": [{"path": "/bria", "auth": False}, "junk"]}, "sk-master", False, False),
|
||||
],
|
||||
)
|
||||
def test_auth_skips_common_checks_names_the_requests_that_never_run_them(
|
||||
|
|
|
|||
|
|
@ -18,6 +18,8 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
count_request_input_tokens,
|
||||
estimate_request_max_cost,
|
||||
merge_budget_reservation,
|
||||
release_budget_reservation_on_cancel,
|
||||
reserve_budget_for_added_tags,
|
||||
reserve_budget_for_request,
|
||||
)
|
||||
|
|
@ -237,6 +239,35 @@ async def test_reserve_budget_for_added_tags_skips_routes_auth_never_reserves(sp
|
|||
assert spend_counter_cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_release_budget_reservation_on_cancel_settles_each_entry_to_its_own_input_cost(
|
||||
spend_counter_cache: DualCache,
|
||||
):
|
||||
"""Auth priced the key before the hook rewrote the prompt; the hook tag was priced after. A cancellation
|
||||
charges each counter the input cost its own estimate saw instead of the auth-time one for both."""
|
||||
spend_counter_cache.in_memory_cache.set_cache(key="spend:key:hashed-hook-tag-key", value=0.5)
|
||||
auth_reservation: Final[dict[str, object]] = {
|
||||
"reserved_cost": 0.5,
|
||||
"entries": [{"counter_key": "spend:key:hashed-hook-tag-key", "entity_type": "Key", "reserved_cost": 0.5}],
|
||||
"finalized": False,
|
||||
"input_cost": 0.2,
|
||||
"input_tokens": 40,
|
||||
}
|
||||
hook_reservation: Final = await _reserve_added_tags(
|
||||
"/v1/chat/completions", _budgeted_tag_prisma((HOOK_TAG,), max_budget=1.0)
|
||||
)
|
||||
assert hook_reservation is not None
|
||||
hook_input_cost: Final = hook_reservation["input_cost"]
|
||||
assert isinstance(hook_input_cost, float) and 0 < hook_input_cost < 0.2
|
||||
|
||||
merge_budget_reservation(existing=auth_reservation, added=hook_reservation)
|
||||
await release_budget_reservation_on_cancel(auth_reservation)
|
||||
|
||||
assert spend_counter_cache.in_memory_cache.get_cache(key="spend:key:hashed-hook-tag-key") == pytest.approx(0.2)
|
||||
assert spend_counter_cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") == pytest.approx(hook_input_cost)
|
||||
assert auth_reservation["finalized"] is True
|
||||
|
||||
|
||||
class _ParkingIncrementCache(DualCache):
|
||||
"""A spend-counter cache whose increment of ``parked_key`` never returns, so a test can cancel mid-reservation."""
|
||||
|
||||
|
|
@ -253,12 +284,15 @@ class _ParkingIncrementCache(DualCache):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.timeout(30)
|
||||
async def test_reserve_budget_for_added_tags_releases_the_reserved_tag_when_cancelled_mid_acquisition(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
"""A client disconnect under SSE keepalives cancels the request while the second tag is being reserved;
|
||||
the first tag's counter must not stay charged for a request that never reached the provider."""
|
||||
the first tag's counter must not stay charged for a request that never reached the provider, and the
|
||||
second tag, whose increment never happened, must not be refunded for it either."""
|
||||
cache: Final = _ParkingIncrementCache(parked_key=f"spend:tag:{SECOND_HOOK_TAG}")
|
||||
cache.in_memory_cache.set_cache(key=f"spend:tag:{SECOND_HOOK_TAG}", value=0.3)
|
||||
monkeypatch.setattr(proxy_server, "spend_counter_cache", cache)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
prisma: Final = _budgeted_tag_prisma((HOOK_TAG, SECOND_HOOK_TAG), max_budget=1.0)
|
||||
|
|
@ -273,6 +307,7 @@ async def test_reserve_budget_for_added_tags_releases_the_reserved_tag_when_canc
|
|||
await reserving
|
||||
|
||||
assert cache.in_memory_cache.get_cache(key=f"spend:tag:{HOOK_TAG}") == pytest.approx(0.0)
|
||||
assert cache.in_memory_cache.get_cache(key=f"spend:tag:{SECOND_HOOK_TAG}") == pytest.approx(0.3)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -656,14 +656,14 @@ class TestProxyBaseLLMRequestProcessing:
|
|||
assert tag_check.await_count == (1 if checked else 0)
|
||||
|
||||
@staticmethod
|
||||
def _reservation(counter_key: str) -> dict:
|
||||
def _reservation(counter_key: str, input_cost: float = 0.1) -> dict:
|
||||
return {
|
||||
"reserved_cost": 0.5,
|
||||
"entries": [
|
||||
{"counter_key": counter_key, "entity_type": "Tag", "entity_id": counter_key, "reserved_cost": 0.5}
|
||||
],
|
||||
"finalized": False,
|
||||
"input_cost": 0.1,
|
||||
"input_cost": input_cost,
|
||||
"input_tokens": 3,
|
||||
}
|
||||
|
||||
|
|
@ -677,7 +677,7 @@ class TestProxyBaseLLMRequestProcessing:
|
|||
mock_request, mock_proxy_logging_obj, _ = self._tag_budget_rig(
|
||||
monkeypatch, request_data={"model": "live-mini", "metadata": {"tags": []}}, pre_call_hook=mock_pre_call_hook
|
||||
)
|
||||
reserve = AsyncMock(return_value=self._reservation("spend:tag:guardrail-tag"))
|
||||
reserve = AsyncMock(return_value=self._reservation("spend:tag:guardrail-tag", input_cost=0.02))
|
||||
monkeypatch.setattr(litellm.proxy.common_request_processing, "reserve_budget_for_added_tags", reserve)
|
||||
await processing_obj.common_processing_pre_call_logic(
|
||||
request=mock_request,
|
||||
|
|
@ -715,6 +715,8 @@ class TestProxyBaseLLMRequestProcessing:
|
|||
"spend:key:test-token",
|
||||
"spend:tag:guardrail-tag",
|
||||
]
|
||||
assert [entry.get("input_cost") for entry in auth_reservation["entries"]] == [None, 0.02]
|
||||
assert auth_reservation["input_cost"] == 0.1
|
||||
assert "user_api_key_budget_reservation" not in processing_obj.data["metadata"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue