diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 6e274f16ce..a61fda91fd 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 10d4a8fb07..2f2b59db67 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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( diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 65f04c290d..9f414f560c 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -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"], ) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 7ed6278808..481d88b8d0 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -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