Litellm batch api background cost calc (#12125)

* feat(check_batch_cost.py): emit spend log on successful request

ensures cost tracked for batch requests

* feat(proxy_server.py): add background job to poll completed batch jobs

used for calculating cost for batch jobs

* fix(proxy_server.py): run batch cost tracking job every hour

batch jobs take time to complete, no need to run every few seconds

* feat(proxy_server.py): run batch cost tracking job every hour
This commit is contained in:
Krish Dholakia 2025-06-27 21:46:25 -07:00 • committed by GitHub
parent 06c86d6130
commit d895ae827d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 92 additions and 9 deletions

View file

@ -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": "<retrieve_batch>"}],
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"},
)

View file

@ -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

View file

@ -18,4 +18,8 @@ model_list:
model: gpt-4o
api_key: os.environ/OPENAI_API_KEY_TEST_2
model_info:
id: 12345679
id: 12345679
general_settings:
check_managed_files_batch_cost: true

View file

@ -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()