diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 937baf9c382..86a85912af3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1479,7 +1479,8 @@ class LiteLLM_VerificationTokenView(LiteLLM_VerificationToken): # Check if the value is None and set the corresponding attribute if getattr(self, attr_name, None) is None: kwargs[attr_name] = value - + if key == "end_user_id" and value is not None and isinstance(value, int): + kwargs[key] = str(value) # Initialize the superclass super().__init__(**kwargs) @@ -2288,7 +2289,6 @@ class ProxyStateVariables(TypedDict): UI_TEAM_ID = "litellm-dashboard" - class JWTAuthBuilderResult(TypedDict): is_proxy_admin: bool team_object: Optional[LiteLLM_TeamTable] @@ -2301,6 +2301,7 @@ class JWTAuthBuilderResult(TypedDict): end_user_id: Optional[str] org_id: Optional[str] + class ClientSideFallbackModel(TypedDict, total=False): """ Dictionary passed when client configuring input diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index a669c592c3b..5a431e66f07 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2435,6 +2435,138 @@ async def reset_budget(prisma_client: PrismaClient): ) +class ProxyUpdateSpend: + @staticmethod + async def update_end_user_spend( + n_retry_times: int, prisma_client: PrismaClient, proxy_logging_obj: ProxyLogging + ): + 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 ( + end_user_id, + response_cost, + ) in prisma_client.end_user_list_transactons.items(): + if litellm.max_end_user_budget is not None: + pass + batcher.litellm_endusertable.upsert( + where={"user_id": end_user_id}, + data={ + "create": { + "user_id": end_user_id, + "spend": response_cost, + "blocked": False, + }, + "update": {"spend": {"increment": response_cost}}, + }, + ) + + 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 + ) + finally: + prisma_client.end_user_list_transactons = ( + {} + ) # reset the end user list transactions - prevent bad data from causing issues + + @staticmethod + async def update_spend_logs( + n_retry_times: int, + prisma_client: PrismaClient, + db_writer_client: Optional[HTTPHandler], + proxy_logging_obj: ProxyLogging, + ): + BATCH_SIZE = 100 # Preferred size of each batch to write to the database + MAX_LOGS_PER_INTERVAL = ( + 1000 # Maximum number of logs to flush in a single interval + ) + for i in range(n_retry_times + 1): + start_time = time.time() + logs_to_process = prisma_client.spend_log_transactions + try: + base_url = os.getenv("SPEND_LOGS_URL", None) + ## WRITE TO SEPARATE SERVER ## + if ( + len(prisma_client.spend_log_transactions) > 0 + and base_url is not None + and db_writer_client is not None + ): + + if not base_url.endswith("/"): + base_url += "/" + verbose_proxy_logger.debug("base_url: {}".format(base_url)) + response = await db_writer_client.post( + url=base_url + "spend/update", + data=json.dumps(prisma_client.spend_log_transactions), # type: ignore + headers={"Content-Type": "application/json"}, + ) + if response.status_code == 200: + prisma_client.spend_log_transactions = [] + else: ## (default) WRITE TO DB ## + logs_to_process = prisma_client.spend_log_transactions[ + :MAX_LOGS_PER_INTERVAL + ] + for j in range(0, len(logs_to_process), BATCH_SIZE): + # Create sublist for current batch, ensuring it doesn't exceed the BATCH_SIZE + batch = logs_to_process[j : j + BATCH_SIZE] + + # Convert datetime strings to Date objects + batch_with_dates = [ + prisma_client.jsonify_object( + { + **entry, + } + ) + for entry in batch + ] + + await prisma_client.db.litellm_spendlogs.create_many( + data=batch_with_dates, skip_duplicates=True # type: ignore + ) + + verbose_proxy_logger.debug( + f"Flushed {len(batch)} logs to the DB." + ) + + verbose_proxy_logger.debug( + f"{len(logs_to_process)} logs processed. Remaining in queue: {len(prisma_client.spend_log_transactions)}" + ) + break + except DB_CONNECTION_ERROR_TYPES as e: + if i is None: + i = 0 + if ( + i >= n_retry_times + ): # If we've reached the maximum number of retries raise the exception + _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) # type: ignore + except Exception as e: + _raise_failed_update_spend_exception( + e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj + ) + finally: + prisma_client.spend_log_transactions = ( + prisma_client.spend_log_transactions[len(logs_to_process) :] + ) + + async def update_spend( # noqa: PLR0915 prisma_client: PrismaClient, db_writer_client: Optional[HTTPHandler], @@ -2493,47 +2625,13 @@ async def update_spend( # noqa: PLR0915 ) ) if len(prisma_client.end_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 ( - end_user_id, - response_cost, - ) in prisma_client.end_user_list_transactons.items(): - if litellm.max_end_user_budget is not None: - pass - batcher.litellm_endusertable.upsert( - where={"user_id": end_user_id}, - data={ - "create": { - "user_id": end_user_id, - "spend": response_cost, - "blocked": False, - }, - "update": {"spend": {"increment": response_cost}}, - }, - ) - - prisma_client.end_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 - ) - + asyncio.create_task( + 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( @@ -2690,80 +2788,15 @@ async def update_spend( # noqa: PLR0915 "Spend Logs transactions: {}".format(len(prisma_client.spend_log_transactions)) ) - BATCH_SIZE = 100 # Preferred size of each batch to write to the database - MAX_LOGS_PER_INTERVAL = 1000 # Maximum number of logs to flush in a single interval - if len(prisma_client.spend_log_transactions) > 0: - for i in range(n_retry_times + 1): - start_time = time.time() - try: - base_url = os.getenv("SPEND_LOGS_URL", None) - ## WRITE TO SEPARATE SERVER ## - if ( - len(prisma_client.spend_log_transactions) > 0 - and base_url is not None - and db_writer_client is not None - ): - if not base_url.endswith("/"): - base_url += "/" - verbose_proxy_logger.debug("base_url: {}".format(base_url)) - response = await db_writer_client.post( - url=base_url + "spend/update", - data=json.dumps(prisma_client.spend_log_transactions), # type: ignore - headers={"Content-Type": "application/json"}, - ) - if response.status_code == 200: - prisma_client.spend_log_transactions = [] - else: ## (default) WRITE TO DB ## - logs_to_process = prisma_client.spend_log_transactions[ - :MAX_LOGS_PER_INTERVAL - ] - for j in range(0, len(logs_to_process), BATCH_SIZE): - # Create sublist for current batch, ensuring it doesn't exceed the BATCH_SIZE - batch = logs_to_process[j : j + BATCH_SIZE] - - # Convert datetime strings to Date objects - batch_with_dates = [ - prisma_client.jsonify_object( - { - **entry, - } - ) - for entry in batch - ] - - await prisma_client.db.litellm_spendlogs.create_many( - data=batch_with_dates, skip_duplicates=True # type: ignore - ) - - verbose_proxy_logger.debug( - f"Flushed {len(batch)} logs to the DB." - ) - # Remove the processed logs from spend_logs - prisma_client.spend_log_transactions = ( - prisma_client.spend_log_transactions[len(logs_to_process) :] - ) - - verbose_proxy_logger.debug( - f"{len(logs_to_process)} logs processed. Remaining in queue: {len(prisma_client.spend_log_transactions)}" - ) - break - except DB_CONNECTION_ERROR_TYPES as e: - if i is None: - i = 0 - if ( - i >= n_retry_times - ): # If we've reached the maximum number of retries raise the exception - _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) # type: ignore - except Exception as e: - _raise_failed_update_spend_exception( - e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj - ) + asyncio.create_task( + ProxyUpdateSpend.update_spend_logs( + n_retry_times=n_retry_times, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + db_writer_client=db_writer_client, + ) + ) def _raise_failed_update_spend_exception(