diff --git a/litellm/proxy/db/redis_update_buffer.py b/litellm/proxy/db/redis_update_buffer.py index ab43dc575d..3d27e41d15 100644 --- a/litellm/proxy/db/redis_update_buffer.py +++ b/litellm/proxy/db/redis_update_buffer.py @@ -80,30 +80,35 @@ 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, + async with prisma_client.in_memory_transaction_lock: + 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, + ) ) - ) - # only store in redis if there are any updates to commit - if ( - self._number_of_transactions_to_store_in_redis(db_spend_update_transactions) - == 0 - ): - return + # only store in redis if there are any updates to commit + if ( + self._number_of_transactions_to_store_in_redis( + db_spend_update_transactions + ) + == 0 + ): + return - list_of_transactions = [safe_dumps(db_spend_update_transactions)] - await self.redis_cache.async_rpush( - key=REDIS_UPDATE_BUFFER_KEY, - values=list_of_transactions, - ) - self._clear_all_in_memory_spend_updates(prisma_client) + list_of_transactions = [safe_dumps(db_spend_update_transactions)] + await self.redis_cache.async_rpush( + key=REDIS_UPDATE_BUFFER_KEY, + 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( @@ -201,12 +206,15 @@ class RedisUpdateBuffer: @staticmethod def _parse_list_of_transactions( - list_of_transactions: List[str], + list_of_transactions: Union[Any, List[Any]], ) -> List[DBSpendUpdateTransactions]: """ Parses the list of transactions from Redis """ - return [json.loads(transaction) for transaction in list_of_transactions] + if isinstance(list_of_transactions, list): + return [json.loads(transaction) for transaction in list_of_transactions] + else: + return [json.loads(list_of_transactions)] @staticmethod def _combine_list_of_transactions( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index d0d4ea9b4e..66dd224dff 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1121,6 +1121,7 @@ class PrismaClient: self.iam_token_db_auth: Optional[bool] = str_to_bool( os.getenv("IAM_TOKEN_DB_AUTH") ) + self.in_memory_transaction_lock = asyncio.Lock() verbose_proxy_logger.debug("Creating Prisma Client..") try: from prisma import Prisma # type: ignore