use spend_update_queue for RedisUpdateBuffer

This commit is contained in:
Ishaan Jaff 2025-03-31 18:40:52 -07:00
parent efe6d375e9
commit bcd49204f6

View file

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