diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 0571ac9f15..7378ed878a 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -1045,3 +1045,109 @@ class RedisCache(BaseCache): except Exception as e: verbose_logger.debug(f"Redis TTL Error: {e}") return None + + async def async_rpush( + self, + key: str, + values: List[Any], + parent_otel_span: Optional[Span] = None, + **kwargs, + ) -> int: + """ + Append one or multiple values to a list stored at key + + Args: + key: The Redis key of the list + values: One or more values to append to the list + parent_otel_span: Optional parent OpenTelemetry span + + Returns: + int: The length of the list after the push operation + """ + _redis_client: Any = self.init_async_client() + start_time = time.time() + try: + response = await _redis_client.rpush(key, *values) + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_success_hook( + service=ServiceTypes.REDIS, + duration=_duration, + call_type="async_rpush", + ) + ) + return response + except Exception as e: + # NON blocking - notify users Redis is throwing an exception + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_failure_hook( + service=ServiceTypes.REDIS, + duration=_duration, + error=e, + call_type="async_rpush", + ) + ) + verbose_logger.error( + f"LiteLLM Redis Cache RPUSH: - Got exception from REDIS : {str(e)}" + ) + raise e + + async def async_lpop( + self, + key: str, + count: Optional[int] = None, + parent_otel_span: Optional[Span] = None, + **kwargs, + ) -> Union[Any, List[Any]]: + _redis_client: Any = self.init_async_client() + start_time = time.time() + print_verbose(f"LPOP from Redis list: key: {key}, count: {count}") + try: + result = await _redis_client.lpop(key, count) + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_success_hook( + service=ServiceTypes.REDIS, + duration=_duration, + call_type="async_lpop", + ) + ) + + # Handle result parsing if needed + if isinstance(result, bytes): + try: + return result.decode("utf-8") + except Exception: + return result + elif isinstance(result, list) and all( + isinstance(item, bytes) for item in result + ): + try: + return [item.decode("utf-8") for item in result] + except Exception: + return result + return result + except Exception as e: + # NON blocking - notify users Redis is throwing an exception + ## LOGGING ## + end_time = time.time() + _duration = end_time - start_time + asyncio.create_task( + self.service_logger_obj.async_service_failure_hook( + service=ServiceTypes.REDIS, + duration=_duration, + error=e, + call_type="async_lpop", + ) + ) + verbose_logger.error( + f"LiteLLM Redis Cache LPOP: - Got exception from REDIS : {str(e)}" + ) + raise e diff --git a/litellm/constants.py b/litellm/constants.py index b355a0f683..e224b3d33e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -18,6 +18,8 @@ DEFAULT_IMAGE_WIDTH = 300 DEFAULT_IMAGE_HEIGHT = 300 MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB = 1024 # 1MB = 1024KB SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD = 1000 # Minimum number of requests to consider "reasonable traffic". Used for single-deployment cooldown logic. +REDIS_UPDATE_BUFFER_KEY = "litellm_spend_update_buffer" +MAX_REDIS_BUFFER_DEQUEUE_COUNT = 100 #### RELIABILITY #### REPEATED_STREAMING_CHUNK_LIMIT = 100 # catch if model starts looping the same chunk while streaming. Uses high default to prevent false positives. #### Networking settings #### @@ -443,3 +445,7 @@ HEALTH_CHECK_TIMEOUT_SECONDS = 60 # 60 seconds UI_SESSION_TOKEN_TEAM_ID = "litellm-dashboard" LITELLM_PROXY_ADMIN_NAME = "default_user_id" + +########################### DB CRON JOB NAMES ########################### +DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job" +DEFAULT_CRON_JOB_LOCK_TTL_SECONDS = 60 # 1 minute diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 3c22aaa601..e7d6bec004 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -144,6 +144,21 @@ class LitellmTableNames(str, enum.Enum): PROXY_MODEL_TABLE_NAME = "LiteLLM_ProxyModelTable" +class Litellm_EntityType(enum.Enum): + """ + Enum for types of entities on litellm + + This enum allows specifying the type of entity that is being tracked in the database. + """ + + KEY = "key" + USER = "user" + END_USER = "end_user" + TEAM = "team" + TEAM_MEMBER = "team_member" + ORGANIZATION = "organization" + + def hash_token(token: str): import hashlib @@ -2719,3 +2734,16 @@ class DailyUserSpendTransaction(TypedDict): completion_tokens: int spend: float api_requests: int + + +class DBSpendUpdateTransactions(TypedDict): + """ + Internal Data Structure for buffering spend updates in Redis or in memory before committing them to the database + """ + + user_list_transactions: Optional[Dict[str, float]] + end_user_list_transactions: Optional[Dict[str, float]] + key_list_transactions: Optional[Dict[str, float]] + team_list_transactions: Optional[Dict[str, float]] + team_member_list_transactions: Optional[Dict[str, float]] + org_list_transactions: Optional[Dict[str, float]] diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py new file mode 100644 index 0000000000..f46b03b57a --- /dev/null +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -0,0 +1,746 @@ +""" +Module responsible for + +1. Writing spend increments to either in memory list of transactions or to redis +2. Reading increments from redis or in memory list of transactions and committing them to db +""" + +import asyncio +import os +import time +import traceback +from datetime import datetime, timedelta +from typing import TYPE_CHECKING, Any, Optional, Union + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache, RedisCache +from litellm.constants import DB_SPEND_UPDATE_JOB_NAME +from litellm.proxy._types import ( + DB_CONNECTION_ERROR_TYPES, + DBSpendUpdateTransactions, + Litellm_EntityType, + LiteLLM_UserTable, + SpendLogsPayload, +) +from litellm.proxy.db.pod_lock_manager import PodLockManager +from litellm.proxy.db.redis_update_buffer import RedisUpdateBuffer + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient, ProxyLogging +else: + PrismaClient = Any + ProxyLogging = Any + + +class DBSpendUpdateWriter: + """ + Module responsible for + + 1. Writing spend increments to either in memory list of transactions or to redis + 2. Reading increments from redis or in memory list of transactions and committing them to db + """ + + def __init__( + self, + redis_cache: Optional[RedisCache] = None, + ): + self.redis_cache = redis_cache + self.redis_update_buffer = RedisUpdateBuffer(redis_cache=self.redis_cache) + self.pod_lock_manager = PodLockManager(cronjob_id=DB_SPEND_UPDATE_JOB_NAME) + + @staticmethod + async def update_database( + # LiteLLM management object fields + token: Optional[str], + user_id: Optional[str], + end_user_id: Optional[str], + team_id: Optional[str], + org_id: Optional[str], + # Completion object fields + kwargs: Optional[dict], + completion_response: Optional[Union[litellm.ModelResponse, Any, Exception]], + start_time: Optional[datetime], + end_time: Optional[datetime], + response_cost: Optional[float], + ): + from litellm.proxy.proxy_server import ( + disable_spend_logs, + litellm_proxy_budget_name, + prisma_client, + user_api_key_cache, + ) + from litellm.proxy.utils import ProxyUpdateSpend, hash_token + + try: + verbose_proxy_logger.debug( + f"Enters prisma db call, response_cost: {response_cost}, token: {token}; user_id: {user_id}; team_id: {team_id}" + ) + if ProxyUpdateSpend.disable_spend_updates() is True: + return + if token is not None and isinstance(token, str) and token.startswith("sk-"): + hashed_token = hash_token(token=token) + else: + hashed_token = token + + asyncio.create_task( + DBSpendUpdateWriter._update_user_db( + response_cost=response_cost, + user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + litellm_proxy_budget_name=litellm_proxy_budget_name, + end_user_id=end_user_id, + ) + ) + asyncio.create_task( + DBSpendUpdateWriter._update_key_db( + response_cost=response_cost, + hashed_token=hashed_token, + prisma_client=prisma_client, + ) + ) + asyncio.create_task( + DBSpendUpdateWriter._update_team_db( + response_cost=response_cost, + team_id=team_id, + user_id=user_id, + prisma_client=prisma_client, + ) + ) + asyncio.create_task( + DBSpendUpdateWriter._update_org_db( + response_cost=response_cost, + org_id=org_id, + prisma_client=prisma_client, + ) + ) + if disable_spend_logs is False: + await DBSpendUpdateWriter._insert_spend_log_to_db( + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + response_cost=response_cost, + prisma_client=prisma_client, + ) + else: + verbose_proxy_logger.info( + "disable_spend_logs=True. Skipping writing spend logs to db. Other spend updates - Key/User/Team table will still occur." + ) + + verbose_proxy_logger.debug("Runs spend update on all tables") + except Exception: + verbose_proxy_logger.debug( + f"Error updating Prisma database: {traceback.format_exc()}" + ) + + @staticmethod + async def _update_transaction_list( + response_cost: Optional[float], + entity_id: Optional[str], + transaction_list: dict, + entity_type: Litellm_EntityType, + debug_msg: Optional[str] = None, + prisma_client: Optional[PrismaClient] = None, + ) -> bool: + """ + Common helper method to update a transaction list for an entity + + Args: + response_cost: The cost to add + entity_id: The ID of the entity to update + transaction_list: The transaction list dictionary to update + entity_type: The type of entity (from EntityType enum) + debug_msg: Optional custom debug message + + Returns: + bool: True if update happened, False otherwise + """ + try: + if debug_msg: + verbose_proxy_logger.debug(debug_msg) + else: + verbose_proxy_logger.debug( + f"adding spend to {entity_type.value} db. Response cost: {response_cost}. {entity_type.value}_id: {entity_id}." + ) + if prisma_client is None: + return False + + if entity_id is None: + verbose_proxy_logger.debug( + f"track_cost_callback: {entity_type.value}_id is None. Not tracking spend for {entity_type.value}" + ) + return False + transaction_list[entity_id] = response_cost + transaction_list.get( + entity_id, 0 + ) + return True + + except Exception as e: + verbose_proxy_logger.info( + f"Update {entity_type.value.capitalize()} DB failed to execute - {str(e)}\n{traceback.format_exc()}" + ) + raise e + + @staticmethod + async def _update_key_db( + response_cost: Optional[float], + hashed_token: Optional[str], + prisma_client: Optional[PrismaClient], + ): + try: + if hashed_token is None or prisma_client is None: + return + + await DBSpendUpdateWriter._update_transaction_list( + response_cost=response_cost, + entity_id=hashed_token, + transaction_list=prisma_client.key_list_transactions, + entity_type=Litellm_EntityType.KEY, + debug_msg=f"adding spend to key db. Response cost: {response_cost}. Token: {hashed_token}.", + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.exception( + f"Update Key DB Call failed to execute - {str(e)}" + ) + raise e + + @staticmethod + async def _update_user_db( + response_cost: Optional[float], + user_id: Optional[str], + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + litellm_proxy_budget_name: Optional[str], + end_user_id: Optional[str] = None, + ): + """ + - Update that user's row + - Update litellm-proxy-budget row (global proxy spend) + """ + ## if an end-user is passed in, do an upsert - we can't guarantee they already exist in db + existing_user_obj = await user_api_key_cache.async_get_cache(key=user_id) + if existing_user_obj is not None and isinstance(existing_user_obj, dict): + existing_user_obj = LiteLLM_UserTable(**existing_user_obj) + try: + if prisma_client is not None: # update + user_ids = [user_id] + if ( + litellm.max_budget > 0 + ): # track global proxy budget, if user set max budget + user_ids.append(litellm_proxy_budget_name) + + for _id in user_ids: + if _id is not None: + await DBSpendUpdateWriter._update_transaction_list( + response_cost=response_cost, + entity_id=_id, + transaction_list=prisma_client.user_list_transactions, + entity_type=Litellm_EntityType.USER, + prisma_client=prisma_client, + ) + + if end_user_id is not None: + await DBSpendUpdateWriter._update_transaction_list( + response_cost=response_cost, + entity_id=end_user_id, + transaction_list=prisma_client.end_user_list_transactions, + entity_type=Litellm_EntityType.END_USER, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.info( + "\033[91m" + + f"Update User DB call failed to execute {str(e)}\n{traceback.format_exc()}" + ) + + @staticmethod + async def _update_team_db( + response_cost: Optional[float], + team_id: Optional[str], + user_id: Optional[str], + prisma_client: Optional[PrismaClient], + ): + try: + if team_id is None or prisma_client is None: + verbose_proxy_logger.debug( + "track_cost_callback: team_id is None or prisma_client is None. Not tracking spend for team" + ) + return + + await DBSpendUpdateWriter._update_transaction_list( + response_cost=response_cost, + entity_id=team_id, + transaction_list=prisma_client.team_list_transactions, + entity_type=Litellm_EntityType.TEAM, + prisma_client=prisma_client, + ) + + try: + # Track spend of the team member within this team + if user_id is not None: + # key is "team_id::::user_id::" + team_member_key = f"team_id::{team_id}::user_id::{user_id}" + await DBSpendUpdateWriter._update_transaction_list( + response_cost=response_cost, + entity_id=team_member_key, + transaction_list=prisma_client.team_member_list_transactions, + entity_type=Litellm_EntityType.TEAM_MEMBER, + prisma_client=prisma_client, + ) + except Exception: + pass + except Exception as e: + verbose_proxy_logger.info( + f"Update Team DB failed to execute - {str(e)}\n{traceback.format_exc()}" + ) + raise e + + @staticmethod + async def _update_org_db( + response_cost: Optional[float], + org_id: Optional[str], + prisma_client: Optional[PrismaClient], + ): + try: + if org_id is None or prisma_client is None: + verbose_proxy_logger.debug( + "track_cost_callback: org_id is None or prisma_client is None. Not tracking spend for org" + ) + return + + await DBSpendUpdateWriter._update_transaction_list( + response_cost=response_cost, + entity_id=org_id, + transaction_list=prisma_client.org_list_transactions, + entity_type=Litellm_EntityType.ORGANIZATION, + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.info( + f"Update Org DB failed to execute - {str(e)}\n{traceback.format_exc()}" + ) + raise e + + @staticmethod + async def _insert_spend_log_to_db( + kwargs: Optional[dict], + completion_response: Optional[Union[litellm.ModelResponse, Any, Exception]], + start_time: Optional[datetime], + end_time: Optional[datetime], + response_cost: Optional[float], + prisma_client: Optional[PrismaClient], + ): + from litellm.proxy.spend_tracking.spend_tracking_utils import ( + get_logging_payload, + ) + + try: + if prisma_client: + payload = get_logging_payload( + kwargs=kwargs, + response_obj=completion_response, + start_time=start_time, + end_time=end_time, + ) + payload["spend"] = response_cost or 0.0 + DBSpendUpdateWriter._set_spend_logs_payload( + payload=payload, + spend_logs_url=os.getenv("SPEND_LOGS_URL"), + prisma_client=prisma_client, + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Update Spend Logs DB failed to execute - {str(e)}\n{traceback.format_exc()}" + ) + raise e + + @staticmethod + def _set_spend_logs_payload( + payload: Union[dict, SpendLogsPayload], + prisma_client: PrismaClient, + spend_logs_url: Optional[str] = None, + ) -> PrismaClient: + verbose_proxy_logger.info( + "Writing spend log to db - request_id: {}, spend: {}".format( + payload.get("request_id"), payload.get("spend") + ) + ) + if prisma_client is not None and spend_logs_url is not None: + if isinstance(payload["startTime"], datetime): + payload["startTime"] = payload["startTime"].isoformat() + if isinstance(payload["endTime"], datetime): + payload["endTime"] = payload["endTime"].isoformat() + prisma_client.spend_log_transactions.append(payload) + elif prisma_client is not None: + prisma_client.spend_log_transactions.append(payload) + + prisma_client.add_spend_log_transaction_to_daily_user_transaction( + payload.copy() + ) + return prisma_client + + async def db_update_spend_transaction_handler( + self, + prisma_client: PrismaClient, + n_retry_times: int, + proxy_logging_obj: ProxyLogging, + ): + """ + Handles commiting update spend transactions to db + + `UPDATES` can lead to deadlocks, hence we handle them separately + + Args: + prisma_client: PrismaClient object + n_retry_times: int, number of retry times + proxy_logging_obj: ProxyLogging object + + How this works: + - Check `general_settings.use_redis_transaction_buffer` + - If enabled, write in-memory transactions to Redis + - Check if this Pod should read from the DB + else: + - Regular flow of this method + """ + if RedisUpdateBuffer._should_commit_spend_updates_to_redis(): + await self._commit_spend_updates_to_db_with_redis( + prisma_client=prisma_client, + n_retry_times=n_retry_times, + proxy_logging_obj=proxy_logging_obj, + ) + + else: + await self._commit_spend_updates_to_db_without_redis_buffer( + prisma_client=prisma_client, + n_retry_times=n_retry_times, + proxy_logging_obj=proxy_logging_obj, + ) + + async def _commit_spend_updates_to_db_with_redis( + self, + prisma_client: PrismaClient, + n_retry_times: int, + proxy_logging_obj: ProxyLogging, + ): + """ + Handler to commit spend updates to Redis and attempt to acquire lock to commit to db + + This is a v2 scalable approach to first commit spend updates to redis, then commit to db + + This minimizes DB Deadlocks since + - All pods only need to write their spend updates to redis + - Only 1 pod will commit to db at a time (based on if it can acquire the lock over writing to DB) + """ + await self.redis_update_buffer.store_in_memory_spend_updates_in_redis( + prisma_client=prisma_client, + ) + + # Only commit from redis to db if this pod is the leader + if await self.pod_lock_manager.acquire_lock(): + verbose_proxy_logger.debug("acquired lock for spend updates") + + try: + db_spend_update_transactions = ( + await self.redis_update_buffer.get_all_update_transactions_from_redis_buffer() + ) + if db_spend_update_transactions is not None: + await DBSpendUpdateWriter._commit_spend_updates_to_db( + prisma_client=prisma_client, + n_retry_times=n_retry_times, + proxy_logging_obj=proxy_logging_obj, + db_spend_update_transactions=db_spend_update_transactions, + ) + except Exception as e: + verbose_proxy_logger.error(f"Error committing spend updates: {e}") + finally: + await self.pod_lock_manager.release_lock() + + async def _commit_spend_updates_to_db_without_redis_buffer( + self, + prisma_client: PrismaClient, + n_retry_times: int, + proxy_logging_obj: ProxyLogging, + ): + """ + Commits all the spend `UPDATE` transactions to the Database + + This is the regular flow of committing to db without using a redis buffer + + Note: This flow causes Deadlocks in production (1K RPS+). Use self._commit_spend_updates_to_db_with_redis() instead if you expect 1K+ RPS. + """ + db_spend_update_transactions = 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, + ) + await DBSpendUpdateWriter._commit_spend_updates_to_db( + prisma_client=prisma_client, + n_retry_times=n_retry_times, + proxy_logging_obj=proxy_logging_obj, + db_spend_update_transactions=db_spend_update_transactions, + ) + + @staticmethod + async def _commit_spend_updates_to_db( # noqa: PLR0915 + prisma_client: PrismaClient, + n_retry_times: int, + proxy_logging_obj: ProxyLogging, + db_spend_update_transactions: DBSpendUpdateTransactions, + ): + """ + Commits all the spend `UPDATE` transactions to the Database + + """ + from litellm.proxy.utils import ( + ProxyUpdateSpend, + _raise_failed_update_spend_exception, + ) + + ### UPDATE USER TABLE ### + user_list_transactions = db_spend_update_transactions["user_list_transactions"] + verbose_proxy_logger.debug( + "User Spend transactions: {}".format(user_list_transactions) + ) + if ( + user_list_transactions is not None + and len(user_list_transactions.keys()) > 0 + ): + for i in range(n_retry_times + 1): + start_time = time.time() + try: + async with prisma_client.db.tx( + timeout=timedelta(seconds=60) + ) as transaction: + async with transaction.batch_() as batcher: + for ( + user_id, + response_cost, + ) in user_list_transactions.items(): + batcher.litellm_usertable.update_many( + where={"user_id": user_id}, + data={"spend": {"increment": response_cost}}, + ) + prisma_client.user_list_transactions = ( + {} + ) # Clear the remaining transactions after processing all batches in the loop. + break + except DB_CONNECTION_ERROR_TYPES as e: + if ( + i >= n_retry_times + ): # If we've reached the maximum number of retries + _raise_failed_update_spend_exception( + e=e, + start_time=start_time, + proxy_logging_obj=proxy_logging_obj, + ) + # Optionally, sleep for a bit before retrying + await asyncio.sleep(2**i) # Exponential backoff + except Exception as e: + _raise_failed_update_spend_exception( + e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj + ) + + ### UPDATE END-USER TABLE ### + end_user_list_transactions = db_spend_update_transactions[ + "end_user_list_transactions" + ] + verbose_proxy_logger.debug( + "End-User Spend transactions: {}".format(end_user_list_transactions) + ) + if ( + end_user_list_transactions is not None + and len(end_user_list_transactions.keys()) > 0 + ): + await ProxyUpdateSpend.update_end_user_spend( + n_retry_times=n_retry_times, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + ) + ### UPDATE KEY TABLE ### + key_list_transactions = db_spend_update_transactions["key_list_transactions"] + verbose_proxy_logger.debug( + "KEY Spend transactions: {}".format(key_list_transactions) + ) + if key_list_transactions is not None and len(key_list_transactions.keys()) > 0: + for i in range(n_retry_times + 1): + start_time = time.time() + try: + async with prisma_client.db.tx( + timeout=timedelta(seconds=60) + ) as transaction: + async with transaction.batch_() as batcher: + for ( + token, + response_cost, + ) in key_list_transactions.items(): + batcher.litellm_verificationtoken.update_many( # 'update_many' prevents error from being raised if no row exists + where={"token": token}, + data={"spend": {"increment": response_cost}}, + ) + prisma_client.key_list_transactions = ( + {} + ) # Clear the remaining transactions after processing all batches in the loop. + break + except DB_CONNECTION_ERROR_TYPES as e: + if ( + i >= n_retry_times + ): # If we've reached the maximum number of retries + _raise_failed_update_spend_exception( + e=e, + start_time=start_time, + proxy_logging_obj=proxy_logging_obj, + ) + # Optionally, sleep for a bit before retrying + await asyncio.sleep(2**i) # Exponential backoff + except Exception as e: + _raise_failed_update_spend_exception( + e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj + ) + + ### UPDATE TEAM TABLE ### + team_list_transactions = db_spend_update_transactions["team_list_transactions"] + verbose_proxy_logger.debug( + "Team Spend transactions: {}".format(team_list_transactions) + ) + if ( + team_list_transactions is not None + and len(team_list_transactions.keys()) > 0 + ): + for i in range(n_retry_times + 1): + start_time = time.time() + try: + async with prisma_client.db.tx( + timeout=timedelta(seconds=60) + ) as transaction: + async with transaction.batch_() as batcher: + for ( + team_id, + response_cost, + ) in team_list_transactions.items(): + verbose_proxy_logger.debug( + "Updating spend for team id={} by {}".format( + team_id, response_cost + ) + ) + batcher.litellm_teamtable.update_many( # 'update_many' prevents error from being raised if no row exists + where={"team_id": team_id}, + data={"spend": {"increment": response_cost}}, + ) + prisma_client.team_list_transactions = ( + {} + ) # Clear the remaining transactions after processing all batches in the loop. + break + except DB_CONNECTION_ERROR_TYPES as e: + if ( + i >= n_retry_times + ): # If we've reached the maximum number of retries + _raise_failed_update_spend_exception( + e=e, + start_time=start_time, + proxy_logging_obj=proxy_logging_obj, + ) + # Optionally, sleep for a bit before retrying + await asyncio.sleep(2**i) # Exponential backoff + except Exception as e: + _raise_failed_update_spend_exception( + e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj + ) + + ### UPDATE TEAM Membership TABLE with spend ### + team_member_list_transactions = db_spend_update_transactions[ + "team_member_list_transactions" + ] + verbose_proxy_logger.debug( + "Team Membership Spend transactions: {}".format( + team_member_list_transactions + ) + ) + if ( + team_member_list_transactions is not None + and len(team_member_list_transactions.keys()) > 0 + ): + for i in range(n_retry_times + 1): + start_time = time.time() + try: + async with prisma_client.db.tx( + timeout=timedelta(seconds=60) + ) as transaction: + async with transaction.batch_() as batcher: + for ( + key, + response_cost, + ) in team_member_list_transactions.items(): + # key is "team_id::::user_id::" + team_id = key.split("::")[1] + user_id = key.split("::")[3] + + batcher.litellm_teammembership.update_many( # 'update_many' prevents error from being raised if no row exists + where={"team_id": team_id, "user_id": user_id}, + data={"spend": {"increment": response_cost}}, + ) + prisma_client.team_member_list_transactions = ( + {} + ) # Clear the remaining transactions after processing all batches in the loop. + break + except DB_CONNECTION_ERROR_TYPES as e: + if ( + i >= n_retry_times + ): # If we've reached the maximum number of retries + _raise_failed_update_spend_exception( + e=e, + start_time=start_time, + proxy_logging_obj=proxy_logging_obj, + ) + # Optionally, sleep for a bit before retrying + await asyncio.sleep(2**i) # Exponential backoff + except Exception as e: + _raise_failed_update_spend_exception( + e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj + ) + + ### UPDATE ORG TABLE ### + org_list_transactions = db_spend_update_transactions["org_list_transactions"] + verbose_proxy_logger.debug( + "Org Spend transactions: {}".format(org_list_transactions) + ) + if org_list_transactions is not None and len(org_list_transactions.keys()) > 0: + for i in range(n_retry_times + 1): + start_time = time.time() + try: + async with prisma_client.db.tx( + timeout=timedelta(seconds=60) + ) as transaction: + async with transaction.batch_() as batcher: + for ( + org_id, + response_cost, + ) in org_list_transactions.items(): + batcher.litellm_organizationtable.update_many( # 'update_many' prevents error from being raised if no row exists + where={"organization_id": org_id}, + data={"spend": {"increment": response_cost}}, + ) + prisma_client.org_list_transactions = ( + {} + ) # Clear the remaining transactions after processing all batches in the loop. + break + except DB_CONNECTION_ERROR_TYPES as e: + if ( + i >= n_retry_times + ): # If we've reached the maximum number of retries + _raise_failed_update_spend_exception( + e=e, + start_time=start_time, + proxy_logging_obj=proxy_logging_obj, + ) + # Optionally, sleep for a bit before retrying + await asyncio.sleep(2**i) # Exponential backoff + except Exception as e: + _raise_failed_update_spend_exception( + e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj + ) diff --git a/litellm/proxy/db/pod_lock_manager.py b/litellm/proxy/db/pod_lock_manager.py new file mode 100644 index 0000000000..84c92c9daa --- /dev/null +++ b/litellm/proxy/db/pod_lock_manager.py @@ -0,0 +1,137 @@ +import uuid +from datetime import datetime, timedelta, timezone +from typing import TYPE_CHECKING, Any + +from litellm._logging import verbose_proxy_logger +from litellm.constants import DEFAULT_CRON_JOB_LOCK_TTL_SECONDS + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient, ProxyLogging +else: + PrismaClient = Any + ProxyLogging = Any + + +class PodLockManager: + """ + Manager for acquiring and releasing locks for cron jobs. + + Ensures that only one pod can run a cron job at a time. + """ + + def __init__(self, cronjob_id: str): + self.pod_id = str(uuid.uuid4()) + self.cronjob_id = cronjob_id + + async def acquire_lock(self) -> bool: + """ + Attempt to acquire the lock for a specific cron job using database locking. + """ + from litellm.proxy.proxy_server import prisma_client + + verbose_proxy_logger.debug( + "Pod %s acquiring lock for cronjob_id=%s", self.pod_id, self.cronjob_id + ) + if not prisma_client: + verbose_proxy_logger.debug("prisma is None, returning False") + return False + + try: + current_time = datetime.now(timezone.utc) + ttl_expiry = current_time + timedelta( + seconds=DEFAULT_CRON_JOB_LOCK_TTL_SECONDS + ) + + # Use Prisma's findUnique with FOR UPDATE lock to prevent race conditions + lock_record = await prisma_client.db.litellm_cronjob.find_unique( + where={"cronjob_id": self.cronjob_id}, + ) + + if lock_record: + # If record exists, only update if it's inactive or expired + if lock_record.status == "ACTIVE" and lock_record.ttl > current_time: + return lock_record.pod_id == self.pod_id + + # Update existing record + updated_lock = await prisma_client.db.litellm_cronjob.update( + where={"cronjob_id": self.cronjob_id}, + data={ + "pod_id": self.pod_id, + "status": "ACTIVE", + "last_updated": current_time, + "ttl": ttl_expiry, + }, + ) + else: + # Create new record if none exists + updated_lock = await prisma_client.db.litellm_cronjob.create( + data={ + "cronjob_id": self.cronjob_id, + "pod_id": self.pod_id, + "status": "ACTIVE", + "last_updated": current_time, + "ttl": ttl_expiry, + } + ) + + return updated_lock.pod_id == self.pod_id + + except Exception as e: + verbose_proxy_logger.error( + f"Error acquiring the lock for {self.cronjob_id}: {e}" + ) + return False + + async def renew_lock(self): + """ + Renew the lock (update the TTL) for the pod holding the lock. + """ + from litellm.proxy.proxy_server import prisma_client + + if not prisma_client: + return False + try: + verbose_proxy_logger.debug( + "renewing lock for cronjob_id=%s", self.cronjob_id + ) + current_time = datetime.now(timezone.utc) + # Extend the TTL for another DEFAULT_CRON_JOB_LOCK_TTL_SECONDS + ttl_expiry = current_time + timedelta( + seconds=DEFAULT_CRON_JOB_LOCK_TTL_SECONDS + ) + + await prisma_client.db.litellm_cronjob.update( + where={"cronjob_id": self.cronjob_id, "pod_id": self.pod_id}, + data={"ttl": ttl_expiry, "last_updated": current_time}, + ) + verbose_proxy_logger.info( + f"Renewed the lock for Pod {self.pod_id} for {self.cronjob_id}" + ) + except Exception as e: + verbose_proxy_logger.error( + f"Error renewing the lock for {self.cronjob_id}: {e}" + ) + + async def release_lock(self): + """ + Release the lock and mark the pod as inactive. + """ + from litellm.proxy.proxy_server import prisma_client + + if not prisma_client: + return False + try: + verbose_proxy_logger.debug( + "Pod %s releasing lock for cronjob_id=%s", self.pod_id, self.cronjob_id + ) + await prisma_client.db.litellm_cronjob.update( + where={"cronjob_id": self.cronjob_id, "pod_id": self.pod_id}, + data={"status": "INACTIVE"}, + ) + verbose_proxy_logger.info( + f"Pod {self.pod_id} has released the lock for {self.cronjob_id}." + ) + except Exception as e: + verbose_proxy_logger.error( + f"Error releasing the lock for {self.cronjob_id}: {e}" + ) diff --git a/litellm/proxy/db/redis_update_buffer.py b/litellm/proxy/db/redis_update_buffer.py new file mode 100644 index 0000000000..f77c839aaf --- /dev/null +++ b/litellm/proxy/db/redis_update_buffer.py @@ -0,0 +1,251 @@ +""" +Handles buffering database `UPDATE` transactions in Redis before committing them to the database + +This is to prevent deadlocks and improve reliability +""" + +import json +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union + +from litellm._logging import verbose_proxy_logger +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.secret_managers.main import str_to_bool + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient +else: + PrismaClient = Any + + +class RedisUpdateBuffer: + """ + Handles buffering database `UPDATE` transactions in Redis before committing them to the database + + This is to prevent deadlocks and improve reliability + """ + + def __init__( + self, + redis_cache: Optional[RedisCache] = None, + ): + self.redis_cache = redis_cache + + @staticmethod + def _should_commit_spend_updates_to_redis() -> bool: + """ + Checks if the Pod should commit spend updates to Redis + + This setting enables buffering database transactions in Redis + to improve reliability and reduce database contention + """ + from litellm.proxy.proxy_server import general_settings + + _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: + return False + return _use_redis_transaction_buffer + + async def store_in_memory_spend_updates_in_redis( + self, + prisma_client: PrismaClient, + ): + """ + Stores the in-memory spend updates to Redis + + Each transaction is a dict stored as following: + - key is the entity id + - value is the spend amount + + ``` + Redis List: + key_list_transactions: + [ + "0929880201": 1.2, + "0929880202": 0.01, + "0929880203": 0.001, + ] + ``` + """ + if self.redis_cache is None: + verbose_proxy_logger.debug( + "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, + ) + ) + + # 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, + ) + + # 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, + ) -> int: + """ + Gets the number of transactions to store in Redis + """ + num_transactions = 0 + for v in db_spend_update_transactions.values(): + if isinstance(v, dict): + 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]: + """ + Removes the specified prefix from the keys of a dictionary. + """ + return {key.replace(prefix, "", 1): value for key, value in data.items()} + + async def get_all_update_transactions_from_redis_buffer( + self, + ) -> Optional[DBSpendUpdateTransactions]: + """ + Gets all the update transactions from Redis + + On Redis we store a list of transactions as a JSON string + + eg. + [ + DBSpendUpdateTransactions( + user_list_transactions={ + "user_id_1": 1.2, + "user_id_2": 0.01, + }, + end_user_list_transactions={}, + key_list_transactions={ + "0929880201": 1.2, + "0929880202": 0.01, + }, + team_list_transactions={}, + team_member_list_transactions={}, + org_list_transactions={}, + ), + DBSpendUpdateTransactions( + user_list_transactions={ + "user_id_3": 1.2, + "user_id_4": 0.01, + }, + end_user_list_transactions={}, + key_list_transactions={ + "key_id_1": 1.2, + "key_id_2": 0.01, + }, + team_list_transactions={}, + team_member_list_transactions={}, + org_list_transactions={}, + ] + """ + if self.redis_cache is None: + return None + list_of_transactions = await self.redis_cache.async_lpop( + key=REDIS_UPDATE_BUFFER_KEY, + count=MAX_REDIS_BUFFER_DEQUEUE_COUNT, + ) + if list_of_transactions is None: + return None + + # Parse the list of transactions from JSON strings + parsed_transactions = self._parse_list_of_transactions(list_of_transactions) + + # If there are no transactions, return None + if len(parsed_transactions) == 0: + return None + + # Combine all transactions into a single transaction + combined_transaction = self._combine_list_of_transactions(parsed_transactions) + + return combined_transaction + + @staticmethod + def _parse_list_of_transactions( + list_of_transactions: Union[Any, List[Any]], + ) -> List[DBSpendUpdateTransactions]: + """ + Parses the list of transactions from Redis + """ + 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( + list_of_transactions: List[DBSpendUpdateTransactions], + ) -> DBSpendUpdateTransactions: + """ + Combines the list of transactions into a single DBSpendUpdateTransactions object + """ + # Initialize a new combined transaction object with empty dictionaries + combined_transaction = DBSpendUpdateTransactions( + user_list_transactions={}, + end_user_list_transactions={}, + key_list_transactions={}, + team_list_transactions={}, + team_member_list_transactions={}, + org_list_transactions={}, + ) + + # Define the transaction fields to process + transaction_fields = [ + "user_list_transactions", + "end_user_list_transactions", + "key_list_transactions", + "team_list_transactions", + "team_member_list_transactions", + "org_list_transactions", + ] + + # Loop through each transaction and combine the values + for transaction in list_of_transactions: + # Process each field type + for field in transaction_fields: + if transaction.get(field): + for entity_id, amount in transaction[field].items(): # type: ignore + combined_transaction[field][entity_id] = ( # type: ignore + combined_transaction[field].get(entity_id, 0) + amount # type: ignore + ) + + return combined_transaction diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index e8a947329d..f205b0146f 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -13,6 +13,7 @@ from litellm.litellm_core_utils.core_helpers import ( from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import log_db_metrics +from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import ( StandardLoggingPayload, @@ -33,8 +34,6 @@ class _ProxyDBLogger(CustomLogger): original_exception: Exception, user_api_key_dict: UserAPIKeyAuth, ): - from litellm.proxy.proxy_server import update_database - if _ProxyDBLogger._should_track_errors_in_db() is False: return @@ -67,7 +66,7 @@ class _ProxyDBLogger(CustomLogger): request_data.get("proxy_server_request") or {} ) request_data["litellm_params"]["metadata"] = existing_metadata - await update_database( + await DBSpendUpdateWriter.update_database( token=user_api_key_dict.api_key, response_cost=0.0, user_id=user_api_key_dict.user_id, @@ -94,7 +93,6 @@ class _ProxyDBLogger(CustomLogger): prisma_client, proxy_logging_obj, update_cache, - update_database, ) verbose_proxy_logger.debug("INSIDE _PROXY_track_cost_callback") @@ -138,7 +136,7 @@ class _ProxyDBLogger(CustomLogger): end_user_id=end_user_id, ): ## UPDATE DATABASE - await update_database( + await DBSpendUpdateWriter.update_database( token=user_api_key, response_cost=response_cost, user_id=user_id, diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 4912a35f89..106003f996 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -5,5 +5,3 @@ model_list: api_key: fake-key api_base: https://exampleopenaiendpoint-production.up.railway.app/ -general_settings: - allow_requests_on_db_unavailable: True \ No newline at end of file diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7c9b60d5cb..89c15f413d 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -24,12 +24,12 @@ from typing import ( get_type_hints, ) +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.types.utils import ( ModelResponse, ModelResponseStream, TextCompletionResponse, ) -from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -897,211 +897,6 @@ def cost_tracking(): litellm.logging_callback_manager.add_litellm_callback(_ProxyDBLogger()) -def _set_spend_logs_payload( - payload: Union[dict, SpendLogsPayload], - prisma_client: PrismaClient, - spend_logs_url: Optional[str] = None, -): - verbose_proxy_logger.info( - "Writing spend log to db - request_id: {}, spend: {}".format( - payload.get("request_id"), payload.get("spend") - ) - ) - if prisma_client is not None and spend_logs_url is not None: - if isinstance(payload["startTime"], datetime): - payload["startTime"] = payload["startTime"].isoformat() - if isinstance(payload["endTime"], datetime): - payload["endTime"] = payload["endTime"].isoformat() - prisma_client.spend_log_transactions.append(payload) - elif prisma_client is not None: - prisma_client.spend_log_transactions.append(payload) - - prisma_client.add_spend_log_transaction_to_daily_user_transaction(payload.copy()) - return prisma_client - - -async def update_database( # noqa: PLR0915 - token, - response_cost, - user_id=None, - end_user_id=None, - team_id=None, - kwargs=None, - completion_response=None, - start_time=None, - end_time=None, - org_id=None, -): - try: - global prisma_client - verbose_proxy_logger.debug( - f"Enters prisma db call, response_cost: {response_cost}, token: {token}; user_id: {user_id}; team_id: {team_id}" - ) - if ProxyUpdateSpend.disable_spend_updates() is True: - return - if token is not None and isinstance(token, str) and token.startswith("sk-"): - hashed_token = hash_token(token=token) - else: - hashed_token = token - - ### UPDATE USER SPEND ### - async def _update_user_db(): - """ - - Update that user's row - - Update litellm-proxy-budget row (global proxy spend) - """ - ## if an end-user is passed in, do an upsert - we can't guarantee they already exist in db - existing_user_obj = await user_api_key_cache.async_get_cache(key=user_id) - if existing_user_obj is not None and isinstance(existing_user_obj, dict): - existing_user_obj = LiteLLM_UserTable(**existing_user_obj) - try: - if prisma_client is not None: # update - user_ids = [user_id] - if ( - litellm.max_budget > 0 - ): # track global proxy budget, if user set max budget - user_ids.append(litellm_proxy_budget_name) - ### KEY CHANGE ### - for _id in user_ids: - if _id is not None: - prisma_client.user_list_transactons[_id] = ( - response_cost - + prisma_client.user_list_transactons.get(_id, 0) - ) - if end_user_id is not None: - prisma_client.end_user_list_transactons[end_user_id] = ( - response_cost - + prisma_client.end_user_list_transactons.get( - end_user_id, 0 - ) - ) - except Exception as e: - verbose_proxy_logger.info( - "\033[91m" - + f"Update User DB call failed to execute {str(e)}\n{traceback.format_exc()}" - ) - - ### UPDATE KEY SPEND ### - async def _update_key_db(): - try: - verbose_proxy_logger.debug( - f"adding spend to key db. Response cost: {response_cost}. Token: {hashed_token}." - ) - if hashed_token is None: - return - if prisma_client is not None: - prisma_client.key_list_transactons[hashed_token] = ( - response_cost - + prisma_client.key_list_transactons.get(hashed_token, 0) - ) - except Exception as e: - verbose_proxy_logger.exception( - f"Update Key DB Call failed to execute - {str(e)}" - ) - raise e - - ### UPDATE SPEND LOGS ### - async def _insert_spend_log_to_db(): - try: - global prisma_client - if prisma_client is not None: - # Helper to generate payload to log - payload = get_logging_payload( - kwargs=kwargs, - response_obj=completion_response, - start_time=start_time, - end_time=end_time, - ) - payload["spend"] = response_cost - prisma_client = _set_spend_logs_payload( - payload=payload, - spend_logs_url=os.getenv("SPEND_LOGS_URL"), - prisma_client=prisma_client, - ) - except Exception as e: - verbose_proxy_logger.debug( - f"Update Spend Logs DB failed to execute - {str(e)}\n{traceback.format_exc()}" - ) - raise e - - ### UPDATE TEAM SPEND ### - async def _update_team_db(): - try: - verbose_proxy_logger.debug( - f"adding spend to team db. Response cost: {response_cost}. team_id: {team_id}." - ) - if team_id is None: - verbose_proxy_logger.debug( - "track_cost_callback: team_id is None. Not tracking spend for team" - ) - return - if prisma_client is not None: - prisma_client.team_list_transactons[team_id] = ( - response_cost - + prisma_client.team_list_transactons.get(team_id, 0) - ) - - try: - # Track spend of the team member within this team - # key is "team_id::::user_id::" - team_member_key = f"team_id::{team_id}::user_id::{user_id}" - prisma_client.team_member_list_transactons[team_member_key] = ( - response_cost - + prisma_client.team_member_list_transactons.get( - team_member_key, 0 - ) - ) - except Exception: - pass - except Exception as e: - verbose_proxy_logger.info( - f"Update Team DB failed to execute - {str(e)}\n{traceback.format_exc()}" - ) - raise e - - ### UPDATE ORG SPEND ### - async def _update_org_db(): - try: - verbose_proxy_logger.debug( - "adding spend to org db. Response cost: {}. org_id: {}.".format( - response_cost, org_id - ) - ) - if org_id is None: - verbose_proxy_logger.debug( - "track_cost_callback: org_id is None. Not tracking spend for org" - ) - return - if prisma_client is not None: - prisma_client.org_list_transactons[org_id] = ( - response_cost - + prisma_client.org_list_transactons.get(org_id, 0) - ) - except Exception as e: - verbose_proxy_logger.info( - f"Update Org DB failed to execute - {str(e)}\n{traceback.format_exc()}" - ) - raise e - - asyncio.create_task(_update_user_db()) - asyncio.create_task(_update_key_db()) - asyncio.create_task(_update_team_db()) - asyncio.create_task(_update_org_db()) - # asyncio.create_task(_insert_spend_log_to_db()) - if disable_spend_logs is False: - await _insert_spend_log_to_db() - else: - verbose_proxy_logger.info( - "disable_spend_logs=True. Skipping writing spend logs to db. Other spend updates - Key/User/Team table will still occur." - ) - - verbose_proxy_logger.debug("Runs spend update on all tables") - except Exception: - verbose_proxy_logger.debug( - f"Error updating Prisma database: {traceback.format_exc()}" - ) - - async def update_cache( # noqa: PLR0915 token: Optional[str], user_id: Optional[str], @@ -1737,7 +1532,7 @@ class ProxyConfig: cache_params = {} if "cache_params" in litellm_settings: cache_params_in_config = litellm_settings["cache_params"] - # overwrie cache_params with cache_params_in_config + # overwrite cache_params with cache_params_in_config cache_params.update(cache_params_in_config) cache_type = cache_params.get("type", "redis") diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 7acb5e2615..d98f1fa981 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -336,3 +336,20 @@ model LiteLLM_DailyUserSpend { @@index([api_key]) @@index([model]) } + + +// Track the status of cron jobs running. Only allow one pod to run the job at a time +model LiteLLM_CronJob { + cronjob_id String @id @default(cuid()) // Unique ID for the record + pod_id String // Unique identifier for the pod acting as the leader + status JobStatus @default(INACTIVE) // Status of the cron job (active or inactive) + last_updated DateTime @default(now()) // Timestamp for the last update of the cron job record + ttl DateTime // Time when the leader's lease expires +} + +enum JobStatus { + ACTIVE + INACTIVE +} + + diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 1c4e625781..4c82f586fb 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -62,6 +62,7 @@ from litellm.proxy.db.create_views import ( create_missing_views, should_create_missing_views, ) +from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.log_db_metrics import log_db_metrics from litellm.proxy.db.prisma_client import PrismaWrapper from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck @@ -264,6 +265,7 @@ class ProxyLogging: ) self.premium_user = premium_user self.service_logging_obj = ServiceLogging() + self.db_spend_update_writer = DBSpendUpdateWriter() def startup_event( self, @@ -336,6 +338,7 @@ class ProxyLogging: if redis_cache is not None: self.internal_usage_cache.dual_cache.redis_cache = redis_cache + self.db_spend_update_writer.redis_update_buffer.redis_cache = redis_cache def _init_litellm_callbacks(self, llm_router: Optional[Router] = None): litellm.logging_callback_manager.add_litellm_callback(self.max_parallel_request_limiter) # type: ignore @@ -1098,12 +1101,12 @@ def jsonify_object(data: dict) -> dict: class PrismaClient: - user_list_transactons: dict = {} - end_user_list_transactons: dict = {} - key_list_transactons: dict = {} - team_list_transactons: dict = {} - team_member_list_transactons: dict = {} # key is ["team_id" + "user_id"] - org_list_transactons: dict = {} + user_list_transactions: dict = {} + end_user_list_transactions: dict = {} + key_list_transactions: dict = {} + team_list_transactions: dict = {} + team_member_list_transactions: dict = {} # key is ["team_id" + "user_id"] + org_list_transactions: dict = {} spend_log_transactions: List = [] daily_user_spend_transactions: Dict[str, DailyUserSpendTransaction] = {} @@ -2432,7 +2435,7 @@ class ProxyUpdateSpend: for ( end_user_id, response_cost, - ) in prisma_client.end_user_list_transactons.items(): + ) in prisma_client.end_user_list_transactions.items(): if litellm.max_end_user_budget is not None: pass batcher.litellm_endusertable.upsert( @@ -2460,7 +2463,7 @@ class ProxyUpdateSpend: e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj ) finally: - prisma_client.end_user_list_transactons = ( + prisma_client.end_user_list_transactions = ( {} ) # reset the end user list transactions - prevent bad data from causing issues @@ -2680,202 +2683,11 @@ async def update_spend( # noqa: PLR0915 spend_logs: list, """ n_retry_times = 3 - i = None - ### UPDATE USER TABLE ### - if len(prisma_client.user_list_transactons.keys()) > 0: - for i in range(n_retry_times + 1): - start_time = time.time() - try: - async with prisma_client.db.tx( - timeout=timedelta(seconds=60) - ) as transaction: - async with transaction.batch_() as batcher: - for ( - user_id, - response_cost, - ) in prisma_client.user_list_transactons.items(): - batcher.litellm_usertable.update_many( - where={"user_id": user_id}, - data={"spend": {"increment": response_cost}}, - ) - prisma_client.user_list_transactons = ( - {} - ) # Clear the remaining transactions after processing all batches in the loop. - break - except DB_CONNECTION_ERROR_TYPES as e: - if i >= n_retry_times: # If we've reached the maximum number of retries - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff - except Exception as e: - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - - ### UPDATE END-USER TABLE ### - verbose_proxy_logger.debug( - "End-User Spend transactions: {}".format( - len(prisma_client.end_user_list_transactons.keys()) - ) + await proxy_logging_obj.db_spend_update_writer.db_update_spend_transaction_handler( + prisma_client=prisma_client, + n_retry_times=n_retry_times, + proxy_logging_obj=proxy_logging_obj, ) - if len(prisma_client.end_user_list_transactons.keys()) > 0: - await ProxyUpdateSpend.update_end_user_spend( - n_retry_times=n_retry_times, - prisma_client=prisma_client, - proxy_logging_obj=proxy_logging_obj, - ) - ### UPDATE KEY TABLE ### - verbose_proxy_logger.debug( - "KEY Spend transactions: {}".format( - len(prisma_client.key_list_transactons.keys()) - ) - ) - if len(prisma_client.key_list_transactons.keys()) > 0: - for i in range(n_retry_times + 1): - start_time = time.time() - try: - async with prisma_client.db.tx( - timeout=timedelta(seconds=60) - ) as transaction: - async with transaction.batch_() as batcher: - for ( - token, - response_cost, - ) in prisma_client.key_list_transactons.items(): - batcher.litellm_verificationtoken.update_many( # 'update_many' prevents error from being raised if no row exists - where={"token": token}, - data={"spend": {"increment": response_cost}}, - ) - prisma_client.key_list_transactons = ( - {} - ) # Clear the remaining transactions after processing all batches in the loop. - break - except DB_CONNECTION_ERROR_TYPES as e: - if i >= n_retry_times: # If we've reached the maximum number of retries - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff - except Exception as e: - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - - ### UPDATE TEAM TABLE ### - verbose_proxy_logger.debug( - "Team Spend transactions: {}".format( - len(prisma_client.team_list_transactons.keys()) - ) - ) - if len(prisma_client.team_list_transactons.keys()) > 0: - for i in range(n_retry_times + 1): - start_time = time.time() - try: - async with prisma_client.db.tx( - timeout=timedelta(seconds=60) - ) as transaction: - async with transaction.batch_() as batcher: - for ( - team_id, - response_cost, - ) in prisma_client.team_list_transactons.items(): - verbose_proxy_logger.debug( - "Updating spend for team id={} by {}".format( - team_id, response_cost - ) - ) - batcher.litellm_teamtable.update_many( # 'update_many' prevents error from being raised if no row exists - where={"team_id": team_id}, - data={"spend": {"increment": response_cost}}, - ) - prisma_client.team_list_transactons = ( - {} - ) # Clear the remaining transactions after processing all batches in the loop. - break - except DB_CONNECTION_ERROR_TYPES as e: - if i >= n_retry_times: # If we've reached the maximum number of retries - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff - except Exception as e: - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - - ### UPDATE TEAM Membership TABLE with spend ### - if len(prisma_client.team_member_list_transactons.keys()) > 0: - for i in range(n_retry_times + 1): - start_time = time.time() - try: - async with prisma_client.db.tx( - timeout=timedelta(seconds=60) - ) as transaction: - async with transaction.batch_() as batcher: - for ( - key, - response_cost, - ) in prisma_client.team_member_list_transactons.items(): - # key is "team_id::::user_id::" - team_id = key.split("::")[1] - user_id = key.split("::")[3] - - batcher.litellm_teammembership.update_many( # 'update_many' prevents error from being raised if no row exists - where={"team_id": team_id, "user_id": user_id}, - data={"spend": {"increment": response_cost}}, - ) - prisma_client.team_member_list_transactons = ( - {} - ) # Clear the remaining transactions after processing all batches in the loop. - break - except DB_CONNECTION_ERROR_TYPES as e: - if i >= n_retry_times: # If we've reached the maximum number of retries - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff - except Exception as e: - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - - ### UPDATE ORG TABLE ### - if len(prisma_client.org_list_transactons.keys()) > 0: - for i in range(n_retry_times + 1): - start_time = time.time() - try: - async with prisma_client.db.tx( - timeout=timedelta(seconds=60) - ) as transaction: - async with transaction.batch_() as batcher: - for ( - org_id, - response_cost, - ) in prisma_client.org_list_transactons.items(): - batcher.litellm_organizationtable.update_many( # 'update_many' prevents error from being raised if no row exists - where={"organization_id": org_id}, - data={"spend": {"increment": response_cost}}, - ) - prisma_client.org_list_transactons = ( - {} - ) # Clear the remaining transactions after processing all batches in the loop. - break - except DB_CONNECTION_ERROR_TYPES as e: - if i >= n_retry_times: # If we've reached the maximum number of retries - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) - # Optionally, sleep for a bit before retrying - await asyncio.sleep(2**i) # Exponential backoff - except Exception as e: - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) ### UPDATE SPEND LOGS ### verbose_proxy_logger.debug( diff --git a/migrations/20250329044652_add_cron_job_table/migration.sql b/migrations/20250329044652_add_cron_job_table/migration.sql new file mode 100644 index 0000000000..edfa99abd5 --- /dev/null +++ b/migrations/20250329044652_add_cron_job_table/migration.sql @@ -0,0 +1,407 @@ +-- CreateEnum +CREATE TYPE "JobStatus" AS ENUM ('ACTIVE', 'INACTIVE'); + +-- CreateTable +CREATE TABLE "LiteLLM_BudgetTable" ( + "budget_id" TEXT NOT NULL, + "max_budget" DOUBLE PRECISION, + "soft_budget" DOUBLE PRECISION, + "max_parallel_requests" INTEGER, + "tpm_limit" BIGINT, + "rpm_limit" BIGINT, + "model_max_budget" JSONB, + "budget_duration" TEXT, + "budget_reset_at" TIMESTAMP(3), + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT NOT NULL, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_by" TEXT NOT NULL, + + CONSTRAINT "LiteLLM_BudgetTable_pkey" PRIMARY KEY ("budget_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_CredentialsTable" ( + "credential_id" TEXT NOT NULL, + "credential_name" TEXT NOT NULL, + "credential_values" JSONB NOT NULL, + "credential_info" JSONB, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT NOT NULL, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_by" TEXT NOT NULL, + + CONSTRAINT "LiteLLM_CredentialsTable_pkey" PRIMARY KEY ("credential_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_ProxyModelTable" ( + "model_id" TEXT NOT NULL, + "model_name" TEXT NOT NULL, + "litellm_params" JSONB NOT NULL, + "model_info" JSONB, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT NOT NULL, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_by" TEXT NOT NULL, + + CONSTRAINT "LiteLLM_ProxyModelTable_pkey" PRIMARY KEY ("model_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_OrganizationTable" ( + "organization_id" TEXT NOT NULL, + "organization_alias" TEXT NOT NULL, + "budget_id" TEXT NOT NULL, + "metadata" JSONB NOT NULL DEFAULT '{}', + "models" TEXT[], + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "model_spend" JSONB NOT NULL DEFAULT '{}', + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT NOT NULL, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_by" TEXT NOT NULL, + + CONSTRAINT "LiteLLM_OrganizationTable_pkey" PRIMARY KEY ("organization_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_ModelTable" ( + "id" SERIAL NOT NULL, + "aliases" JSONB, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT NOT NULL, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_by" TEXT NOT NULL, + + CONSTRAINT "LiteLLM_ModelTable_pkey" PRIMARY KEY ("id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_TeamTable" ( + "team_id" TEXT NOT NULL, + "team_alias" TEXT, + "organization_id" TEXT, + "admins" TEXT[], + "members" TEXT[], + "members_with_roles" JSONB NOT NULL DEFAULT '{}', + "metadata" JSONB NOT NULL DEFAULT '{}', + "max_budget" DOUBLE PRECISION, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "models" TEXT[], + "max_parallel_requests" INTEGER, + "tpm_limit" BIGINT, + "rpm_limit" BIGINT, + "budget_duration" TEXT, + "budget_reset_at" TIMESTAMP(3), + "blocked" BOOLEAN NOT NULL DEFAULT false, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "model_spend" JSONB NOT NULL DEFAULT '{}', + "model_max_budget" JSONB NOT NULL DEFAULT '{}', + "model_id" INTEGER, + + CONSTRAINT "LiteLLM_TeamTable_pkey" PRIMARY KEY ("team_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_UserTable" ( + "user_id" TEXT NOT NULL, + "user_alias" TEXT, + "team_id" TEXT, + "sso_user_id" TEXT, + "organization_id" TEXT, + "password" TEXT, + "teams" TEXT[] DEFAULT ARRAY[]::TEXT[], + "user_role" TEXT, + "max_budget" DOUBLE PRECISION, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "user_email" TEXT, + "models" TEXT[], + "metadata" JSONB NOT NULL DEFAULT '{}', + "max_parallel_requests" INTEGER, + "tpm_limit" BIGINT, + "rpm_limit" BIGINT, + "budget_duration" TEXT, + "budget_reset_at" TIMESTAMP(3), + "allowed_cache_controls" TEXT[] DEFAULT ARRAY[]::TEXT[], + "model_spend" JSONB NOT NULL DEFAULT '{}', + "model_max_budget" JSONB NOT NULL DEFAULT '{}', + "created_at" TIMESTAMP(3) DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) DEFAULT CURRENT_TIMESTAMP, + + CONSTRAINT "LiteLLM_UserTable_pkey" PRIMARY KEY ("user_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_VerificationToken" ( + "token" TEXT NOT NULL, + "key_name" TEXT, + "key_alias" TEXT, + "soft_budget_cooldown" BOOLEAN NOT NULL DEFAULT false, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "expires" TIMESTAMP(3), + "models" TEXT[], + "aliases" JSONB NOT NULL DEFAULT '{}', + "config" JSONB NOT NULL DEFAULT '{}', + "user_id" TEXT, + "team_id" TEXT, + "permissions" JSONB NOT NULL DEFAULT '{}', + "max_parallel_requests" INTEGER, + "metadata" JSONB NOT NULL DEFAULT '{}', + "blocked" BOOLEAN, + "tpm_limit" BIGINT, + "rpm_limit" BIGINT, + "max_budget" DOUBLE PRECISION, + "budget_duration" TEXT, + "budget_reset_at" TIMESTAMP(3), + "allowed_cache_controls" TEXT[] DEFAULT ARRAY[]::TEXT[], + "model_spend" JSONB NOT NULL DEFAULT '{}', + "model_max_budget" JSONB NOT NULL DEFAULT '{}', + "budget_id" TEXT, + "organization_id" TEXT, + "created_at" TIMESTAMP(3) DEFAULT CURRENT_TIMESTAMP, + "created_by" TEXT, + "updated_at" TIMESTAMP(3) DEFAULT CURRENT_TIMESTAMP, + "updated_by" TEXT, + + CONSTRAINT "LiteLLM_VerificationToken_pkey" PRIMARY KEY ("token") +); + +-- CreateTable +CREATE TABLE "LiteLLM_EndUserTable" ( + "user_id" TEXT NOT NULL, + "alias" TEXT, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "allowed_model_region" TEXT, + "default_model" TEXT, + "budget_id" TEXT, + "blocked" BOOLEAN NOT NULL DEFAULT false, + + CONSTRAINT "LiteLLM_EndUserTable_pkey" PRIMARY KEY ("user_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_Config" ( + "param_name" TEXT NOT NULL, + "param_value" JSONB, + + CONSTRAINT "LiteLLM_Config_pkey" PRIMARY KEY ("param_name") +); + +-- CreateTable +CREATE TABLE "LiteLLM_SpendLogs" ( + "request_id" TEXT NOT NULL, + "call_type" TEXT NOT NULL, + "api_key" TEXT NOT NULL DEFAULT '', + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "total_tokens" INTEGER NOT NULL DEFAULT 0, + "prompt_tokens" INTEGER NOT NULL DEFAULT 0, + "completion_tokens" INTEGER NOT NULL DEFAULT 0, + "startTime" TIMESTAMP(3) NOT NULL, + "endTime" TIMESTAMP(3) NOT NULL, + "completionStartTime" TIMESTAMP(3), + "model" TEXT NOT NULL DEFAULT '', + "model_id" TEXT DEFAULT '', + "model_group" TEXT DEFAULT '', + "custom_llm_provider" TEXT DEFAULT '', + "api_base" TEXT DEFAULT '', + "user" TEXT DEFAULT '', + "metadata" JSONB DEFAULT '{}', + "cache_hit" TEXT DEFAULT '', + "cache_key" TEXT DEFAULT '', + "request_tags" JSONB DEFAULT '[]', + "team_id" TEXT, + "end_user" TEXT, + "requester_ip_address" TEXT, + "messages" JSONB DEFAULT '{}', + "response" JSONB DEFAULT '{}', + + CONSTRAINT "LiteLLM_SpendLogs_pkey" PRIMARY KEY ("request_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_ErrorLogs" ( + "request_id" TEXT NOT NULL, + "startTime" TIMESTAMP(3) NOT NULL, + "endTime" TIMESTAMP(3) NOT NULL, + "api_base" TEXT NOT NULL DEFAULT '', + "model_group" TEXT NOT NULL DEFAULT '', + "litellm_model_name" TEXT NOT NULL DEFAULT '', + "model_id" TEXT NOT NULL DEFAULT '', + "request_kwargs" JSONB NOT NULL DEFAULT '{}', + "exception_type" TEXT NOT NULL DEFAULT '', + "exception_string" TEXT NOT NULL DEFAULT '', + "status_code" TEXT NOT NULL DEFAULT '', + + CONSTRAINT "LiteLLM_ErrorLogs_pkey" PRIMARY KEY ("request_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_UserNotifications" ( + "request_id" TEXT NOT NULL, + "user_id" TEXT NOT NULL, + "models" TEXT[], + "justification" TEXT NOT NULL, + "status" TEXT NOT NULL, + + CONSTRAINT "LiteLLM_UserNotifications_pkey" PRIMARY KEY ("request_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_TeamMembership" ( + "user_id" TEXT NOT NULL, + "team_id" TEXT NOT NULL, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "budget_id" TEXT, + + CONSTRAINT "LiteLLM_TeamMembership_pkey" PRIMARY KEY ("user_id","team_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_OrganizationMembership" ( + "user_id" TEXT NOT NULL, + "organization_id" TEXT NOT NULL, + "user_role" TEXT, + "spend" DOUBLE PRECISION DEFAULT 0.0, + "budget_id" TEXT, + "created_at" TIMESTAMP(3) DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) DEFAULT CURRENT_TIMESTAMP, + + CONSTRAINT "LiteLLM_OrganizationMembership_pkey" PRIMARY KEY ("user_id","organization_id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_InvitationLink" ( + "id" TEXT NOT NULL, + "user_id" TEXT NOT NULL, + "is_accepted" BOOLEAN NOT NULL DEFAULT false, + "accepted_at" TIMESTAMP(3), + "expires_at" TIMESTAMP(3) NOT NULL, + "created_at" TIMESTAMP(3) NOT NULL, + "created_by" TEXT NOT NULL, + "updated_at" TIMESTAMP(3) NOT NULL, + "updated_by" TEXT NOT NULL, + + CONSTRAINT "LiteLLM_InvitationLink_pkey" PRIMARY KEY ("id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_AuditLog" ( + "id" TEXT NOT NULL, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "changed_by" TEXT NOT NULL DEFAULT '', + "changed_by_api_key" TEXT NOT NULL DEFAULT '', + "action" TEXT NOT NULL, + "table_name" TEXT NOT NULL, + "object_id" TEXT NOT NULL, + "before_value" JSONB, + "updated_values" JSONB, + + CONSTRAINT "LiteLLM_AuditLog_pkey" PRIMARY KEY ("id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_DailyUserSpend" ( + "id" TEXT NOT NULL, + "user_id" TEXT NOT NULL, + "date" TEXT NOT NULL, + "api_key" TEXT NOT NULL, + "model" TEXT NOT NULL, + "model_group" TEXT, + "custom_llm_provider" TEXT, + "prompt_tokens" INTEGER NOT NULL DEFAULT 0, + "completion_tokens" INTEGER NOT NULL DEFAULT 0, + "spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0, + "api_requests" INTEGER NOT NULL DEFAULT 0, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL, + + CONSTRAINT "LiteLLM_DailyUserSpend_pkey" PRIMARY KEY ("id") +); + +-- CreateTable +CREATE TABLE "LiteLLM_CronJob" ( + "cronjob_id" TEXT NOT NULL, + "pod_id" TEXT NOT NULL, + "status" "JobStatus" NOT NULL DEFAULT 'INACTIVE', + "last_updated" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "ttl" TIMESTAMP(3) NOT NULL, + + CONSTRAINT "LiteLLM_CronJob_pkey" PRIMARY KEY ("cronjob_id") +); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_CredentialsTable_credential_name_key" ON "LiteLLM_CredentialsTable"("credential_name"); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_TeamTable_model_id_key" ON "LiteLLM_TeamTable"("model_id"); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_UserTable_sso_user_id_key" ON "LiteLLM_UserTable"("sso_user_id"); + +-- CreateIndex +CREATE INDEX "LiteLLM_SpendLogs_startTime_idx" ON "LiteLLM_SpendLogs"("startTime"); + +-- CreateIndex +CREATE INDEX "LiteLLM_SpendLogs_end_user_idx" ON "LiteLLM_SpendLogs"("end_user"); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_OrganizationMembership_user_id_organization_id_key" ON "LiteLLM_OrganizationMembership"("user_id", "organization_id"); + +-- CreateIndex +CREATE INDEX "LiteLLM_DailyUserSpend_date_idx" ON "LiteLLM_DailyUserSpend"("date"); + +-- CreateIndex +CREATE INDEX "LiteLLM_DailyUserSpend_user_id_idx" ON "LiteLLM_DailyUserSpend"("user_id"); + +-- CreateIndex +CREATE INDEX "LiteLLM_DailyUserSpend_api_key_idx" ON "LiteLLM_DailyUserSpend"("api_key"); + +-- CreateIndex +CREATE INDEX "LiteLLM_DailyUserSpend_model_idx" ON "LiteLLM_DailyUserSpend"("model"); + +-- CreateIndex +CREATE UNIQUE INDEX "LiteLLM_DailyUserSpend_user_id_date_api_key_model_custom_ll_key" ON "LiteLLM_DailyUserSpend"("user_id", "date", "api_key", "model", "custom_llm_provider"); + +-- AddForeignKey +ALTER TABLE "LiteLLM_OrganizationTable" ADD CONSTRAINT "LiteLLM_OrganizationTable_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE RESTRICT ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_TeamTable" ADD CONSTRAINT "LiteLLM_TeamTable_organization_id_fkey" FOREIGN KEY ("organization_id") REFERENCES "LiteLLM_OrganizationTable"("organization_id") ON DELETE SET NULL ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_TeamTable" ADD CONSTRAINT "LiteLLM_TeamTable_model_id_fkey" FOREIGN KEY ("model_id") REFERENCES "LiteLLM_ModelTable"("id") ON DELETE SET NULL ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_UserTable" ADD CONSTRAINT "LiteLLM_UserTable_organization_id_fkey" FOREIGN KEY ("organization_id") REFERENCES "LiteLLM_OrganizationTable"("organization_id") ON DELETE SET NULL ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_VerificationToken" ADD CONSTRAINT "LiteLLM_VerificationToken_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE SET NULL ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_VerificationToken" ADD CONSTRAINT "LiteLLM_VerificationToken_organization_id_fkey" FOREIGN KEY ("organization_id") REFERENCES "LiteLLM_OrganizationTable"("organization_id") ON DELETE SET NULL ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_EndUserTable" ADD CONSTRAINT "LiteLLM_EndUserTable_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE SET NULL ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_TeamMembership" ADD CONSTRAINT "LiteLLM_TeamMembership_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE SET NULL ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_OrganizationMembership" ADD CONSTRAINT "LiteLLM_OrganizationMembership_user_id_fkey" FOREIGN KEY ("user_id") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE RESTRICT ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_OrganizationMembership" ADD CONSTRAINT "LiteLLM_OrganizationMembership_organization_id_fkey" FOREIGN KEY ("organization_id") REFERENCES "LiteLLM_OrganizationTable"("organization_id") ON DELETE RESTRICT ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_OrganizationMembership" ADD CONSTRAINT "LiteLLM_OrganizationMembership_budget_id_fkey" FOREIGN KEY ("budget_id") REFERENCES "LiteLLM_BudgetTable"("budget_id") ON DELETE SET NULL ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_InvitationLink" ADD CONSTRAINT "LiteLLM_InvitationLink_user_id_fkey" FOREIGN KEY ("user_id") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE RESTRICT ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_InvitationLink" ADD CONSTRAINT "LiteLLM_InvitationLink_created_by_fkey" FOREIGN KEY ("created_by") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE RESTRICT ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "LiteLLM_InvitationLink" ADD CONSTRAINT "LiteLLM_InvitationLink_updated_by_fkey" FOREIGN KEY ("updated_by") REFERENCES "LiteLLM_UserTable"("user_id") ON DELETE RESTRICT ON UPDATE CASCADE; diff --git a/schema.prisma b/schema.prisma index 7acb5e2615..39ba0b10d2 100644 --- a/schema.prisma +++ b/schema.prisma @@ -336,3 +336,18 @@ model LiteLLM_DailyUserSpend { @@index([api_key]) @@index([model]) } + + +// Track the status of cron jobs running. Only allow one pod to run the job at a time +model LiteLLM_CronJob { + cronjob_id String @id @default(cuid()) // Unique ID for the record + pod_id String // Unique identifier for the pod acting as the leader + status JobStatus @default(INACTIVE) // Status of the cron job (active or inactive) + last_updated DateTime @default(now()) // Timestamp for the last update of the cron job record + ttl DateTime // Time when the leader's lease expires +} + +enum JobStatus { + ACTIVE + INACTIVE +} diff --git a/tests/litellm/proxy/db/test_pod_lock_manager.py b/tests/litellm/proxy/db/test_pod_lock_manager.py new file mode 100644 index 0000000000..bce7b66409 --- /dev/null +++ b/tests/litellm/proxy/db/test_pod_lock_manager.py @@ -0,0 +1,320 @@ +import json +import os +import sys +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path + +from litellm.constants import DEFAULT_CRON_JOB_LOCK_TTL_SECONDS +from litellm.proxy.db.pod_lock_manager import PodLockManager + + +# Mock Prisma client class +class MockPrismaClient: + def __init__(self): + self.db = MagicMock() + self.db.litellm_cronjob = AsyncMock() + + +@pytest.fixture +def mock_prisma(monkeypatch): + mock_client = MockPrismaClient() + + # Mock the prisma_client import in proxy_server + def mock_get_prisma(): + return mock_client + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_client) + return mock_client + + +@pytest.fixture +def pod_lock_manager(): + return PodLockManager(cronjob_id="test_job") + + +@pytest.mark.asyncio +async def test_acquire_lock_success(pod_lock_manager, mock_prisma): + """ + Test that the lock is acquired successfully when no existing lock exists + """ + # Mock find_unique to return None (no existing lock) + mock_prisma.db.litellm_cronjob.find_unique.return_value = None + + # Mock successful creation of new lock + mock_response = AsyncMock() + mock_response.status = "ACTIVE" + mock_response.pod_id = pod_lock_manager.pod_id + mock_prisma.db.litellm_cronjob.create.return_value = mock_response + + result = await pod_lock_manager.acquire_lock() + assert result == True + + # Verify find_unique was called + mock_prisma.db.litellm_cronjob.find_unique.assert_called_once() + # Verify create was called with correct parameters + mock_prisma.db.litellm_cronjob.create.assert_called_once() + call_args = mock_prisma.db.litellm_cronjob.create.call_args[1] + assert call_args["data"]["cronjob_id"] == "test_job" + assert call_args["data"]["pod_id"] == pod_lock_manager.pod_id + assert call_args["data"]["status"] == "ACTIVE" + + +@pytest.mark.asyncio +async def test_acquire_lock_existing_active(pod_lock_manager, mock_prisma): + """ + Test that the lock is not acquired if there's an active lock by different pod + """ + # Mock existing active lock + mock_existing = AsyncMock() + mock_existing.status = "ACTIVE" + mock_existing.pod_id = "different_pod_id" + mock_existing.ttl = datetime.now(timezone.utc) + timedelta(seconds=30) # Future TTL + mock_prisma.db.litellm_cronjob.find_unique.return_value = mock_existing + + result = await pod_lock_manager.acquire_lock() + assert result == False + + # Verify find_unique was called but update/create were not + mock_prisma.db.litellm_cronjob.find_unique.assert_called_once() + mock_prisma.db.litellm_cronjob.update.assert_not_called() + mock_prisma.db.litellm_cronjob.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_acquire_lock_expired(pod_lock_manager, mock_prisma): + """ + Test that the lock can be acquired if existing lock is expired + """ + # Mock existing expired lock + mock_existing = AsyncMock() + mock_existing.status = "ACTIVE" + mock_existing.pod_id = "different_pod_id" + mock_existing.ttl = datetime.now(timezone.utc) - timedelta(seconds=30) # Past TTL + mock_prisma.db.litellm_cronjob.find_unique.return_value = mock_existing + + # Mock successful update + mock_updated = AsyncMock() + mock_updated.pod_id = pod_lock_manager.pod_id + mock_prisma.db.litellm_cronjob.update.return_value = mock_updated + + result = await pod_lock_manager.acquire_lock() + assert result == True + + # Verify both find_unique and update were called + mock_prisma.db.litellm_cronjob.find_unique.assert_called_once() + mock_prisma.db.litellm_cronjob.update.assert_called_once() + + +@pytest.mark.asyncio +async def test_renew_lock(pod_lock_manager, mock_prisma): + """ + Test that the renew lock calls the DB update method with the correct parameters + """ + mock_prisma.db.litellm_cronjob.update.return_value = AsyncMock() + + await pod_lock_manager.renew_lock() + + # Verify update was called with correct parameters + mock_prisma.db.litellm_cronjob.update.assert_called_once() + call_args = mock_prisma.db.litellm_cronjob.update.call_args[1] + assert call_args["where"]["cronjob_id"] == "test_job" + assert call_args["where"]["pod_id"] == pod_lock_manager.pod_id + assert "ttl" in call_args["data"] + assert "last_updated" in call_args["data"] + + +@pytest.mark.asyncio +async def test_release_lock(pod_lock_manager, mock_prisma): + """ + Test that the release lock calls the DB update method with the correct parameters + + specifically, the status should be set to INACTIVE + """ + mock_prisma.db.litellm_cronjob.update.return_value = AsyncMock() + + await pod_lock_manager.release_lock() + + # Verify update was called with correct parameters + mock_prisma.db.litellm_cronjob.update.assert_called_once() + call_args = mock_prisma.db.litellm_cronjob.update.call_args[1] + assert call_args["where"]["cronjob_id"] == "test_job" + assert call_args["where"]["pod_id"] == pod_lock_manager.pod_id + assert call_args["data"]["status"] == "INACTIVE" + + +@pytest.mark.asyncio +async def test_prisma_client_none(pod_lock_manager, monkeypatch): + # Mock prisma_client as None + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + + # Test all methods with None client + assert await pod_lock_manager.acquire_lock() == False + assert await pod_lock_manager.renew_lock() == False + assert await pod_lock_manager.release_lock() == False + + +@pytest.mark.asyncio +async def test_database_error_handling(pod_lock_manager, mock_prisma): + # Mock database errors + mock_prisma.db.litellm_cronjob.upsert.side_effect = Exception("Database error") + mock_prisma.db.litellm_cronjob.update.side_effect = Exception("Database error") + + # Test error handling in all methods + assert await pod_lock_manager.acquire_lock() == False + await pod_lock_manager.renew_lock() # Should not raise exception + await pod_lock_manager.release_lock() # Should not raise exception + + +@pytest.mark.asyncio +async def test_acquire_lock_inactive_status(pod_lock_manager, mock_prisma): + """ + Test that the lock can be acquired if existing lock is INACTIVE + """ + # Mock existing inactive lock + mock_existing = AsyncMock() + mock_existing.status = "INACTIVE" + mock_existing.pod_id = "different_pod_id" + mock_existing.ttl = datetime.now(timezone.utc) + timedelta(seconds=30) + mock_prisma.db.litellm_cronjob.find_unique.return_value = mock_existing + + # Mock successful update + mock_updated = AsyncMock() + mock_updated.pod_id = pod_lock_manager.pod_id + mock_prisma.db.litellm_cronjob.update.return_value = mock_updated + + result = await pod_lock_manager.acquire_lock() + assert result == True + + mock_prisma.db.litellm_cronjob.update.assert_called_once() + + +@pytest.mark.asyncio +async def test_acquire_lock_same_pod(pod_lock_manager, mock_prisma): + """ + Test that the lock returns True if the same pod already holds the lock + """ + # Mock existing active lock held by same pod + mock_existing = AsyncMock() + mock_existing.status = "ACTIVE" + mock_existing.pod_id = pod_lock_manager.pod_id + mock_existing.ttl = datetime.now(timezone.utc) + timedelta(seconds=30) + mock_prisma.db.litellm_cronjob.find_unique.return_value = mock_existing + + result = await pod_lock_manager.acquire_lock() + assert result == True + + # Verify no update was needed + mock_prisma.db.litellm_cronjob.update.assert_not_called() + mock_prisma.db.litellm_cronjob.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_acquire_lock_race_condition(pod_lock_manager, mock_prisma): + """ + Test handling of potential race conditions during lock acquisition + """ + # First find_unique returns None + mock_prisma.db.litellm_cronjob.find_unique.return_value = None + + # But create raises unique constraint violation + mock_prisma.db.litellm_cronjob.create.side_effect = Exception( + "Unique constraint violation" + ) + + result = await pod_lock_manager.acquire_lock() + assert result == False + + +@pytest.mark.asyncio +async def test_ttl_calculation(pod_lock_manager, mock_prisma): + """ + Test that TTL is calculated correctly when acquiring lock + """ + mock_prisma.db.litellm_cronjob.find_unique.return_value = None + mock_prisma.db.litellm_cronjob.create.return_value = AsyncMock() + + await pod_lock_manager.acquire_lock() + + call_args = mock_prisma.db.litellm_cronjob.create.call_args[1] + ttl = call_args["data"]["ttl"] + + # Verify TTL is in the future by DEFAULT_CRON_JOB_LOCK_TTL_SECONDS + expected_ttl = datetime.now(timezone.utc) + timedelta( + seconds=DEFAULT_CRON_JOB_LOCK_TTL_SECONDS + ) + assert abs((ttl - expected_ttl).total_seconds()) < 1 # Allow 1 second difference + + +@pytest.mark.asyncio +async def test_concurrent_lock_acquisition_simulation(mock_prisma): + """ + Simulate multiple pods trying to acquire the lock simultaneously + """ + pod1 = PodLockManager(cronjob_id="test_job") + pod2 = PodLockManager(cronjob_id="test_job") + pod3 = PodLockManager(cronjob_id="test_job") + + # Simulate first pod getting the lock + mock_prisma.db.litellm_cronjob.find_unique.return_value = None + mock_response = AsyncMock() + mock_response.pod_id = pod1.pod_id + mock_response.status = "ACTIVE" + mock_prisma.db.litellm_cronjob.create.return_value = mock_response + + # First pod should get the lock + result1 = await pod1.acquire_lock() + assert result1 == True + + # Simulate other pods trying to acquire same lock immediately after + mock_existing = AsyncMock() + mock_existing.status = "ACTIVE" + mock_existing.pod_id = pod1.pod_id + mock_existing.ttl = datetime.now(timezone.utc) + timedelta(seconds=30) + mock_prisma.db.litellm_cronjob.find_unique.return_value = mock_existing + + # Other pods should fail to acquire + result2 = await pod2.acquire_lock() + result3 = await pod3.acquire_lock() + assert result2 == False + assert result3 == False + + +@pytest.mark.asyncio +async def test_lock_takeover_race_condition(mock_prisma): + """ + Test scenario where multiple pods try to take over an expired lock + """ + pod1 = PodLockManager(cronjob_id="test_job") + pod2 = PodLockManager(cronjob_id="test_job") + + # Simulate expired lock + mock_existing = AsyncMock() + mock_existing.status = "ACTIVE" + mock_existing.pod_id = "old_pod" + mock_existing.ttl = datetime.now(timezone.utc) - timedelta(seconds=30) + mock_prisma.db.litellm_cronjob.find_unique.return_value = mock_existing + + # Simulate pod1's update succeeding + mock_update1 = AsyncMock() + mock_update1.pod_id = pod1.pod_id + mock_prisma.db.litellm_cronjob.update.return_value = mock_update1 + + # First pod should successfully take over + result1 = await pod1.acquire_lock() + assert result1 == True + + # Simulate pod2's update failing due to race condition + mock_prisma.db.litellm_cronjob.update.side_effect = Exception( + "Row was updated by another transaction" + ) + + # Second pod should fail to take over + result2 = await pod2.acquire_lock() + assert result2 == False diff --git a/tests/litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/litellm/proxy/hooks/test_proxy_track_cost_callback.py index 1e3b22ae2d..8850436329 100644 --- a/tests/litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -47,7 +47,8 @@ async def test_async_post_call_failure_hook(): # Mock update_database function with patch( - "litellm.proxy.proxy_server.update_database", new_callable=AsyncMock + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, ) as mock_update_database: # Call the method await logger.async_post_call_failure_hook( diff --git a/tests/litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 609fbbcc52..c1e94138a3 100644 --- a/tests/litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -421,7 +421,8 @@ class TestSpendLogsPayload: # litellm._turn_on_debug() with patch.object( - litellm.proxy.proxy_server, "_set_spend_logs_payload" + litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter, + "_set_spend_logs_payload", ) as mock_client, patch.object(litellm.proxy.proxy_server, "prisma_client"): response = await litellm.acompletion( model="gpt-4o", @@ -514,7 +515,8 @@ class TestSpendLogsPayload: client = AsyncHTTPHandler() with patch.object( - litellm.proxy.proxy_server, "_set_spend_logs_payload" + litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter, + "_set_spend_logs_payload", ) as mock_client, patch.object( litellm.proxy.proxy_server, "prisma_client" ), patch.object( @@ -609,7 +611,8 @@ class TestSpendLogsPayload: ) with patch.object( - litellm.proxy.proxy_server, "_set_spend_logs_payload" + litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter, + "_set_spend_logs_payload", ) as mock_client, patch.object( litellm.proxy.proxy_server, "prisma_client" ), patch.object( diff --git a/tests/local_testing/test_update_spend.py b/tests/local_testing/test_update_spend.py index 6aeae851ab..fffa3062d7 100644 --- a/tests/local_testing/test_update_spend.py +++ b/tests/local_testing/test_update_spend.py @@ -93,7 +93,7 @@ def prisma_client(): @pytest.mark.asyncio async def test_batch_update_spend(prisma_client): - prisma_client.user_list_transactons["test-litellm-user-5"] = 23 + prisma_client.user_list_transactions["test-litellm-user-5"] = 23 setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") await litellm.proxy.proxy_server.prisma_client.connect() diff --git a/tests/otel_tests/local_test_spend_accuracy_tests.py b/tests/otel_tests/local_test_spend_accuracy_tests.py new file mode 100644 index 0000000000..6d756219c7 --- /dev/null +++ b/tests/otel_tests/local_test_spend_accuracy_tests.py @@ -0,0 +1,279 @@ +import pytest +import asyncio +import aiohttp +import json +from httpx import AsyncClient +from typing import Any, Optional +import uuid + +""" +Tests to run + +Basic Tests: +1. Basic Spend Accuracy Test: + - 1 Request costs $0.037 + - Make 12 requests + - Expect the spend for each of the following to be 12 * $0.037 + Key: $0.444 (call /info endpoint for each object to validate) + Team: $0.444 + User: $0.444 + Org: $0.444 + End User: $0.444 + +2. Long term spend accuracy test (with 2 bursts of requests) + - 1 Request costs $0.037 + - Burst 1: 12 requests + - Burst 2: 22 requests + + - Expect the spend for each of the following to be (12 + 22) * $0.037 + Key: $1.296 + Team: $1.296 + User: $1.296 + Org: $1.296 + End User: $1.296 + +Additional Test Scenarios: + +3. Concurrent Request Accuracy Test: + - Make 20 concurrent requests + - Verify total spend is 20 * $0.037 + - Check for race conditions in spend tracking + +4. Error Case Test: + - Make 10 successful requests ($0.037 each) + - Make 5 failed requests + - Verify spend is only counted for successful requests (10 * $0.037) + +5. Mixed Request Type Test: + - Make different types of requests with varying costs + - Verify accurate total spend calculation +""" + + +async def create_organization(session, organization_alias: str): + """Helper function to create a new organization""" + url = "http://0.0.0.0:4002/organization/new" + headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + data = {"organization_alias": organization_alias} + async with session.post(url, headers=headers, json=data) as response: + return await response.json() + + +async def create_team(session, org_id: str): + """Helper function to create a new team under an organization""" + url = "http://0.0.0.0:4002/team/new" + headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + data = {"organization_id": org_id, "team_alias": f"test-team-{uuid.uuid4()}"} + async with session.post(url, headers=headers, json=data) as response: + return await response.json() + + +async def create_user(session, org_id: str): + """Helper function to create a new user""" + url = "http://0.0.0.0:4002/user/new" + headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + data = {"user_name": f"test-user-{uuid.uuid4()}"} + async with session.post(url, headers=headers, json=data) as response: + return await response.json() + + +async def generate_key(session, user_id: str, team_id: str): + """Helper function to generate a key for a specific user and team""" + url = "http://0.0.0.0:4002/key/generate" + headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + data = {"user_id": user_id, "team_id": team_id} + async with session.post(url, headers=headers, json=data) as response: + return await response.json() + + +async def chat_completion(session, key: str): + """Make a chat completion request""" + from openai import AsyncOpenAI + import uuid + + client = AsyncOpenAI(api_key=key, base_url="http://0.0.0.0:4002/v1") + + response = await client.chat.completions.create( + model="fake-openai-endpoint", + messages=[{"role": "user", "content": f"Test message {uuid.uuid4()}"}], + ) + return response + + +async def get_spend_info(session, entity_type: str, entity_id: str): + """Helper function to get spend information for an entity""" + url = f"http://0.0.0.0:4002/{entity_type}/info" + headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"} + if entity_type == "key": + data = {"key": entity_id} + else: + data = {f"{entity_type}_id": entity_id} + + async with session.get(url, headers=headers, params=data) as response: + return await response.json() + + +@pytest.mark.asyncio +async def test_basic_spend_accuracy(): + """ + Test basic spend accuracy across different entities: + 1. Create org, team, user, and key + 2. Make 12 requests at $0.037 each + 3. Verify spend accuracy for key, team, user, org, and end user + """ + SPEND_PER_REQUEST = 3.75 * 10**-5 + NUM_LLM_REQUESTS = 20 + expected_spend = NUM_LLM_REQUESTS * SPEND_PER_REQUEST # 12 requests at $0.037 each + + # Add tolerance constant at the top of the test + TOLERANCE = 1e-10 # Small number to account for floating-point precision + + async with aiohttp.ClientSession() as session: + # Create organization + org_response = await create_organization( + session=session, organization_alias=f"test-org-{uuid.uuid4()}" + ) + print("org_response: ", org_response) + org_id = org_response["organization_id"] + + # Create team under organization + team_response = await create_team(session, org_id) + print("team_response: ", team_response) + team_id = team_response["team_id"] + + # Create user + user_response = await create_user(session, org_id) + print("user_response: ", user_response) + user_id = user_response["user_id"] + + # Generate key + key_response = await generate_key(session, user_id, team_id) + print("key_response: ", key_response) + key = key_response["key"] + + # Make 12 requests + for _ in range(NUM_LLM_REQUESTS): + response = await chat_completion(session, key) + print("response: ", response) + + # wait 10 seconds for spend to be updated + await asyncio.sleep(15) + + # Get spend information for each entity + key_info = await get_spend_info(session, "key", key) + print("key_info: ", key_info) + team_info = await get_spend_info(session, "team", team_id) + print("team_info: ", team_info) + user_info = await get_spend_info(session, "user", user_id) + print("user_info: ", user_info) + org_info = await get_spend_info(session, "organization", org_id) + print("org_info: ", org_info) + + # Verify spend for each entity + assert ( + abs(key_info["info"]["spend"] - expected_spend) < TOLERANCE + ), f"Key spend {key_info['info']['spend']} does not match expected {expected_spend}" + + assert ( + abs(user_info["user_info"]["spend"] - expected_spend) < TOLERANCE + ), f"User spend {user_info['info']['spend']} does not match expected {expected_spend}" + + assert ( + abs(team_info["team_info"]["spend"] - expected_spend) < TOLERANCE + ), f"Team spend {team_info['team_info']['spend']} does not match expected {expected_spend}" + + assert ( + abs(org_info["spend"] - expected_spend) < TOLERANCE + ), f"Organization spend {org_info['spend']} does not match expected {expected_spend}" + + +@pytest.mark.asyncio +async def test_long_term_spend_accuracy_with_bursts(): + """ + Test long-term spend accuracy with multiple bursts of requests: + 1. Create org, team, user, and key + 2. Burst 1: Make 12 requests + 3. Burst 2: Make 22 more requests + 4. Verify the total spend (34 requests) is tracked accurately across all entities + """ + SPEND_PER_REQUEST = 3.75 * 10**-5 # Cost per request + BURST_1_REQUESTS = 22 # Number of requests in first burst + BURST_2_REQUESTS = 12 # Number of requests in second burst + TOTAL_REQUESTS = BURST_1_REQUESTS + BURST_2_REQUESTS + expected_spend = TOTAL_REQUESTS * SPEND_PER_REQUEST + + # Tolerance for floating-point comparison + TOLERANCE = 1e-10 + + async with aiohttp.ClientSession() as session: + # Create organization + org_response = await create_organization( + session=session, organization_alias=f"test-org-{uuid.uuid4()}" + ) + print("org_response: ", org_response) + org_id = org_response["organization_id"] + + # Create team under organization + team_response = await create_team(session, org_id) + print("team_response: ", team_response) + team_id = team_response["team_id"] + + # Create user + user_response = await create_user(session, org_id) + print("user_response: ", user_response) + user_id = user_response["user_id"] + + # Generate key + key_response = await generate_key(session, user_id, team_id) + print("key_response: ", key_response) + key = key_response["key"] + + # First burst: 12 requests + print(f"Starting first burst of {BURST_1_REQUESTS} requests...") + for i in range(BURST_1_REQUESTS): + response = await chat_completion(session, key) + print(f"Burst 1 - Request {i+1}/{BURST_1_REQUESTS} completed") + + # Wait for spend to be updated + await asyncio.sleep(8) + + # Check intermediate spend + intermediate_key_info = await get_spend_info(session, "key", key) + print(f"After Burst 1 - Key spend: {intermediate_key_info['info']['spend']}") + + # Second burst: 22 requests + print(f"Starting second burst of {BURST_2_REQUESTS} requests...") + for i in range(BURST_2_REQUESTS): + response = await chat_completion(session, key) + print(f"Burst 2 - Request {i+1}/{BURST_2_REQUESTS} completed") + + # Wait for spend to be updated + await asyncio.sleep(8) + + # Get final spend information for each entity + key_info = await get_spend_info(session, "key", key) + team_info = await get_spend_info(session, "team", team_id) + user_info = await get_spend_info(session, "user", user_id) + org_info = await get_spend_info(session, "organization", org_id) + + print(f"Final key spend: {key_info['info']['spend']}") + print(f"Final team spend: {team_info['team_info']['spend']}") + print(f"Final user spend: {user_info['user_info']['spend']}") + print(f"Final org spend: {org_info['spend']}") + + # Verify total spend for each entity + assert ( + abs(key_info["info"]["spend"] - expected_spend) < TOLERANCE + ), f"Key spend {key_info['info']['spend']} does not match expected {expected_spend}" + + assert ( + abs(user_info["user_info"]["spend"] - expected_spend) < TOLERANCE + ), f"User spend {user_info['user_info']['spend']} does not match expected {expected_spend}" + + assert ( + abs(team_info["team_info"]["spend"] - expected_spend) < TOLERANCE + ), f"Team spend {team_info['team_info']['spend']} does not match expected {expected_spend}" + + assert ( + abs(org_info["spend"] - expected_spend) < TOLERANCE + ), f"Organization spend {org_info['spend']} does not match expected {expected_spend}" diff --git a/tests/proxy_unit_tests/test_e2e_pod_lock_manager.py b/tests/proxy_unit_tests/test_e2e_pod_lock_manager.py new file mode 100644 index 0000000000..7d36bb4791 --- /dev/null +++ b/tests/proxy_unit_tests/test_e2e_pod_lock_manager.py @@ -0,0 +1,354 @@ +import os +import sys +import traceback +import uuid +from datetime import datetime, timezone, timedelta + +from dotenv import load_dotenv +from fastapi import Request +from fastapi.routing import APIRoute +import httpx + +load_dotenv() +import io +import os +import time + +# this file is to test litellm/proxy + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path +import asyncio +import logging + +import pytest +from litellm.proxy.db.pod_lock_manager import PodLockManager +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.proxy.management_endpoints.internal_user_endpoints import ( + new_user, + user_info, + user_update, +) +from litellm.proxy.auth.auth_checks import get_key_object +from litellm.proxy.management_endpoints.key_management_endpoints import ( + delete_key_fn, + generate_key_fn, + generate_key_helper_fn, + info_key_fn, + list_keys, + regenerate_key_fn, + update_key_fn, +) +from litellm.proxy.management_endpoints.team_endpoints import ( + new_team, + team_info, + update_team, +) +from litellm.proxy.proxy_server import ( + LitellmUserRoles, + audio_transcriptions, + chat_completion, + completion, + embeddings, + image_generation, + model_list, + moderations, + user_api_key_auth, +) +from litellm.proxy.management_endpoints.customer_endpoints import ( + new_end_user, +) +from litellm.proxy.spend_tracking.spend_management_endpoints import ( + global_spend, + spend_key_fn, + spend_user_fn, + view_spend_logs, +) +from litellm.proxy.utils import PrismaClient, ProxyLogging, hash_token, update_spend + +verbose_proxy_logger.setLevel(level=logging.DEBUG) + +from starlette.datastructures import URL + +from litellm.caching.caching import DualCache +from litellm.proxy._types import ( + DynamoDBArgs, + GenerateKeyRequest, + KeyRequest, + LiteLLM_UpperboundKeyGenerateParams, + NewCustomerRequest, + NewTeamRequest, + NewUserRequest, + ProxyErrorTypes, + ProxyException, + UpdateKeyRequest, + UpdateTeamRequest, + UpdateUserRequest, + UserAPIKeyAuth, +) + +proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + + +request_data = { + "model": "azure-gpt-3.5", + "messages": [ + {"role": "user", "content": "this is my new test. respond in 50 lines"} + ], +} + + +@pytest.fixture +def prisma_client(): + from litellm.proxy.proxy_cli import append_query_params + + ### add connection pool + pool timeout args + params = {"connection_limit": 100, "pool_timeout": 60} + database_url = os.getenv("DATABASE_URL") + modified_url = append_query_params(database_url, params) + os.environ["DATABASE_URL"] = modified_url + + # Assuming PrismaClient is a class that needs to be instantiated + prisma_client = PrismaClient( + database_url=os.environ["DATABASE_URL"], proxy_logging_obj=proxy_logging_obj + ) + + # Reset litellm.proxy.proxy_server.prisma_client to None + litellm.proxy.proxy_server.litellm_proxy_budget_name = ( + f"litellm-proxy-budget-{time.time()}" + ) + litellm.proxy.proxy_server.user_custom_key_generate = None + + return prisma_client + + +async def setup_db_connection(prisma_client): + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + await litellm.proxy.proxy_server.prisma_client.connect() + + +@pytest.mark.asyncio +async def test_pod_lock_acquisition_when_no_active_lock(prisma_client): + """Test if a pod can acquire a lock when no lock is active""" + await setup_db_connection(prisma_client) + + cronjob_id = str(uuid.uuid4()) + lock_manager = PodLockManager(cronjob_id=cronjob_id) + + # Attempt to acquire lock + result = await lock_manager.acquire_lock() + + assert result == True, "Pod should be able to acquire lock when no lock exists" + + # Verify in database + lock_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + assert lock_record.status == "ACTIVE" + assert lock_record.pod_id == lock_manager.pod_id + + +@pytest.mark.asyncio +async def test_pod_lock_acquisition_after_completion(prisma_client): + """Test if a new pod can acquire lock after previous pod completes""" + await setup_db_connection(prisma_client) + + cronjob_id = str(uuid.uuid4()) + # First pod acquires and releases lock + first_lock_manager = PodLockManager(cronjob_id=cronjob_id) + await first_lock_manager.acquire_lock() + await first_lock_manager.release_lock() + + # Second pod attempts to acquire lock + second_lock_manager = PodLockManager(cronjob_id=cronjob_id) + result = await second_lock_manager.acquire_lock() + + assert result == True, "Second pod should acquire lock after first pod releases it" + + # Verify in database + lock_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + assert lock_record.status == "ACTIVE" + assert lock_record.pod_id == second_lock_manager.pod_id + + +@pytest.mark.asyncio +async def test_pod_lock_acquisition_after_expiry(prisma_client): + """Test if a new pod can acquire lock after previous pod's lock expires""" + await setup_db_connection(prisma_client) + + cronjob_id = str(uuid.uuid4()) + # First pod acquires lock + first_lock_manager = PodLockManager(cronjob_id=cronjob_id) + await first_lock_manager.acquire_lock() + + # release the lock from the first pod + await first_lock_manager.release_lock() + + # Second pod attempts to acquire lock + second_lock_manager = PodLockManager(cronjob_id=cronjob_id) + result = await second_lock_manager.acquire_lock() + + assert ( + result == True + ), "Second pod should acquire lock after first pod's lock expires" + + # Verify in database + lock_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + assert lock_record.status == "ACTIVE" + assert lock_record.pod_id == second_lock_manager.pod_id + + +@pytest.mark.asyncio +async def test_pod_lock_release(prisma_client): + """Test if a pod can successfully release its lock""" + await setup_db_connection(prisma_client) + + cronjob_id = str(uuid.uuid4()) + lock_manager = PodLockManager(cronjob_id=cronjob_id) + + # Acquire and then release lock + await lock_manager.acquire_lock() + await lock_manager.release_lock() + + # Verify in database + lock_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + assert lock_record.status == "INACTIVE" + + +@pytest.mark.asyncio +async def test_concurrent_lock_acquisition(prisma_client): + """Test that only one pod can acquire the lock when multiple pods try simultaneously""" + await setup_db_connection(prisma_client) + + cronjob_id = str(uuid.uuid4()) + # Create multiple lock managers simulating different pods + lock_manager1 = PodLockManager(cronjob_id=cronjob_id) + lock_manager2 = PodLockManager(cronjob_id=cronjob_id) + lock_manager3 = PodLockManager(cronjob_id=cronjob_id) + + # Try to acquire locks concurrently + results = await asyncio.gather( + lock_manager1.acquire_lock(), + lock_manager2.acquire_lock(), + lock_manager3.acquire_lock(), + ) + + # Only one should succeed + print("all results=", results) + assert sum(results) == 1, "Only one pod should acquire the lock" + + # Verify in database + lock_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + assert lock_record.status == "ACTIVE" + assert lock_record.pod_id in [ + lock_manager1.pod_id, + lock_manager2.pod_id, + lock_manager3.pod_id, + ] + + +@pytest.mark.asyncio +async def test_lock_renewal(prisma_client): + """Test that a pod can successfully renew its lock""" + await setup_db_connection(prisma_client) + + cronjob_id = str(uuid.uuid4()) + lock_manager = PodLockManager(cronjob_id=cronjob_id) + + # Acquire initial lock + await lock_manager.acquire_lock() + + # Get initial TTL + initial_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + initial_ttl = initial_record.ttl + + # Wait a short time + await asyncio.sleep(1) + + # Renew the lock + await lock_manager.renew_lock() + + # Get updated record + renewed_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + + assert renewed_record.ttl > initial_ttl, "Lock TTL should be extended after renewal" + assert renewed_record.status == "ACTIVE" + assert renewed_record.pod_id == lock_manager.pod_id + + +@pytest.mark.asyncio +async def test_lock_acquisition_with_expired_ttl(prisma_client): + """Test that a pod can acquire a lock when existing lock has expired TTL""" + await setup_db_connection(prisma_client) + + cronjob_id = str(uuid.uuid4()) + first_lock_manager = PodLockManager(cronjob_id=cronjob_id) + + # First pod acquires lock + await first_lock_manager.acquire_lock() + + # Manually expire the TTL + expired_time = datetime.now(timezone.utc) - timedelta(seconds=10) + await prisma_client.db.litellm_cronjob.update( + where={"cronjob_id": cronjob_id}, data={"ttl": expired_time} + ) + + # Second pod tries to acquire without explicit release + second_lock_manager = PodLockManager(cronjob_id=cronjob_id) + result = await second_lock_manager.acquire_lock() + + assert result == True, "Should acquire lock when existing lock has expired TTL" + + # Verify in database + lock_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + assert lock_record.status == "ACTIVE" + assert lock_record.pod_id == second_lock_manager.pod_id + + +@pytest.mark.asyncio +async def test_release_expired_lock(prisma_client): + """Test that a pod cannot release a lock that has been taken over by another pod""" + await setup_db_connection(prisma_client) + + cronjob_id = str(uuid.uuid4()) + first_lock_manager = PodLockManager(cronjob_id=cronjob_id) + + # First pod acquires lock + await first_lock_manager.acquire_lock() + + # Manually expire the TTL + expired_time = datetime.now(timezone.utc) - timedelta(seconds=10) + await prisma_client.db.litellm_cronjob.update( + where={"cronjob_id": cronjob_id}, data={"ttl": expired_time} + ) + + # Second pod acquires the lock + second_lock_manager = PodLockManager(cronjob_id=cronjob_id) + await second_lock_manager.acquire_lock() + + # First pod attempts to release its lock + await first_lock_manager.release_lock() + + # Verify that second pod's lock is still active + lock_record = await prisma_client.db.litellm_cronjob.find_first( + where={"cronjob_id": cronjob_id} + ) + assert lock_record.status == "ACTIVE" + assert lock_record.pod_id == second_lock_manager.pod_id diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index dccf0d1842..d613118fc8 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -1509,7 +1509,7 @@ from litellm.proxy.utils import ProxyUpdateSpend async def test_end_user_transactions_reset(): # Setup mock_client = MagicMock() - mock_client.end_user_list_transactons = {"1": 10.0} # Bad log + mock_client.end_user_list_transactions = {"1": 10.0} # Bad log mock_client.db.tx = AsyncMock(side_effect=Exception("DB Error")) # Call function - should raise error @@ -1520,7 +1520,7 @@ async def test_end_user_transactions_reset(): # Verify cleanup happened assert ( - mock_client.end_user_list_transactons == {} + mock_client.end_user_list_transactions == {} ), "Transactions list should be empty after error" diff --git a/tests/proxy_unit_tests/test_unit_test_proxy_hooks.py b/tests/proxy_unit_tests/test_unit_test_proxy_hooks.py index 535f5bf019..129be6d754 100644 --- a/tests/proxy_unit_tests/test_unit_test_proxy_hooks.py +++ b/tests/proxy_unit_tests/test_unit_test_proxy_hooks.py @@ -5,6 +5,7 @@ from unittest.mock import Mock, patch, AsyncMock import pytest from fastapi import Request from litellm.proxy.utils import _get_redoc_url, _get_docs_url +from datetime import datetime sys.path.insert(0, os.path.abspath("../..")) import litellm @@ -22,16 +23,20 @@ async def test_disable_spend_logs(): with patch("litellm.proxy.proxy_server.disable_spend_logs", True), patch( "litellm.proxy.proxy_server.prisma_client", mock_prisma_client ): - from litellm.proxy.proxy_server import update_database + from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter # Call update_database with disable_spend_logs=True - await update_database( + await DBSpendUpdateWriter.update_database( token="fake-token", response_cost=0.1, user_id="user123", completion_response=None, - start_time="2024-01-01", - end_time="2024-01-01", + start_time=datetime.now(), + end_time=datetime.now(), + end_user_id="end_user_id", + team_id="team_id", + org_id="org_id", + kwargs={}, ) # Verify no spend logs were added assert len(mock_prisma_client.spend_log_transactions) == 0