mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
8d485f2403
6 changed files with 654 additions and 11 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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue