Add cost tacking and usage info in call_type=aretrieve_batch

This commit is contained in:
Sameer Kankute 2026-01-29 15:27:41 +05:30
parent 4b385e5b32
commit fce26352b6
3 changed files with 365 additions and 4 deletions

View file

@ -192,6 +192,9 @@ async def _get_batch_output_file_content_as_dictionary(
Get the batch output file content as a list of dictionaries
"""
from litellm.files.main import afile_content
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
if custom_llm_provider == "vertex_ai":
raise ValueError("Vertex AI does not support file content retrieval")
@ -199,8 +202,17 @@ async def _get_batch_output_file_content_as_dictionary(
if batch.output_file_id is None:
raise ValueError("Output file id is None cannot retrieve file content")
file_id = batch.output_file_id
is_base64_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
if is_base64_unified_file_id:
try:
file_id = is_base64_unified_file_id.split("llm_output_file_id,")[1].split(";")[0]
verbose_logger.debug(f"Extracted LLM output file ID from unified file ID: {file_id}")
except (IndexError, AttributeError) as e:
verbose_logger.error(f"Failed to extract LLM output file ID from unified file ID: {batch.output_file_id}, error: {e}")
_file_content = await afile_content(
file_id=batch.output_file_id,
file_id=file_id,
custom_llm_provider=custom_llm_provider,
)
return _get_file_content_as_dictionary(_file_content.content)

View file

@ -2336,18 +2336,28 @@ class Logging(LiteLLMLoggingBaseClass):
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:
has_explicit_batch_data = all(
x is not None for x in (batch_cost, batch_usage, batch_models)
)
should_compute_batch_data = (
not is_base64_unified_file_id
or not has_explicit_batch_data
and result.status == "completed"
)
if has_explicit_batch_data:
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
elif should_compute_batch_data:
(
response_cost,
batch_usage,
batch_models,
) = await _handle_completed_batch(
batch=result, custom_llm_provider=self.custom_llm_provider
batch=result,
custom_llm_provider=self.custom_llm_provider,
)
result._hidden_params["response_cost"] = response_cost

View file

@ -169,3 +169,342 @@ def test_get_response_from_batch_job_output_file(sample_file_content_dict):
assert result["id"] == "chatcmpl-AhjSMl7oZ79yIPHLRYgmgXSixTJr7"
assert result["object"] == "chat.completion"
assert result["usage"]["total_tokens"] == 30
@pytest.mark.asyncio
async def test_batch_retrieve_cost_tracking_with_completed_batch_no_explicit_cost():
"""
Test that cost is calculated for completed batches when no explicit cost data is provided.
Regression test for: When batch status is "completed" and explicit batch_cost/batch_usage/batch_models
are not provided, the system should compute batch data by calling _handle_completed_batch.
"""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.utils import CallTypes
from litellm.types.utils import LiteLLMBatch
from unittest.mock import AsyncMock, patch
# Mock batch result with completed status
mock_batch = LiteLLMBatch(
id="batch-test-123",
object="batch",
endpoint="/v1/chat/completions",
errors=None,
input_file_id="file-input-123",
completion_window="24h",
status="completed",
output_file_id="file-output-123",
error_file_id=None,
created_at=1234567890,
in_progress_at=1234567900,
expires_at=1234654290,
finalizing_at=1234568000,
completed_at=1234568100,
failed_at=None,
expired_at=None,
cancelling_at=None,
cancelled_at=None,
request_counts={
"total": 10,
"completed": 10,
"failed": 0,
},
metadata=None,
)
mock_batch._hidden_params = {}
# Create logging object
logging_obj = Logging(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "test"}],
stream=False,
call_type=CallTypes.aretrieve_batch.value,
litellm_call_id="test-call-123",
function_id="test-function",
start_time=time.time(),
dynamic_success_callbacks=[],
)
logging_obj.custom_llm_provider = "openai"
# Mock _handle_completed_batch to return cost data
expected_cost = 0.05
expected_usage = litellm.Usage(
prompt_tokens=100,
completion_tokens=50,
total_tokens=150,
)
expected_models = ["gpt-4o-mini"]
with patch(
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch",
new=AsyncMock(return_value=(expected_cost, expected_usage, expected_models))
) as mock_handle_batch:
# Call async_success_handler
await logging_obj.async_success_handler(
result=mock_batch,
start_time=time.time(),
end_time=time.time() + 1,
)
# Verify _handle_completed_batch was called
mock_handle_batch.assert_called_once()
# Verify cost and usage were set on the batch result
assert mock_batch._hidden_params["response_cost"] == expected_cost
assert mock_batch._hidden_params["batch_models"] == expected_models
assert mock_batch.usage == expected_usage
@pytest.mark.asyncio
async def test_batch_retrieve_cost_tracking_with_explicit_cost_data():
"""
Test that explicit cost data is used when provided, skipping computation.
Regression test for: When batch_cost, batch_usage, and batch_models are explicitly
provided in kwargs, they should be used directly without calling _handle_completed_batch.
"""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.utils import CallTypes
from litellm.types.utils import LiteLLMBatch
from unittest.mock import AsyncMock, patch
# Mock batch result with completed status
mock_batch = LiteLLMBatch(
id="batch-test-456",
object="batch",
endpoint="/v1/chat/completions",
errors=None,
input_file_id="file-input-456",
completion_window="24h",
status="completed",
output_file_id="file-output-456",
error_file_id=None,
created_at=1234567890,
in_progress_at=1234567900,
expires_at=1234654290,
finalizing_at=1234568000,
completed_at=1234568100,
failed_at=None,
expired_at=None,
cancelling_at=None,
cancelled_at=None,
request_counts={
"total": 5,
"completed": 5,
"failed": 0,
},
metadata=None,
)
mock_batch._hidden_params = {}
# Create logging object
logging_obj = Logging(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "test"}],
stream=False,
call_type=CallTypes.aretrieve_batch.value,
litellm_call_id="test-call-456",
function_id="test-function",
start_time=time.time(),
dynamic_success_callbacks=[],
)
logging_obj.custom_llm_provider = "openai"
# Explicit cost data to pass in kwargs
explicit_cost = 0.10
explicit_usage = litellm.Usage(
prompt_tokens=200,
completion_tokens=100,
total_tokens=300,
)
explicit_models = ["gpt-4o-mini", "gpt-3.5-turbo"]
with patch(
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch",
new=AsyncMock()
) as mock_handle_batch:
# Call async_success_handler with explicit cost data
await logging_obj.async_success_handler(
result=mock_batch,
start_time=time.time(),
end_time=time.time() + 1,
batch_cost=explicit_cost,
batch_usage=explicit_usage,
batch_models=explicit_models,
)
# Verify _handle_completed_batch was NOT called (since explicit data provided)
mock_handle_batch.assert_not_called()
# Verify explicit cost data was used
assert mock_batch._hidden_params["response_cost"] == explicit_cost
assert mock_batch._hidden_params["batch_models"] == explicit_models
assert mock_batch.usage == explicit_usage
@pytest.mark.asyncio
async def test_batch_retrieve_cost_tracking_with_unified_file_id_incomplete_batch():
"""
Test that cost computation is skipped for unified file IDs with non-completed batches.
Regression test for: For unified file IDs (base64 encoded), cost should only be computed
when batch status is "completed" and explicit data is not provided.
"""
import base64
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.utils import CallTypes, SpecialEnums
from litellm.types.utils import LiteLLMBatch
from unittest.mock import AsyncMock, patch
# Create a proper unified file ID by encoding the correct prefix
unified_id_str = f"{SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value}:test_file_789;unified_id:batch-789"
encoded_unified_id = base64.urlsafe_b64encode(unified_id_str.encode()).decode().rstrip("=")
# Mock batch result with in_progress status and unified file ID
mock_batch = LiteLLMBatch(
id=encoded_unified_id, # Properly encoded unified ID
object="batch",
endpoint="/v1/chat/completions",
errors=None,
input_file_id="file-input-789",
completion_window="24h",
status="in_progress", # Not completed
output_file_id=None,
error_file_id=None,
created_at=1234567890,
in_progress_at=1234567900,
expires_at=1234654290,
finalizing_at=None,
completed_at=None,
failed_at=None,
expired_at=None,
cancelling_at=None,
cancelled_at=None,
request_counts={
"total": 10,
"completed": 3,
"failed": 0,
},
metadata=None,
)
mock_batch._hidden_params = {}
# Create logging object
logging_obj = Logging(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "test"}],
stream=False,
call_type=CallTypes.aretrieve_batch.value,
litellm_call_id="test-call-789",
function_id="test-function",
start_time=time.time(),
dynamic_success_callbacks=[],
)
logging_obj.custom_llm_provider = "openai"
with patch(
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch",
new=AsyncMock()
) as mock_handle_batch:
# Call async_success_handler with in_progress batch (unified file ID)
await logging_obj.async_success_handler(
result=mock_batch,
start_time=time.time(),
end_time=time.time() + 1,
)
# Verify _handle_completed_batch was NOT called (batch not completed and is unified file ID)
mock_handle_batch.assert_not_called()
# Verify cost data was not set
assert "response_cost" not in mock_batch._hidden_params
assert "batch_models" not in mock_batch._hidden_params
assert not hasattr(mock_batch, "usage") or mock_batch.usage is None
@pytest.mark.asyncio
async def test_batch_retrieve_cost_tracking_with_partial_explicit_data():
"""
Test that cost is computed when only partial explicit data is provided.
Regression test for: If batch_cost, batch_usage, or batch_models is missing
(not all three provided), and batch is completed, system should compute the data.
"""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.types.utils import CallTypes
from litellm.types.utils import LiteLLMBatch
from unittest.mock import AsyncMock, patch
# Mock batch result with completed status
mock_batch = LiteLLMBatch(
id="batch-test-partial",
object="batch",
endpoint="/v1/chat/completions",
errors=None,
input_file_id="file-input-partial",
completion_window="24h",
status="completed",
output_file_id="file-output-partial",
error_file_id=None,
created_at=1234567890,
in_progress_at=1234567900,
expires_at=1234654290,
finalizing_at=1234568000,
completed_at=1234568100,
failed_at=None,
expired_at=None,
cancelling_at=None,
cancelled_at=None,
request_counts={
"total": 8,
"completed": 8,
"failed": 0,
},
metadata=None,
)
mock_batch._hidden_params = {}
# Create logging object
logging_obj = Logging(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "test"}],
stream=False,
call_type=CallTypes.aretrieve_batch.value,
litellm_call_id="test-call-partial",
function_id="test-function",
start_time=time.time(),
dynamic_success_callbacks=[],
)
logging_obj.custom_llm_provider = "openai"
# Only provide batch_cost, missing batch_usage and batch_models
partial_cost = 0.08
expected_cost = 0.06
expected_usage = litellm.Usage(
prompt_tokens=150,
completion_tokens=75,
total_tokens=225,
)
expected_models = ["gpt-4o-mini"]
with patch(
"litellm.litellm_core_utils.litellm_logging._handle_completed_batch",
new=AsyncMock(return_value=(expected_cost, expected_usage, expected_models))
) as mock_handle_batch:
# Call async_success_handler with partial explicit data
await logging_obj.async_success_handler(
result=mock_batch,
start_time=time.time(),
end_time=time.time() + 1,
batch_cost=partial_cost, # Only cost provided, not usage or models
)
# Verify _handle_completed_batch WAS called (since not all data provided)
mock_handle_batch.assert_called_once()
# Verify computed cost data was used (not partial explicit data)
assert mock_batch._hidden_params["response_cost"] == expected_cost
assert mock_batch._hidden_params["batch_models"] == expected_models
assert mock_batch.usage == expected_usage