diff --git a/.circleci/config.yml b/.circleci/config.yml index feb425a38e0..93470ada8a3 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -562,6 +562,7 @@ jobs: pip install "Pillow==10.3.0" pip install "jsonschema==4.22.0" pip install "pytest-postgresql==7.0.1" + pip install "fakeredis==2.28.1" - save_cache: paths: - ./venv diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 61ea930387d..57b5280b505 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -444,6 +444,17 @@ class DBSpendUpdateWriter: proxy_logging_obj=proxy_logging_obj, daily_spend_transactions=daily_team_spend_update_transactions, ) + + daily_tag_spend_update_transactions = ( + await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer() + ) + if daily_tag_spend_update_transactions is not None: + await DBSpendUpdateWriter.update_daily_tag_spend( + n_retry_times=n_retry_times, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + daily_spend_transactions=daily_tag_spend_update_transactions, + ) except Exception as e: verbose_proxy_logger.error(f"Error committing spend updates: {e}") finally: diff --git a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py index 03bd9dca9ee..b08a8517ed4 100644 --- a/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py +++ b/litellm/proxy/db/db_transaction_queue/redis_update_buffer.py @@ -19,6 +19,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import ( + DailyTagSpendTransaction, DailyTeamSpendTransaction, DailyUserSpendTransaction, DBSpendUpdateTransactions, @@ -60,9 +61,9 @@ class RedisUpdateBuffer: """ from litellm.proxy.proxy_server import general_settings - _use_redis_transaction_buffer: Optional[ - Union[bool, str] - ] = general_settings.get("use_redis_transaction_buffer", False) + _use_redis_transaction_buffer: Optional[Union[bool, str]] = ( + general_settings.get("use_redis_transaction_buffer", False) + ) if isinstance(_use_redis_transaction_buffer, str): _use_redis_transaction_buffer = str_to_bool(_use_redis_transaction_buffer) if _use_redis_transaction_buffer is None: @@ -176,14 +177,6 @@ class RedisUpdateBuffer: "ALL DAILY SPEND UPDATE TRANSACTIONS: %s", daily_spend_update_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 - - # Store all transaction types using the helper method await self._store_transactions_in_redis( transactions=db_spend_update_transactions, redis_key=REDIS_UPDATE_BUFFER_KEY, @@ -336,6 +329,30 @@ class RedisUpdateBuffer: ), ) + async def get_all_daily_tag_spend_update_transactions_from_redis_buffer( + self, + ) -> Optional[Dict[str, DailyTagSpendTransaction]]: + """ + Gets all the daily tag spend update transactions from Redis + """ + if self.redis_cache is None: + return None + list_of_transactions = await self.redis_cache.async_lpop( + key=REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY, + count=MAX_REDIS_BUFFER_DEQUEUE_COUNT, + ) + if list_of_transactions is None: + return None + list_of_daily_spend_update_transactions = [ + json.loads(transaction) for transaction in list_of_transactions + ] + return cast( + Dict[str, DailyTagSpendTransaction], + DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions( + list_of_daily_spend_update_transactions + ), + ) + @staticmethod def _parse_list_of_transactions( list_of_transactions: Union[Any, List[Any]], diff --git a/tests/proxy_unit_tests/test_e2e_pod_lock_manager.py b/tests/proxy_unit_tests/test_e2e_pod_lock_manager.py index 061da8c186c..1ade9948a15 100644 --- a/tests/proxy_unit_tests/test_e2e_pod_lock_manager.py +++ b/tests/proxy_unit_tests/test_e2e_pod_lock_manager.py @@ -2,6 +2,7 @@ import os import sys import traceback import uuid +from typing import List from datetime import datetime, timezone, timedelta from dotenv import load_dotenv @@ -9,11 +10,12 @@ from fastapi import Request from fastapi.routing import APIRoute import httpx import json - +from unittest.mock import MagicMock, patch load_dotenv() import io import os import time +import fakeredis # this file is to test litellm/proxy @@ -337,4 +339,63 @@ async def test_release_expired_lock(): # Verify that second pod's lock is still active lock_record = await global_redis_cache.async_get_cache(lock_key) - assert lock_record == second_lock_manager.pod_id \ No newline at end of file + assert lock_record == second_lock_manager.pod_id + +@pytest.mark.asyncio +async def test_e2e_size_of_redis_buffer(): + """ + Ensure that all elements from the redis queue's get flushed to the DB + + Goal of this is to ensure Redis does not blow up in size + """ + from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter + from litellm.proxy.db.db_transaction_queue.base_update_queue import BaseUpdateQueue + from litellm.caching import RedisCache + import uuid + + + redis_cache = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT"), password=os.getenv("REDIS_PASSWORD")) + fake_redis_client = fakeredis.FakeAsyncRedis() + redis_cache.redis_async_client = fake_redis_client + + setattr(litellm.proxy.proxy_server, "use_redis_transaction_buffer", True) + db_writer = DBSpendUpdateWriter(redis_cache=redis_cache) + + # get all the queues + initialized_queues: List[BaseUpdateQueue] = [] + for attr in dir(db_writer): + if isinstance(getattr(db_writer, attr), BaseUpdateQueue): + initialized_queues.append(getattr(db_writer, attr)) + + # add mock data to each queue + new_keys_added = [] + for queue in initialized_queues: + key = f"test_key_{queue.__class__.__name__}_{uuid.uuid4()}" + new_keys_added.append(key) + await queue.add_update({key: {"spend": 1.0}}) + + print("initialized_queues=", initialized_queues) + print("new_keys_added=", new_keys_added) + + # get the size of each queue + for queue in initialized_queues: + assert queue.update_queue.qsize() == 1, f"Queue {queue.__class__.__name__} was not initialized with mock data. Expected size 1, got {queue.update_queue.qsize()}" + + + # flush from in-memory -> redis -> to DB + with patch("litellm.proxy.db.db_spend_update_writer.PodLockManager.acquire_lock", return_value=True): + await db_writer._commit_spend_updates_to_db_with_redis( + prisma_client=MagicMock(), + n_retry_times=3, + proxy_logging_obj=MagicMock() + ) + + # Verify all the keys were looked up in Redis + keys = await fake_redis_client.keys("*") + print("found keys even after flushing to DB", keys) + assert len(keys) == 0, f"Expected Redis to be empty, but found keys: {keys}" + + + + +