mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Add cost tacking and usage info in call_type=aretrieve_batch
This commit is contained in:
parent
4b385e5b32
commit
fce26352b6
3 changed files with 365 additions and 4 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue