From 71e772dd4a5a6547c4ca97374d1ba3898e4660e9 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 31 Mar 2025 19:28:17 -0700 Subject: [PATCH] use typed data structure for queue --- litellm/proxy/_types.py | 6 ++ litellm/proxy/db/db_spend_update_writer.py | 73 +++++++++----------- litellm/proxy/db/redis_update_buffer.py | 17 ++--- litellm/proxy/db/spend_update_queue.py | 78 ++++++++++++++++------ 4 files changed, 101 insertions(+), 73 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 7f13717e299..16d302aa9a6 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2749,3 +2749,9 @@ class DBSpendUpdateTransactions(TypedDict): team_list_transactions: Optional[Dict[str, float]] team_member_list_transactions: Optional[Dict[str, float]] org_list_transactions: Optional[Dict[str, float]] + + +class SpendUpdateQueueItem(TypedDict, total=False): + entity_type: Litellm_EntityType + entity_id: str + response_cost: Optional[float] diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 5899ab3416a..5bf255feae2 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -22,6 +22,7 @@ from litellm.proxy._types import ( Litellm_EntityType, LiteLLM_UserTable, SpendLogsPayload, + SpendUpdateQueueItem, ) from litellm.proxy.db.pod_lock_manager import PodLockManager from litellm.proxy.db.redis_update_buffer import RedisUpdateBuffer @@ -148,11 +149,11 @@ class DBSpendUpdateWriter: return await self.spend_update_queue.add_update( - update={ - "entity_type": Litellm_EntityType.KEY.value, - "entity_id": hashed_token, - "amount": response_cost, - } + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.KEY, + entity_id=hashed_token, + response_cost=response_cost, + ) ) except Exception as e: verbose_proxy_logger.exception( @@ -188,20 +189,20 @@ class DBSpendUpdateWriter: for _id in user_ids: if _id is not None: await self.spend_update_queue.add_update( - update={ - "entity_type": Litellm_EntityType.USER.value, - "entity_id": _id, - "amount": response_cost, - } + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.USER, + entity_id=_id, + response_cost=response_cost, + ) ) if end_user_id is not None: await self.spend_update_queue.add_update( - update={ - "entity_type": Litellm_EntityType.END_USER.value, - "entity_id": end_user_id, - "amount": response_cost, - } + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.END_USER, + entity_id=end_user_id, + response_cost=response_cost, + ) ) except Exception as e: verbose_proxy_logger.info( @@ -224,11 +225,11 @@ class DBSpendUpdateWriter: return await self.spend_update_queue.add_update( - update={ - "entity_type": Litellm_EntityType.TEAM.value, - "entity_id": team_id, - "amount": response_cost, - } + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.TEAM, + entity_id=team_id, + response_cost=response_cost, + ) ) try: @@ -237,11 +238,11 @@ class DBSpendUpdateWriter: # key is "team_id::::user_id::" team_member_key = f"team_id::{team_id}::user_id::{user_id}" await self.spend_update_queue.add_update( - update={ - "entity_type": Litellm_EntityType.TEAM_MEMBER.value, - "entity_id": team_member_key, - "amount": response_cost, - } + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.TEAM_MEMBER, + entity_id=team_member_key, + response_cost=response_cost, + ) ) except Exception: pass @@ -265,11 +266,11 @@ class DBSpendUpdateWriter: return await self.spend_update_queue.add_update( - update={ - "entity_type": Litellm_EntityType.ORGANIZATION.value, - "entity_id": org_id, - "amount": response_cost, - } + update=SpendUpdateQueueItem( + entity_type=Litellm_EntityType.ORGANIZATION, + entity_id=org_id, + response_cost=response_cost, + ) ) except Exception as e: verbose_proxy_logger.info( @@ -424,16 +425,8 @@ class DBSpendUpdateWriter: Note: This flow causes Deadlocks in production (1K RPS+). Use self._commit_spend_updates_to_db_with_redis() instead if you expect 1K+ RPS. """ - aggregated_updates = ( - await self.spend_update_queue.flush_and_get_all_aggregated_updates_by_entity_type() - ) - db_spend_update_transactions = DBSpendUpdateTransactions( - user_list_transactions=aggregated_updates.get("user", {}), - end_user_list_transactions=aggregated_updates.get("end_user", {}), - key_list_transactions=aggregated_updates.get("key", {}), - team_list_transactions=aggregated_updates.get("team", {}), - team_member_list_transactions=aggregated_updates.get("team_member", {}), - org_list_transactions=aggregated_updates.get("organization", {}), + db_spend_update_transactions = ( + await self.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() ) await self._commit_spend_updates_to_db( prisma_client=prisma_client, diff --git a/litellm/proxy/db/redis_update_buffer.py b/litellm/proxy/db/redis_update_buffer.py index 0dfaa72a16f..1a3fd3d42d1 100644 --- a/litellm/proxy/db/redis_update_buffer.py +++ b/litellm/proxy/db/redis_update_buffer.py @@ -80,20 +80,11 @@ class RedisUpdateBuffer: ) return - aggregated_updates = ( - await spend_update_queue.flush_and_get_all_aggregated_updates_by_entity_type() + db_spend_update_transactions = ( + await spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions() ) - verbose_proxy_logger.debug("ALL AGGREGATED UPDATES: %s", aggregated_updates) - - db_spend_update_transactions: DBSpendUpdateTransactions = ( - DBSpendUpdateTransactions( - user_list_transactions=aggregated_updates.get("user", {}), - end_user_list_transactions=aggregated_updates.get("end_user", {}), - key_list_transactions=aggregated_updates.get("key", {}), - team_list_transactions=aggregated_updates.get("team", {}), - team_member_list_transactions=aggregated_updates.get("team_member", {}), - org_list_transactions=aggregated_updates.get("organization", {}), - ) + verbose_proxy_logger.debug( + "ALL DB SPEND UPDATE TRANSACTIONS: %s", db_spend_update_transactions ) # only store in redis if there are any updates to commit diff --git a/litellm/proxy/db/spend_update_queue.py b/litellm/proxy/db/spend_update_queue.py index 2a1e336877c..2d9792ac854 100644 --- a/litellm/proxy/db/spend_update_queue.py +++ b/litellm/proxy/db/spend_update_queue.py @@ -2,6 +2,11 @@ import asyncio from typing import TYPE_CHECKING, Any, Dict, List from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import ( + DBSpendUpdateTransactions, + Litellm_EntityType, + SpendUpdateQueueItem, +) if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient @@ -17,40 +22,73 @@ class SpendUpdateQueue: def __init__( self, ): - self.update_queue: asyncio.Queue[Dict[str, Any]] = asyncio.Queue() + self.update_queue: asyncio.Queue[SpendUpdateQueueItem] = asyncio.Queue() - async def add_update(self, update: Dict[str, Any]) -> None: + async def add_update(self, update: SpendUpdateQueueItem) -> None: """Enqueue an update. Each update might be a dict like {'entity_type': 'user', 'entity_id': '123', 'amount': 1.2}.""" verbose_proxy_logger.debug("Adding update to queue: %s", update) await self.update_queue.put(update) - async def flush_all_updates_from_in_memory_queue(self) -> List[Dict[str, Any]]: + async def flush_all_updates_from_in_memory_queue( + self, + ) -> List[SpendUpdateQueueItem]: """Get all updates from the queue.""" - updates: List[Dict[str, Any]] = [] + updates: List[SpendUpdateQueueItem] = [] while not self.update_queue.empty(): updates.append(await self.update_queue.get()) return updates - async def flush_and_get_all_aggregated_updates_by_entity_type( + async def flush_and_get_aggregated_db_spend_update_transactions( self, - ) -> Dict[str, Any]: + ) -> DBSpendUpdateTransactions: """Flush all updates from the queue and return all updates aggregated by entity type.""" updates = await self.flush_all_updates_from_in_memory_queue() verbose_proxy_logger.debug("Aggregating updates by entity type: %s", updates) - return self.aggregate_updates_by_entity_type(updates) + return self.get_aggregated_db_spend_update_transactions(updates) - def aggregate_updates_by_entity_type( - self, updates: List[Dict[str, Any]] - ) -> Dict[str, Any]: + def get_aggregated_db_spend_update_transactions( + self, updates: List[SpendUpdateQueueItem] + ) -> DBSpendUpdateTransactions: """Aggregate updates by entity type.""" - aggregated_updates = {} + # Initialize all transaction lists as empty dicts + db_spend_update_transactions = DBSpendUpdateTransactions( + user_list_transactions={}, + end_user_list_transactions={}, + key_list_transactions={}, + team_list_transactions={}, + team_member_list_transactions={}, + org_list_transactions={}, + ) + + # Map entity types to their corresponding transaction dictionary keys + entity_type_to_dict_key = { + Litellm_EntityType.USER: "user_list_transactions", + Litellm_EntityType.END_USER: "end_user_list_transactions", + Litellm_EntityType.KEY: "key_list_transactions", + Litellm_EntityType.TEAM: "team_list_transactions", + Litellm_EntityType.TEAM_MEMBER: "team_member_list_transactions", + Litellm_EntityType.ORGANIZATION: "org_list_transactions", + } + for update in updates: - entity_type = update["entity_type"] - entity_id = update["entity_id"] - amount = update["amount"] - if entity_type not in aggregated_updates: - aggregated_updates[entity_type] = {} - if entity_id not in aggregated_updates[entity_type]: - aggregated_updates[entity_type][entity_id] = 0 - aggregated_updates[entity_type][entity_id] += amount - return aggregated_updates + entity_type = update.get("entity_type") + entity_id = update.get("entity_id") + response_cost = update.get("response_cost") + + if entity_type is None or entity_id is None or response_cost is None: + raise ValueError("Invalid update: %s", update) + + dict_key = entity_type_to_dict_key.get(entity_type) + if dict_key is None: + continue # Skip unknown entity types + + transactions_dict = db_spend_update_transactions[dict_key] + if transactions_dict is None: + transactions_dict = {} + db_spend_update_transactions[dict_key] = transactions_dict + + transactions_dict[entity_id] = ( + transactions_dict.get(entity_id, 0) + response_cost + ) + + return db_spend_update_transactions