fix budget tag spend counter reconciliation

This commit is contained in:
user
2026-04-30 16:09:38 -07:00
parent 96a283ed0f
commit 1373ae1021
4 changed files with 160 additions and 21 deletions
@@ -178,17 +178,18 @@ class _ProxyDBLogger(CustomLogger):
sl_object: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object", None
)
response_cost = (
sl_object.get("response_cost", None)
if sl_object is not None
else kwargs.get("response_cost", None)
)
tags: Optional[List[str]] = (
sl_object.get("request_tags", None) if sl_object is not None else None
)
if response_cost is not None:
user_api_key = metadata.get("user_api_key", None)
response_cost = (
sl_object.get("response_cost", None)
if sl_object is not None
else kwargs.get("response_cost", None)
)
tags = _get_request_tags_for_cost_tracking(
sl_object=sl_object,
metadata=metadata,
)
if response_cost is not None:
user_api_key = metadata.get("user_api_key", None)
if kwargs.get("cache_hit", False) is True:
response_cost = 0.0
verbose_proxy_logger.debug(
@@ -219,6 +220,7 @@ class _ProxyDBLogger(CustomLogger):
end_time=end_time,
response_cost=response_cost,
budget_reservation=budget_reservation,
request_tags=tags,
)
# update cache (fire-and-forget for backward compat:
@@ -407,6 +409,22 @@ def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]:
return getattr(user_api_key_auth_obj, "budget_reservation", None)
def _get_request_tags_for_cost_tracking(
sl_object: Optional[StandardLoggingPayload],
metadata: dict,
) -> Optional[List[str]]:
if sl_object is not None:
request_tags = sl_object.get("request_tags", None)
if isinstance(request_tags, list):
return request_tags
metadata_tags = metadata.get("tags", None)
if isinstance(metadata_tags, list):
return metadata_tags
return None
async def _update_database_and_spend_counters(
proxy_logging_obj: Any,
increment_spend_counters: Any,
@@ -421,6 +439,7 @@ async def _update_database_and_spend_counters(
end_time: Any,
response_cost: float,
budget_reservation: Optional[dict],
request_tags: Optional[List[str]] = None,
) -> None:
try:
await proxy_logging_obj.db_spend_update_writer.update_database(
@@ -461,6 +480,8 @@ async def _update_database_and_spend_counters(
response_cost=response_cost,
org_id=org_id,
budget_reservation=budget_reservation,
end_user_id=end_user_id,
tags=request_tags,
)
except Exception:
if budget_reservation is not None:
+88 -9
View File
@@ -1836,6 +1836,8 @@ async def increment_spend_counters(
response_cost: Optional[float],
org_id: Optional[str] = None,
budget_reservation: Optional[dict] = None,
end_user_id: Optional[str] = None,
tags: Optional[List[str]] = None,
):
"""
Atomically increment spend counters for budget enforcement.
@@ -1847,7 +1849,7 @@ async def increment_spend_counters(
Awaited (not create_task) in the cost callback, so the counter is
updated before the next request's auth check runs.
"""
reserved_counter_keys = set()
reserved_counter_keys: Set[str] = set()
if budget_reservation is not None:
from litellm.proxy.spend_tracking.budget_reservation import (
get_reserved_counter_keys,
@@ -1970,14 +1972,91 @@ async def increment_spend_counters(
increment=response_cost,
)
if org_id is not None:
org_counter_key = f"spend:org:{org_id}"
if org_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=org_counter_key,
source_cache_key=f"org_id:{org_id}:with_budget",
increment=response_cost,
)
await _increment_end_user_and_tag_spend_counters(
end_user_id=end_user_id,
tags=tags,
response_cost=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
await _increment_org_spend_counter(
org_id=org_id,
response_cost=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
async def _increment_end_user_and_tag_spend_counters(
end_user_id: Optional[str],
tags: Optional[List[str]],
response_cost: float,
reserved_counter_keys: Set[str],
) -> None:
if end_user_id is not None:
await _increment_warm_unreserved_spend_counter(
counter_key=f"spend:end_user:{end_user_id}",
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
if tags is None:
return
seen_tags: Set[str] = set()
for tag_name in tags:
if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags:
continue
seen_tags.add(tag_name)
await _increment_warm_unreserved_spend_counter(
counter_key=f"spend:tag:{tag_name}",
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
async def _increment_warm_unreserved_spend_counter(
counter_key: str,
increment: float,
reserved_counter_keys: Set[str],
) -> None:
if counter_key in reserved_counter_keys:
return
if await spend_counter_cache.async_get_cache(key=counter_key) is None:
return
await _increment_spend_counter_cache(counter_key=counter_key, increment=increment)
async def _increment_org_spend_counter(
org_id: Optional[str],
response_cost: float,
reserved_counter_keys: Set[str],
) -> None:
if org_id is None:
return
await _init_and_increment_unreserved_spend_counter(
counter_key=f"spend:org:{org_id}",
source_cache_key=f"org_id:{org_id}:with_budget",
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
async def _init_and_increment_unreserved_spend_counter(
counter_key: str,
source_cache_key: str,
increment: float,
reserved_counter_keys: Set[str],
) -> None:
if counter_key in reserved_counter_keys:
return
await _init_and_increment_spend_counter(
counter_key=counter_key,
source_cache_key=source_cache_key,
increment=increment,
)
async def _init_and_increment_spend_counter(
@@ -391,7 +391,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
increment_spend_counters=increment_spend_counters,
user_api_key="test_api_key",
user_id="test_user_id",
end_user_id=None,
end_user_id="test_end_user_id",
team_id="test_team_id",
org_id="test_org_id",
kwargs={},
@@ -400,6 +400,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
end_time=datetime.now(),
response_cost=0.2,
budget_reservation=budget_reservation,
request_tags=["tag-a"],
)
proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once()
@@ -410,6 +411,8 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
response_cost=0.2,
org_id="test_org_id",
budget_reservation=budget_reservation,
end_user_id="test_end_user_id",
tags=["tag-a"],
)
@@ -184,6 +184,7 @@ async def test_should_prevent_second_end_user_reservation_over_budget(
user_id=None,
response_cost=0.2,
budget_reservation=reservation,
end_user_id="end-user-budget-race",
)
assert counter_cache.in_memory_cache.get_cache(
@@ -278,6 +279,7 @@ async def test_should_prevent_second_tag_reservation_over_budget(
user_id=None,
response_cost=0.2,
budget_reservation=reservation,
tags=["tag-budget-race"],
)
assert counter_cache.in_memory_cache.get_cache(
@@ -285,6 +287,40 @@ async def test_should_prevent_second_tag_reservation_over_budget(
) == pytest.approx(0.2)
@pytest.mark.asyncio
async def test_should_update_warm_end_user_and_tag_counters_without_reservation(
spend_counter_state,
):
counter_cache, _ = spend_counter_state
counter_cache.in_memory_cache.set_cache(
key="spend:end_user:customer-1",
value=4.0,
)
counter_cache.in_memory_cache.set_cache(key="spend:tag:paid-tag", value=7.0)
counter_cache.in_memory_cache.set_cache(key="spend:tag:other-tag", value=2.0)
from litellm.proxy.proxy_server import increment_spend_counters
await increment_spend_counters(
token=None,
team_id=None,
user_id=None,
response_cost=0.50,
end_user_id="customer-1",
tags=["paid-tag", "paid-tag", "other-tag", ""],
)
assert counter_cache.in_memory_cache.get_cache(
key="spend:end_user:customer-1"
) == pytest.approx(4.50)
assert counter_cache.in_memory_cache.get_cache(
key="spend:tag:paid-tag"
) == pytest.approx(7.50)
assert counter_cache.in_memory_cache.get_cache(
key="spend:tag:other-tag"
) == pytest.approx(2.50)
@pytest.mark.asyncio
async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_state):
counter_cache, key_cache = spend_counter_state