use typed data structure for queue

This commit is contained in:
Ishaan Jaff 2025-03-31 19:28:17 -07:00
parent a753fc9d9f
commit 71e772dd4a
4 changed files with 101 additions and 73 deletions

View file

@ -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]

View file

@ -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::<value>::user_id::<value>"
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,

View file

@ -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

View file

@ -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