Merge pull request #19986 from BerriAI/litellm_batch_cost_tracking_jan29

[Feat]Add cost tracking and usage object in aretrieve_batch call type
This commit is contained in:
Sameer Kankute 2026-01-30 17:00:42 +05:30 • committed by GitHub
commit 8d485f2403
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 654 additions and 11 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

@ -298,6 +298,103 @@ async def get_request_body(request: Request) -> Dict[str, Any]:
return {}
def extract_nested_form_metadata(
form_data: Dict[str, Any],
prefix: str = "litellm_metadata["
) -> Dict[str, Any]:
"""
Extract nested metadata from form data with bracket notation.
Handles form data that uses bracket notation to represent nested dictionaries,
such as litellm_metadata[spend_logs_metadata][owner] = "value".
This is commonly encountered when SDKs or clients send form data with nested
structures using bracket notation instead of JSON.
Args:
form_data: Dictionary containing form data (from request.form())
prefix: The prefix to look for in form keys (default: "litellm_metadata[")
Returns:
Dictionary with nested structure reconstructed from bracket notation
Example:
Input form_data:
{
"litellm_metadata[spend_logs_metadata][owner]": "john",
"litellm_metadata[spend_logs_metadata][team]": "engineering",
"litellm_metadata[tags]": "production",
"other_field": "value"
}
Output:
{
"spend_logs_metadata": {
"owner": "john",
"team": "engineering"
},
"tags": "production"
}
"""
if not form_data:
return {}
metadata: Dict[str, Any] = {}
for key, value in form_data.items():
# Skip keys that don't start with the prefix
if not isinstance(key, str) or not key.startswith(prefix):
continue
# Skip UploadFile objects - they should not be in metadata
if isinstance(value, UploadFile):
verbose_proxy_logger.warning(
f"Skipping UploadFile in metadata extraction for key: {key}"
)
continue
# Extract the nested path from bracket notation
# Example: "litellm_metadata[spend_logs_metadata][owner]" -> ["spend_logs_metadata", "owner"]
try:
# Remove the prefix and strip trailing ']'
path_string = key.replace(prefix, "").rstrip("]")
# Split by "][" to get individual path parts
parts = path_string.split("][")
if not parts or not parts[0]:
verbose_proxy_logger.warning(
f"Invalid metadata key format (empty path): {key}"
)
continue
# Navigate/create nested dictionary structure
current = metadata
for part in parts[:-1]:
if not isinstance(current, dict):
verbose_proxy_logger.warning(
f"Cannot create nested path - intermediate value is not a dict at: {part}"
)
break
current = current.setdefault(part, {})
else:
# Set the final value (only if we didn't break out of the loop)
if isinstance(current, dict):
current[parts[-1]] = value
else:
verbose_proxy_logger.warning(
f"Cannot set value - parent is not a dict for key: {key}"
)
except Exception as e:
verbose_proxy_logger.error(
f"Error parsing metadata key '{key}': {str(e)}"
)
continue
return metadata
def get_tags_from_request_body(request_body: dict) -> List[str]:
"""
Extract tags from request body metadata.

View file

@ -29,7 +29,10 @@ from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
extract_nested_form_metadata,
)
from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_body,
get_custom_llm_provider_from_request_headers,
@ -354,16 +357,20 @@ async def create_file( # noqa: PLR0915
data = {}
# Parse expires_after if provided
expires_after = None
form_data = await request.form()
litellm_metadata = extract_nested_form_metadata(
form_data=form_data,
prefix="litellm_metadata["
)
expires_after_anchor = form_data.get("expires_after[anchor]")
expires_after_seconds_str = form_data.get("expires_after[seconds]")
# Add litellm_metadata to data if provided (from form field)
if litellm_metadata is not None:
data["litellm_metadata"] = litellm_metadata
# Parse expires_after if provided
expires_after = None
form_data = await request.form()
expires_after_anchor = form_data.get("expires_after[anchor]")
expires_after_seconds_str = form_data.get("expires_after[seconds]")
if expires_after_anchor is not None or expires_after_seconds_str is not None:
if expires_after_anchor is None or expires_after_seconds_str is None:
raise HTTPException(

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

View file

@ -950,3 +950,181 @@ def test_managed_files_with_loadbalancing(mocker: MockerFixture, monkeypatch, ll
assert router_acreate_file_calls[0]["model"] == "azure-gpt-3-5-turbo"
assert router_acreate_file_calls[1]["model"] == "gpt-3.5-turbo"
assert all(call["via_router"] for call in router_acreate_file_calls), "All calls should go through router"
def test_create_file_with_nested_litellm_metadata(
mocker: MockerFixture, monkeypatch, llm_router: Router
):
"""
Test that nested litellm_metadata is correctly parsed from form data in bracket notation.
Regression test for: litellm_metadata[spend_logs_metadata][owner] format should be
correctly parsed into nested dictionary structure.
"""
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.types.llms.openai import OpenAIFileObject
proxy_logging_obj = ProxyLogging(
user_api_key_cache=DualCache(default_in_memory_ttl=1)
)
proxy_logging_obj._add_proxy_hooks(llm_router)
captured_litellm_metadata = {}
class DummyManagedFiles(BaseFileEndpoints):
async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
# Capture litellm_metadata for verification
if isinstance(create_file_request, dict):
captured_litellm_metadata.update(
create_file_request.get("litellm_metadata", {})
)
else:
captured_litellm_metadata.update(
getattr(create_file_request, "litellm_metadata", {})
)
return OpenAIFileObject(
id="file-test-123",
object="file",
bytes=100,
created_at=1234567890,
filename="test.jsonl",
purpose="fine-tune",
status="uploaded",
)
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
raise NotImplementedError("Not implemented for test")
async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
raise NotImplementedError("Not implemented for test")
proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
)
test_file_content = b'{"prompt": "Hello", "completion": "Hi"}'
test_file = ("test.jsonl", test_file_content, "application/jsonl")
# Test with nested litellm_metadata in bracket notation
response = client.post(
"/v1/files",
files={"file": test_file},
data={
"purpose": "fine-tune",
"target_model_names": "gpt-3.5-turbo",
"litellm_metadata[spend_logs_metadata][owner]": "john_doe",
"litellm_metadata[spend_logs_metadata][team]": "engineering",
"litellm_metadata[tags]": "production",
"litellm_metadata[environment]": "prod",
},
headers={"Authorization": "Bearer test-key"},
)
# Verify success
assert response.status_code == 200
result = response.json()
assert result["id"] == "file-test-123"
# Verify nested metadata was correctly parsed
assert "spend_logs_metadata" in captured_litellm_metadata
assert captured_litellm_metadata["spend_logs_metadata"]["owner"] == "john_doe"
assert captured_litellm_metadata["spend_logs_metadata"]["team"] == "engineering"
assert captured_litellm_metadata["tags"] == "production"
assert captured_litellm_metadata["environment"] == "prod"
def test_create_file_with_deep_nested_litellm_metadata(
mocker: MockerFixture, monkeypatch, llm_router: Router
):
"""
Test that deeply nested litellm_metadata is correctly parsed from form data.
Regression test for: litellm_metadata[a][b][c] format should be correctly parsed.
"""
from litellm.llms.base_llm.files.transformation import BaseFileEndpoints
from litellm.types.llms.openai import OpenAIFileObject
proxy_logging_obj = ProxyLogging(
user_api_key_cache=DualCache(default_in_memory_ttl=1)
)
proxy_logging_obj._add_proxy_hooks(llm_router)
captured_litellm_metadata = {}
class DummyManagedFiles(BaseFileEndpoints):
async def acreate_file(self, llm_router, create_file_request, target_model_names_list, litellm_parent_otel_span, user_api_key_dict):
if isinstance(create_file_request, dict):
captured_litellm_metadata.update(
create_file_request.get("litellm_metadata", {})
)
else:
captured_litellm_metadata.update(
getattr(create_file_request, "litellm_metadata", {})
)
return OpenAIFileObject(
id="file-test-456",
object="file",
bytes=50,
created_at=1234567890,
filename="nested.jsonl",
purpose="batch",
status="uploaded",
)
async def afile_retrieve(self, file_id, litellm_parent_otel_span, llm_router):
raise NotImplementedError("Not implemented for test")
async def afile_list(self, purpose, litellm_parent_otel_span):
raise NotImplementedError("Not implemented for test")
async def afile_delete(self, file_id, litellm_parent_otel_span, llm_router, **data):
raise NotImplementedError("Not implemented for test")
async def afile_content(self, file_id, litellm_parent_otel_span, llm_router, **data):
raise NotImplementedError("Not implemented for test")
proxy_logging_obj.proxy_hook_mapping["managed_files"] = DummyManagedFiles()
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
)
test_file_content = b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo"}}'
test_file = ("nested.jsonl", test_file_content, "application/jsonl")
# Test with deeply nested metadata
response = client.post(
"/v1/files",
files={"file": test_file},
data={
"purpose": "batch",
"target_model_names": "gpt-3.5-turbo",
"litellm_metadata[config][database][host]": "localhost",
"litellm_metadata[config][database][port]": "5432",
"litellm_metadata[config][cache][enabled]": "true",
},
headers={"Authorization": "Bearer test-key"},
)
# Verify success
assert response.status_code == 200
result = response.json()
assert result["id"] == "file-test-456"
# Verify deeply nested metadata was correctly parsed
assert "config" in captured_litellm_metadata
assert "database" in captured_litellm_metadata["config"]
assert captured_litellm_metadata["config"]["database"]["host"] == "localhost"
assert captured_litellm_metadata["config"]["database"]["port"] == "5432"
assert "cache" in captured_litellm_metadata["config"]
assert captured_litellm_metadata["config"]["cache"]["enabled"] == "true"