mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-11 16:26:07 +00:00
fix budget tag spend counter reconciliation
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user