diff --git a/litellm/proxy/db/redis_update_buffer.py b/litellm/proxy/db/redis_update_buffer.py index f98fc9300f8..7370b0ed08f 100644 --- a/litellm/proxy/db/redis_update_buffer.py +++ b/litellm/proxy/db/redis_update_buffer.py @@ -12,6 +12,7 @@ from litellm.caching import RedisCache from litellm.constants import MAX_REDIS_BUFFER_DEQUEUE_COUNT, REDIS_UPDATE_BUFFER_KEY from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import DBSpendUpdateTransactions +from litellm.proxy.db.spend_update_queue import SpendUpdateQueue from litellm.secret_managers.main import str_to_bool if TYPE_CHECKING: @@ -32,6 +33,7 @@ class RedisUpdateBuffer: redis_cache: Optional[RedisCache] = None, ): self.redis_cache = redis_cache + self.spend_update_queue = SpendUpdateQueue() @staticmethod def _should_commit_spend_updates_to_redis() -> bool: @@ -54,7 +56,6 @@ class RedisUpdateBuffer: async def store_in_memory_spend_updates_in_redis( self, - prisma_client: PrismaClient, ): """ Stores the in-memory spend updates to Redis @@ -78,13 +79,21 @@ class RedisUpdateBuffer: "redis_cache is None, skipping store_in_memory_spend_updates_in_redis" ) return - db_spend_update_transactions: DBSpendUpdateTransactions = DBSpendUpdateTransactions( - user_list_transactions=prisma_client.user_list_transactions, - end_user_list_transactions=prisma_client.end_user_list_transactions, - key_list_transactions=prisma_client.key_list_transactions, - team_list_transactions=prisma_client.team_list_transactions, - team_member_list_transactions=prisma_client.team_member_list_transactions, - org_list_transactions=prisma_client.org_list_transactions, + + aggregated_updates = ( + await self.spend_update_queue.flush_and_get_all_aggregated_updates_by_entity_type() + ) + verbose_proxy_logger.debug("ALL AGGREGATED UPDATES: ", 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("org", {}), + ) ) # only store in redis if there are any updates to commit @@ -100,9 +109,6 @@ class RedisUpdateBuffer: values=list_of_transactions, ) - # clear the in-memory spend updates - RedisUpdateBuffer._clear_all_in_memory_spend_updates(prisma_client) - @staticmethod def _number_of_transactions_to_store_in_redis( db_spend_update_transactions: DBSpendUpdateTransactions, @@ -116,20 +122,6 @@ class RedisUpdateBuffer: num_transactions += len(v) return num_transactions - @staticmethod - def _clear_all_in_memory_spend_updates( - prisma_client: PrismaClient, - ): - """ - Clears all in-memory spend updates - """ - prisma_client.user_list_transactions = {} - prisma_client.end_user_list_transactions = {} - prisma_client.key_list_transactions = {} - prisma_client.team_list_transactions = {} - prisma_client.team_member_list_transactions = {} - prisma_client.org_list_transactions = {} - @staticmethod def _remove_prefix_from_keys(data: Dict[str, Any], prefix: str) -> Dict[str, Any]: """