mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
use typed data structure for queue
This commit is contained in:
parent
a753fc9d9f
commit
71e772dd4a
4 changed files with 101 additions and 73 deletions
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue