mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
[Reliability fix] Redis transaction buffer - ensure all redis queues are periodically flushed (#10393)
* test_e2e_size_of_redis_buffer * fix store_in_memory_spend_updates_in_redis * _commit_spend_updates_to_db_with_redis * daily_tag_spend_update_transactions * pip install fakeredis==2.28.1
This commit is contained in:
parent
72de453cc0
commit
f984089b01
4 changed files with 103 additions and 13 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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]],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
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}"
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue