diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 9fb20224322..d8b8efeef47 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -2,6 +2,8 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if the cost has been tracked. """ +import uuid +from datetime import datetime from typing import TYPE_CHECKING, Optional, cast from litellm._logging import verbose_proxy_logger @@ -42,6 +44,7 @@ class CheckBatchCost: calculate_batch_cost_and_usage, ) from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, get_batch_id_from_unified_batch_id, @@ -55,6 +58,8 @@ class CheckBatchCost: } ) + completed_jobs = [] + for job in jobs: # get the model from the job unified_object_id = job.unified_object_id @@ -81,6 +86,10 @@ class CheckBatchCost: response = await self.llm_router.aretrieve_batch( model=model_id, batch_id=batch_id, + litellm_metadata={ + "user_api_key_user_id": job.created_by or "default-user-id", + "batch_ignore_default_logging": True, + }, ) ## RETRIEVE THE BATCH JOB OUTPUT FILE @@ -129,6 +138,38 @@ class CheckBatchCost: ) ) - if response.status != "validating": - # mark for updating - pass + logging_obj = LiteLLMLogging( + model=batch_models[0], + messages=[{"role": "user", "content": ""}], + stream=False, + call_type="aretrieve_batch", + start_time=datetime.now(), + litellm_call_id=str(uuid.uuid4()), + function_id=str(uuid.uuid4()), + ) + + logging_obj.update_environment_variables( + litellm_params={ + "metadata": { + "user_api_key_user_id": job.created_by or "default-user-id", + } + }, + optional_params={}, + ) + + await logging_obj.async_success_handler( + result=response, + batch_cost=batch_cost, + batch_usage=batch_usage, + batch_models=batch_models, + ) + + # mark the job as complete + completed_jobs.append(job) + + if len(completed_jobs) > 0: + # mark the jobs as complete + await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={"id": {"in": [job.id for job in completed_jobs]}}, + data={"status": "complete"}, + ) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index a8b9ac8bfa1..94962e81b9b 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -152,6 +152,7 @@ from .initialize_dynamic_callback_params import ( from .specialty_caches.dynamic_logging_cache import DynamicLoggingCache if TYPE_CHECKING: + from litellm.llms.base_llm.files.transformation import BaseFileEndpoints from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig try: from litellm_enterprise.enterprise_callbacks.callback_controls import ( @@ -1997,18 +1998,32 @@ class Logging(LiteLLMLoggingBaseClass): return ## CALCULATE COST FOR BATCH JOBS - if ( - self.call_type == CallTypes.aretrieve_batch.value - and isinstance(result, LiteLLMBatch) - and result.status == "completed" + if self.call_type == CallTypes.aretrieve_batch.value and isinstance( + result, LiteLLMBatch ): + litellm_params = self.litellm_params or {} + litellm_metadata = litellm_params.get("litellm_metadata", {}) + if ( + litellm_metadata.get("batch_ignore_default_logging", False) is True + ): # polling job will query these frequently, don't spam db logs + return + from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, ) # check if file id is a unified file id is_base64_unified_file_id = _is_base64_encoded_unified_file_id(result.id) - if not is_base64_unified_file_id: # only run for non-unified file ids + + batch_cost = kwargs.get("batch_cost", None) + batch_usage = kwargs.get("batch_usage", None) + batch_models = kwargs.get("batch_models", None) + if all([batch_cost, batch_usage, batch_models]) is not None: + result._hidden_params["response_cost"] = batch_cost + result._hidden_params["batch_models"] = batch_models + result.usage = batch_usage + + elif not is_base64_unified_file_id: # only run for non-unified file ids response_cost, batch_usage, batch_models = ( await _handle_completed_batch( batch=result, custom_llm_provider=self.custom_llm_provider diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 21c8a6148c9..4c84ff99a83 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -18,4 +18,8 @@ model_list: model: gpt-4o api_key: os.environ/OPENAI_API_KEY_TEST_2 model_info: - id: 12345679 \ No newline at end of file + id: 12345679 + + +general_settings: + check_managed_files_batch_cost: true \ No newline at end of file diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2bbcae46abd..e2476a2e684 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -3437,6 +3437,29 @@ class ProxyStartupEvent: verbose_proxy_logger.error( "Invalid maximum_spend_logs_retention_interval value" ) + ### CHECK BATCH COST ### + if llm_router is not None: + try: + from litellm_enterprise.proxy.common_utils.check_batch_cost import ( + CheckBatchCost, + ) + + check_batch_cost_job = CheckBatchCost( + proxy_logging_obj=proxy_logging_obj, + prisma_client=prisma_client, + llm_router=llm_router, + ) + scheduler.add_job( + check_batch_cost_job.check_batch_cost, + "interval", + seconds=3600, # these can run infrequently, as batch jobs take time to complete + ) + + except Exception: + verbose_proxy_logger.debug( + "Checking batch cost for LiteLLM Managed Files is an Enterprise Feature. Skipping..." + ) + pass scheduler.start()