litellm/tests/proxy_unit_tests/test_check_batch_cost.py

2579 lines
107 KiB
Python

"""
Unit tests for CheckBatchCost class.
Covers: stale-row cleanup (file_purpose scoping), paginated find_many,
the batch_processed-column fallback query, and routing of unmanaged
Vertex (raw gs:// input_file_id) and Bedrock (raw s3:// input_file_id,
ARN unified_object_id) batches with no managed unified id.
"""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
_IS_B64 = "litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id"
def _unmanaged_vertex_file_object(
input_file_id="gs://bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-flash/abc.jsonl",
status="validating",
):
"""A LiteLLMBatch JSON blob shaped like what the managed-files hook stores for an
unmanaged Vertex batch (raw gs:// input_file_id)."""
from litellm.types.utils import LiteLLMBatch
return LiteLLMBatch(
id="8823717160934178816",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id=input_file_id,
object="batch",
status=status,
).model_dump_json()
def _unmanaged_bedrock_file_object(
input_file_id=(
"s3://bucket/litellm-bedrock-files-us.anthropic.claude-sonnet-4-20250514-v1-0"
"-74b61828-9191-4d80-addb-5a0f9ab0ec6a.jsonl"
),
status="validating",
):
"""A LiteLLMBatch JSON blob shaped like what gets stored for an unmanaged Bedrock
batch (raw s3:// input_file_id, ARN unified_object_id)."""
from litellm.types.utils import LiteLLMBatch
return LiteLLMBatch(
id="arn:aws:bedrock:us-east-1:298249409318:model-invocation-job/1ofb47x17jua",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id=input_file_id,
object="batch",
status=status,
).model_dump_json()
class TestCheckBatchCost:
"""Test suite for CheckBatchCost class"""
@pytest.fixture
def mock_prisma_client(self):
client = MagicMock()
client.db = MagicMock()
client.db.litellm_managedobjecttable = MagicMock()
client.db.litellm_usertable = MagicMock()
return client
@pytest.fixture
def mock_proxy_logging_obj(self):
mock = MagicMock()
mock.get_proxy_hook.return_value = None
return mock
@pytest.fixture
def mock_llm_router(self):
return MagicMock()
@pytest.fixture
def check_batch_cost_instance(
self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router
):
from litellm_enterprise.proxy.common_utils.check_batch_cost import (
CheckBatchCost,
)
return CheckBatchCost(
proxy_logging_obj=mock_proxy_logging_obj,
prisma_client=mock_prisma_client,
llm_router=mock_llm_router,
)
@pytest.mark.asyncio
async def test_cleanup_scoped_to_batch_file_purpose(
self, check_batch_cost_instance, mock_prisma_client
):
"""_cleanup_stale_managed_objects scopes its update to file_purpose='batch' only."""
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
# Return empty so the main poll loop exits immediately
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[]
)
await check_batch_cost_instance.check_batch_cost()
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
stale_call = calls[0]
assert stale_call[1]["data"] == {"status": "stale_expired"}
where = stale_call[1]["where"]
assert where["file_purpose"] == "batch"
assert "stale_expired" in where["status"]["not_in"]
assert "created_at" in where
@pytest.mark.asyncio
async def test_startup_probe_confirms_batch_processed_support(
self, check_batch_cost_instance, mock_prisma_client
):
mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=None)
await check_batch_cost_instance.confirm_batch_processed_support()
probe_where = mock_prisma_client.db.litellm_managedobjecttable.find_first.call_args[1]["where"]
assert probe_where["batch_processed"] is False
assert check_batch_cost_instance.batch_processed_support_confirmed is True
assert check_batch_cost_instance._has_batch_processed_column is True
@pytest.mark.asyncio
async def test_startup_probe_marks_column_absent(
self, check_batch_cost_instance, mock_prisma_client
):
mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
side_effect=Exception("column batch_processed does not exist")
)
await check_batch_cost_instance.confirm_batch_processed_support()
assert check_batch_cost_instance.batch_processed_support_confirmed is False
assert check_batch_cost_instance._has_batch_processed_column is False
@pytest.mark.asyncio
async def test_startup_probe_transient_error_defers_to_poll_cycle(
self, check_batch_cost_instance, mock_prisma_client
):
mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
side_effect=Exception("connection reset by peer")
)
await check_batch_cost_instance.confirm_batch_processed_support()
assert check_batch_cost_instance.batch_processed_support_confirmed is False
assert check_batch_cost_instance._has_batch_processed_column is True
@pytest.mark.asyncio
async def test_find_many_uses_pagination_and_excludes_stale(
self, check_batch_cost_instance, mock_prisma_client
):
"""find_many is called with take, order, and all terminal statuses excluded."""
from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[]
)
await check_batch_cost_instance.check_batch_cost()
find_call = mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args
assert find_call[1]["take"] == MAX_OBJECTS_PER_POLL_CYCLE
assert find_call[1]["order"] == {"created_at": "asc"}
not_in = find_call[1]["where"]["status"]["not_in"]
assert "stale_expired" in not_in
# "complete"/"completed" are intentionally NOT excluded from the
# primary query — the batch_processed=False filter is sufficient.
# This allows CheckBatchCost to pick up batches that were
# transitioned to "complete" by the retrieve_batch endpoint
# before CheckBatchCost had a chance to process them.
assert "complete" not in not_in
assert "completed" not in not_in
assert find_call[1]["where"]["batch_processed"] is False
assert check_batch_cost_instance.batch_processed_support_confirmed is True
@pytest.mark.asyncio
async def test_fallback_query_used_when_batch_processed_missing(
self, check_batch_cost_instance, mock_prisma_client
):
"""Falls back to query without batch_processed when primary query raises."""
from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
# First find_many (primary query) raises with a schema error; second (fallback) returns empty
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
side_effect=[Exception("column batch_processed does not exist"), []]
)
await check_batch_cost_instance.check_batch_cost()
calls = (
mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args_list
)
assert len(calls) == 2
fallback_where = calls[1][1]["where"]
assert "batch_processed" not in fallback_where
assert "stale_expired" in fallback_where["status"]["not_in"]
assert calls[1][1]["take"] == MAX_OBJECTS_PER_POLL_CYCLE
# Column absence is now cached — next call should go straight to fallback
assert check_batch_cost_instance._has_batch_processed_column is False
assert check_batch_cost_instance.batch_processed_support_confirmed is False
@pytest.mark.asyncio
async def test_column_absence_cached_across_cycles(
self, check_batch_cost_instance, mock_prisma_client
):
"""After column absence is discovered, subsequent cycles skip the primary query entirely."""
from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
# Simulate column already known absent from a previous cycle
check_batch_cost_instance._has_batch_processed_column = False
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[]
)
await check_batch_cost_instance.check_batch_cost()
# Only one find_many call — the fallback directly, no primary query attempt
assert (
mock_prisma_client.db.litellm_managedobjecttable.find_many.call_count == 1
)
fallback_where = (
mock_prisma_client.db.litellm_managedobjecttable.find_many.call_args[1][
"where"
]
)
assert "batch_processed" not in fallback_where
@pytest.mark.asyncio
async def test_fallback_completion_update_omits_batch_processed(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""When batch_processed column is absent, completion update must not include it.
If it did, the update would fail silently, the job would never be marked done,
and every subsequent poll cycle would re-log the cost (duplicate billing).
"""
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-fallback-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
# Simulate column already known absent (e.g. discovered on a previous cycle)
check_batch_cost_instance._has_batch_processed_column = False
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
# Build a fake batch response whose status triggers the completion branch
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "file-output-123"
mock_response.model_dump_json.return_value = (
'{"id":"batch-1","status":"completed"}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "openai"
mock_deployment.litellm_params.model = "gpt-4"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1"}'
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(
0.01,
{"prompt_tokens": 10, "completion_tokens": 5},
["gpt-4"],
),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("gpt-4", "openai", None, None),
),
patch(
"litellm.litellm_core_utils.litellm_logging.Logging"
) as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await check_batch_cost_instance.check_batch_cost()
# The update must have been called — this is the core assertion.
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
), "Expected update() to be called exactly once for the completed job"
update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[
1
]["data"]
assert (
"batch_processed" not in update_data
), "update() must NOT include batch_processed when column is absent"
assert update_data["status"] == "complete"
@pytest.mark.asyncio
async def test_output_fetch_passes_deployment_credentials_as_trusted_snapshot(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""Bedrock resolves the output bucket ONLY from the immutable snapshot kwarg.
Spreading the credentials as plain kwargs is not enough: get_litellm_params drops
s3_bucket_name, so without _litellm_internal_model_credentials the cost poller
cannot read the output file and every completed Bedrock batch stays unbilled.
"""
from types import MappingProxyType
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_job = MagicMock()
mock_job.id = "job-bedrock-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "file-output-123"
mock_response.model_dump_json.return_value = '{"id":"batch-1","status":"completed"}'
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={
"custom_llm_provider": "bedrock",
"s3_bucket_name": "configured-batch-bucket",
"aws_region_name": "us-east-1",
}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "bedrock"
mock_deployment.litellm_params.model = "bedrock/anthropic.claude-haiku-4-5-20251001-v1:0"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
mock_file_content = MagicMock()
mock_file_content.content = b'{"recordId":"req-1"}'
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
) as mock_afile_content,
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"recordId": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(0.01, {"prompt_tokens": 10, "completion_tokens": 5}, ["claude-haiku-4-5"]),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("anthropic.claude-haiku-4-5-20251001-v1:0", "bedrock", None, None),
),
patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await check_batch_cost_instance.check_batch_cost()
mock_afile_content.assert_awaited()
passed_kwargs = mock_afile_content.await_args[1]
snapshot = passed_kwargs.get("_litellm_internal_model_credentials")
assert snapshot is not None, "cost poller must pass the trusted credential snapshot"
assert isinstance(
snapshot, MappingProxyType
), "snapshot must be a MappingProxyType; a plain dict is rejected by get_configured_s3_bucket_name"
assert snapshot["s3_bucket_name"] == "configured-batch-bucket"
@pytest.mark.asyncio
async def test_poller_prices_with_deployment_registered_batch_rates(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""The cost poller must price with the rates the router registered for the deployment.
The deployment's raw model_info dict carries no litellm_params pricing, so passing
its model_dump() made the poller bill custom-rate batches at the public cost-map
price while the inline retrieve path billed the declared rate.
"""
from unittest.mock import patch
import litellm
deployment_id = "deploy-poller-registered-rates-1"
litellm.model_cost[deployment_id] = {
"id": deployment_id,
"input_cost_per_token_batches": 2e-06,
"output_cost_per_token_batches": 4e-06,
"litellm_provider": "bedrock",
"mode": "chat",
}
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_job = MagicMock()
mock_job.id = "job-poller-rates-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "file-output-123"
mock_response.model_dump_json.return_value = '{"id":"batch-1","status":"completed"}'
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"custom_llm_provider": "bedrock", "aws_region_name": "us-east-1"}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "bedrock"
mock_deployment.litellm_params.model = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
mock_file_content = MagicMock()
mock_file_content.content = b'{"recordId":"req-1"}'
decoded_id = f"llm_model_id,{deployment_id};llm_batch_id,batch-456;"
try:
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value=deployment_id,
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"recordId": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(0.0052, {"prompt_tokens": 1400, "completion_tokens": 600}, ["claude-haiku-4-5"]),
) as mock_calculate,
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("us.anthropic.claude-haiku-4-5-20251001-v1:0", "bedrock", None, None),
),
patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await check_batch_cost_instance.check_batch_cost()
finally:
litellm.model_cost.pop(deployment_id, None)
mock_calculate.assert_awaited_once()
passed_model_info = mock_calculate.await_args.kwargs["model_info"]
assert passed_model_info is not None, "poller must pass the deployment's registered pricing"
assert passed_model_info["input_cost_per_token_batches"] == 2e-06
assert passed_model_info["output_cost_per_token_batches"] == 4e-06
@pytest.mark.asyncio
async def test_primary_path_completion_update_includes_batch_processed(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""When batch_processed column IS present, completion update must set it to True.
This is the symmetric counterpart to test_fallback_completion_update_omits_batch_processed
and proves the conditional on _has_batch_processed_column governs the update data.
"""
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-primary-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "file-output-123"
mock_response.model_dump_json.return_value = (
'{"id":"batch-1","status":"completed"}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "openai"
mock_deployment.litellm_params.model = "gpt-4"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1"}'
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(
0.01,
{"prompt_tokens": 10, "completion_tokens": 5},
["gpt-4"],
),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("gpt-4", "openai", None, None),
),
patch(
"litellm.litellm_core_utils.litellm_logging.Logging"
) as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await check_batch_cost_instance.check_batch_cost()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
), "Expected update() to be called exactly once for the completed job"
update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[
1
]["data"]
assert (
update_data["batch_processed"] is True
), "update() must include batch_processed=True when column is present"
assert update_data["status"] == "complete"
@pytest.mark.asyncio
async def test_completed_batch_with_no_attributable_owner_still_writes_spend_log(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""Regression: a batch created with the master key or a team-less key has
created_by=None and team_id=None on LiteLLM_ManagedObjectTable (the table
never stores the raw key hash). CheckBatchCost's synthetic logging_obj for
such a batch then carries no attributable key/user/team/end-user, and
before the fix _should_track_cost_callback silently skipped the DB write
with no error or warning: batch_processed still became True, but no
LiteLLM_SpendLogs row was ever written.
Unlike the other tests in this file, this one does NOT mock
litellm_logging.Logging or async_success_handler -- it runs the real
logging pipeline through to _ProxyDBLogger, which is the exact gap that
let the original bug ship undetected.
"""
import litellm
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_job = MagicMock()
mock_job.id = "job-unattributed-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = None
mock_job.team_id = None
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
# A real LiteLLMBatch (not a bare MagicMock): this test runs the real
# litellm_logging.Logging pipeline, which type-checks the result via
# isinstance(..., LiteLLMBatch) before it will compute/attach a cost.
from litellm.types.utils import LiteLLMBatch
mock_response = LiteLLMBatch(
id="batch-1",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="file-input-123",
object="batch",
status="completed",
output_file_id="file-output-123",
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"})
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "openai"
mock_deployment.litellm_params.model = "gpt-4"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1"}'
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
db_logger = _ProxyDBLogger()
mock_update_database = AsyncMock()
# Unlike the other tests in this file, this one runs the real
# litellm_logging.Logging pipeline, which calls
# _is_base64_encoded_unified_file_id an extra time (checking result.id
# after it's reset to job.unified_object_id). Key off the argument
# instead of a fixed-length side_effect list so the exact call count
# doesn't matter.
def _fake_is_base64_encoded(file_id):
return decoded_id if file_id == mock_job.unified_object_id else None
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=_fake_is_base64_encoded,
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(
0.01,
{"prompt_tokens": 10, "completion_tokens": 5},
["gpt-4"],
),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("gpt-4", "openai", None, None),
),
patch.object(litellm, "_async_success_callback", [db_logger]),
patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
MagicMock(
db_spend_update_writer=MagicMock(update_database=mock_update_database),
slack_alerting_instance=MagicMock(customer_spend_alert=AsyncMock()),
),
),
patch("litellm.proxy.proxy_server.increment_spend_counters", AsyncMock()),
patch("litellm.proxy.proxy_server.update_cache", AsyncMock()),
):
await check_batch_cost_instance.check_batch_cost()
mock_update_database.assert_awaited_once()
assert mock_update_database.call_args.kwargs["response_cost"] == 0.01
assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1, (
"the job must still be marked processed once cost tracking succeeds"
)
@pytest.mark.asyncio
async def test_cost_tracking_failure_leaves_job_unprocessed(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""LIT-4008 regression: when fetching a completed batch's results fails
(e.g. Anthropic rejecting a msgbatch_ id on the Files API), the job must
NOT be marked complete/batch_processed. Pre-fix the $0 spend row was
written and batch_processed=True made it permanent; the failure must
instead leave the row untouched so the next poll retries, without
aborting the poll cycle.
"""
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-anthropic-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "msgbatch_01WA5hdsa2Xx8w4zyPjV1frs"
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test", "custom_llm_provider": "anthropic"}
)
decoded_id = "llm_model_id,model-123;llm_batch_id,msgbatch_01WA5hdsa2Xx8w4zyPjV1frs;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="msgbatch_01WA5hdsa2Xx8w4zyPjV1frs",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
side_effect=Exception("File id must have `file_` prefix."),
),
):
await check_batch_cost_instance.check_batch_cost()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0
), "a failed cost tracking attempt must not mark the job processed"
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal_status", ["failed", "expired", "cancelled"])
async def test_terminal_status_marks_job_processed(
self,
check_batch_cost_instance,
mock_prisma_client,
mock_llm_router,
terminal_status,
):
"""When the provider reports a terminal status with nothing to bill
(failed/cancelled, or expired with no output file), the row must be written back
with that status and batch_processed=True so it stops being polled forever.
"""
import base64
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-terminal-1"
mock_job.unified_object_id = base64.urlsafe_b64encode(
b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456"
).decode()
mock_job.created_by = "user-1"
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_response = MagicMock()
mock_response.status = terminal_status
mock_response.output_file_id = None
mock_response.model_dump_json.return_value = (
f'{{"id":"batch-1","status":"{terminal_status}"}}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
await check_batch_cost_instance.check_batch_cost()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
), f"Expected update() to be called exactly once for a {terminal_status} job"
update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[
1
]["data"]
assert update_data["status"] == terminal_status
assert (
update_data["batch_processed"] is True
), "terminal-status update() must set batch_processed=True so polling stops"
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal_status", ["failed", "cancelled"])
async def test_terminal_status_persists_managed_output_file_ids(
self,
check_batch_cost_instance,
mock_prisma_client,
mock_llm_router,
terminal_status,
):
"""A cancelled/failed batch with a provider error file (and no output file) must
be persisted with unified managed file IDs, never raw provider IDs. Raw IDs
written here leak to every later GET /batches/{id} and GET /batches because the
terminal row is final (batch_processed=True) and read paths only resolve, never
mint. (Any terminal status with an output file is billed through the completed
path instead, covered by test_terminal_status_with_output_file_is_billed.)
"""
import base64
import json
from litellm.types.utils import LiteLLMBatch
unified_batch_uid = base64.urlsafe_b64encode(
b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456"
).decode()
raw_error_file_id = "file-terminal-err-xyz"
raw_input_file_id = "file-terminal-in-123"
unified_input_file_id = base64.urlsafe_b64encode(
b"litellm_proxy:application/octet-stream;unified_id,in-1;target_model_names,gpt-5-batch"
).decode()
unified_error_file_id = base64.urlsafe_b64encode(
f"litellm_proxy:application/octet-stream;unified_id,u-2;llm_output_file_id,{raw_error_file_id}".encode()
).decode()
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
input_file_row = MagicMock()
input_file_row.unified_file_id = unified_input_file_id
def find_managed_file(where):
if where["flat_model_file_ids"]["has"] == raw_input_file_id:
return input_file_row
return None
mock_prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(
side_effect=find_managed_file
)
mock_job = MagicMock()
mock_job.id = "job-terminal-mint-1"
mock_job.unified_object_id = unified_batch_uid
mock_job.created_by = "user-1"
mock_job.team_id = "team-1"
check_batch_cost_instance._has_batch_processed_column = True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
response = LiteLLMBatch(
id="batch-456",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id=raw_input_file_id,
object="batch",
status=terminal_status,
output_file_id=None,
error_file_id=raw_error_file_id,
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=response)
mock_hook = MagicMock()
mock_hook.get_unified_output_file_id.side_effect = [unified_error_file_id]
mock_hook.store_unified_file_id = AsyncMock()
check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = (
mock_hook
)
await check_batch_cost_instance.check_batch_cost()
mock_hook.get_unified_output_file_id.assert_called_once_with(
output_file_id=raw_error_file_id,
model_id="model-123",
model_name="gpt-5-batch",
)
stored = {
next(iter(c.kwargs["model_mappings"].values())): c.kwargs["file_id"]
for c in mock_hook.store_unified_file_id.call_args_list
}
assert stored == {raw_error_file_id: unified_error_file_id}
for store_call in mock_hook.store_unified_file_id.call_args_list:
assert store_call.kwargs["user_api_key_dict"].user_id == "user-1"
assert store_call.kwargs["user_api_key_dict"].team_id == "team-1"
assert mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
update_call = mock_prisma_client.db.litellm_managedobjecttable.update.call_args
assert update_call.kwargs["where"] == {"id": "job-terminal-mint-1"}
update_data = update_call.kwargs["data"]
assert update_data["status"] == terminal_status
assert update_data["batch_processed"] is True
persisted = json.loads(update_data["file_object"])
assert persisted["id"] == unified_batch_uid
assert persisted["input_file_id"] == unified_input_file_id
assert persisted["output_file_id"] is None
assert persisted["error_file_id"] == unified_error_file_id
assert raw_error_file_id not in update_data["file_object"]
@pytest.mark.asyncio
@pytest.mark.parametrize("completed_status", ["completed", "complete"])
async def test_completed_without_output_file_marked_processed_without_billing(
self,
check_batch_cost_instance,
mock_prisma_client,
mock_llm_router,
completed_status,
):
"""#35354 regression: a terminal completed batch whose request lines all failed
reaches `completed` with output_file_id=None (only an error_file_id).
Pre-fix it matched neither the completed-with-output branch nor the
failed/expired/cancelled branch, so batch_processed stayed False and the row
was re-selected on every poll cycle forever. It must now be marked terminal
exactly once, without being billed (no output means nothing to bill).
"""
import base64
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-completed-no-output-1"
mock_job.unified_object_id = base64.urlsafe_b64encode(
b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456"
).decode()
mock_job.created_by = "user-1"
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_response = MagicMock()
mock_response.status = completed_status
mock_response.output_file_id = None
mock_response.error_file_id = "file-error-123"
mock_response.model_dump_json.return_value = (
f'{{"id":"batch-1","status":"{completed_status}"}}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
# Billing reads credentials off the router; if it is touched we billed a batch
# that has no output, which is the behaviour this test guards against.
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
with patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
) as mock_afile_content:
await check_batch_cost_instance.check_batch_cost()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
), "a completed batch with no output file must be marked processed exactly once"
update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[
1
]["data"]
assert update_data["status"] == completed_status
assert (
update_data["batch_processed"] is True
), "completed-without-output update() must set batch_processed=True so polling stops"
assert (
mock_afile_content.await_count == 0
), "a batch with no output file must not be billed"
assert (
mock_llm_router.get_deployment_credentials_with_provider.call_count == 0
), "a batch with no output file must not enter the cost-tracking path"
@pytest.mark.asyncio
async def test_non_terminal_status_left_unprocessed(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""A batch still validating/in_progress must NOT be treated as terminal: no DB
write, so it keeps being polled until it actually reaches a terminal status.
"""
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_job = MagicMock()
mock_job.id = "job-in-progress-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_response = MagicMock()
mock_response.status = "in_progress"
mock_response.output_file_id = None
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
):
await check_batch_cost_instance.check_batch_cost()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 0
), "a non-terminal batch must not be written back (would stop polling prematurely)"
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal_status", ["expired", "cancelled", "failed"])
async def test_terminal_status_with_output_file_is_billed(
self,
check_batch_cost_instance,
mock_prisma_client,
mock_llm_router,
terminal_status,
):
"""A terminal (expired/cancelled/failed) batch that still produced an output file
served real request lines, so it must be billed (cost tracked) and then marked
processed, not silently marked terminal without billing.
"""
from unittest.mock import patch
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-terminal-with-output-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_response = MagicMock()
mock_response.status = terminal_status
mock_response.output_file_id = "file-output-123"
mock_response.model_dump_json.return_value = (
f'{{"id":"batch-1","status":"{terminal_status}"}}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "openai"
mock_deployment.litellm_params.model = "gpt-4"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1"}'
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
side_effect=[decoded_id, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
) as mock_afile_content,
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(
0.01,
{"prompt_tokens": 10, "completion_tokens": 5},
["gpt-4"],
),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("gpt-4", "openai", None, None),
),
patch(
"litellm.litellm_core_utils.litellm_logging.Logging"
) as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await check_batch_cost_instance.check_batch_cost()
assert (
mock_afile_content.await_count == 1
), f"{terminal_status} batch with an output file must fetch results and be billed"
mock_logging_obj.async_success_handler.assert_awaited_once()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
)
update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[
1
]["data"]
assert update_data["batch_processed"] is True
assert (
update_data["status"] == terminal_status
), f"billed {terminal_status} batch must keep its real terminal status in the DB"
@pytest.mark.asyncio
async def test_terminal_batch_with_missing_output_file_is_retired_unbilled(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""A terminal batch whose advertised output file 404s at the provider has
nothing to fetch on this or any later poll (Vertex AI advertises an output
path for every batch, even ones that never wrote it), so the job must be
retired as terminal on the first cycle instead of retrying until the
staleness sweep gives up on it.
"""
import base64
from unittest.mock import patch
from litellm.exceptions import NotFoundError
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-output-gone-1"
mock_job.unified_object_id = base64.urlsafe_b64encode(
b"litellm_proxy;model_id:model-123;llm_batch_id:batch-456"
).decode()
mock_job.created_by = "user-1"
assert check_batch_cost_instance._has_batch_processed_column is True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
missing_output_file_id = "gs://batch-out/job-1/predictions.jsonl"
mock_response = MagicMock()
mock_response.status = "failed"
mock_response.output_file_id = missing_output_file_id
mock_response.error_file_id = None
mock_response.model_dump_json.return_value = (
'{"id":"batch-1","status":"failed"}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
with (
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
side_effect=NotFoundError(
message=f"404: output file {missing_output_file_id} does not exist",
model="gemini-2.5-pro",
llm_provider="vertex_ai",
),
) as mock_afile_content,
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
) as mock_calculate,
):
await check_batch_cost_instance.check_batch_cost()
assert mock_afile_content.await_count == 1
mock_calculate.assert_not_awaited()
assert (
mock_prisma_client.db.litellm_managedobjecttable.update.call_count == 1
), "a terminal batch with a 404ing output file must be retired, not retried forever"
update_data = mock_prisma_client.db.litellm_managedobjecttable.update.call_args[
1
]["data"]
assert update_data["status"] == "failed"
assert update_data["batch_processed"] is True
@pytest.mark.asyncio
async def test_raw_output_file_id_converted_to_managed_id(
self, check_batch_cost_instance, mock_prisma_client, mock_llm_router
):
"""CheckBatchCost must convert a raw provider output_file_id to a managed base64 ID.
Without this, GET /batches/{id} returns a raw file ID that cannot be routed
through the proxy, causing API_KEY errors when clients call GET /files/{id}/content.
"""
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=None
)
mock_job = MagicMock()
mock_job.id = "job-raw-file-1"
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
mock_job.created_by = "user-1"
mock_job.team_id = None
check_batch_cost_instance._has_batch_processed_column = True
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
raw_output_file_id = "file-batch-output-abc123"
raw_error_file_id = "file-batch-error-xyz456"
fake_managed_output_id = "bGl0ZWxsbV9wcm94eTo6b3V0cHV0"
fake_managed_error_id = "bGl0ZWxsbV9wcm94eTo6ZXJyb3I="
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = raw_output_file_id
mock_response.error_file_id = raw_error_file_id
mock_response.model_dump_json.return_value = (
'{"id":"batch-1","status":"completed"}'
)
mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response)
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
mock_deployment = MagicMock()
mock_deployment.litellm_params.custom_llm_provider = "azure"
mock_deployment.litellm_params.model = "azure/gpt-5-mini"
mock_deployment.model_name = "gpt-5-batch"
mock_deployment.model_info.model_dump.return_value = {}
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
mock_hook = MagicMock()
mock_hook.get_unified_output_file_id.side_effect = [
fake_managed_output_id,
fake_managed_error_id,
]
mock_hook.store_unified_file_id = AsyncMock()
check_batch_cost_instance.proxy_logging_obj.get_proxy_hook.return_value = (
mock_hook
)
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1"}'
decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;"
with (
patch(
"litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id",
# call 1: job unified_object_id decode, call 2: existing raw check for output_file_id,
# call 3: fix guard for output_file_id, call 4: fix guard for error_file_id
side_effect=[decoded_id, None, None, None],
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id",
return_value="model-123",
),
patch(
"litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id",
return_value="batch-456",
),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(
0.01,
{"prompt_tokens": 10, "completion_tokens": 5},
["gpt-4"],
),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("gpt-5-mini", "azure", None, None),
),
patch(
"litellm.litellm_core_utils.litellm_logging.Logging"
) as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await check_batch_cost_instance.check_batch_cost()
assert mock_hook.get_unified_output_file_id.call_count == 2
mock_hook.get_unified_output_file_id.assert_any_call(
output_file_id=raw_output_file_id,
model_id="model-123",
model_name="gpt-5-batch",
)
mock_hook.get_unified_output_file_id.assert_any_call(
output_file_id=raw_error_file_id,
model_id="model-123",
model_name="gpt-5-batch",
)
assert mock_hook.store_unified_file_id.await_count == 2
# {raw_file_id: managed_file_id} for each store call
stored = {
next(iter(c[1]["model_mappings"].values())): c[1]["file_id"]
for c in mock_hook.store_unified_file_id.call_args_list
}
assert stored == {
raw_output_file_id: fake_managed_output_id,
raw_error_file_id: fake_managed_error_id,
}
assert mock_response.output_file_id == fake_managed_output_id
assert mock_response.error_file_id == fake_managed_error_id
class TestUnmanagedVertexRouting:
"""Routing of unmanaged Vertex batches whose unified_object_id is a raw provider job id."""
def _instance(self, track_unmanaged, router):
from litellm_enterprise.proxy.common_utils.check_batch_cost import (
CheckBatchCost,
)
return CheckBatchCost(
proxy_logging_obj=MagicMock(),
prisma_client=MagicMock(),
llm_router=router,
track_unmanaged_batch_cost=track_unmanaged,
)
def _job(self, file_object=None):
job = MagicMock()
job.unified_object_id = "8823717160934178816"
job.file_object = (
file_object if file_object is not None else _unmanaged_vertex_file_object()
)
return job
def test_flag_off_skips_unmanaged_id_unchanged(self):
"""Default (flag off): a raw numeric unified_object_id is skipped exactly as before;
no model derivation or router lookup happens."""
router = MagicMock()
instance = self._instance(track_unmanaged=False, router=router)
prom = MagicMock()
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), prom)
assert result is None
prom.record_check_batch_cost_error.assert_called_once_with("invalid_unified_id")
router.resolve_model_name_from_model_id.assert_not_called()
router.get_model_ids.assert_not_called()
def _vertex_deployment(self):
deployment = MagicMock()
deployment.litellm_params.custom_llm_provider = "vertex_ai"
deployment.litellm_params.model = "vertex_ai/gemini-2.5-flash"
return deployment
def test_flag_on_routes_to_vertex_deployment(self):
"""Flag on: derive the bare model from the gs:// path, resolve it to a deployment id,
and use the raw unified_object_id as the provider batch id."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = "gemini-2.5-flash"
router.get_model_ids.return_value = ["deploy-1"]
router.get_deployment = MagicMock(return_value=self._vertex_deployment())
instance = self._instance(track_unmanaged=True, router=router)
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), MagicMock())
assert result == ("deploy-1", "8823717160934178816")
# bare model name (trailing GCS segment), not the full publishers/.. path
router.resolve_model_name_from_model_id.assert_called_once_with(
"gemini-2.5-flash"
)
router.get_model_ids.assert_called_once_with(model_name="gemini-2.5-flash")
def test_flag_on_skips_non_vertex_deployment_sharing_model_group(self):
"""Flag on, but the only deployment for the model group is a non-vertex_ai
provider: must not be selected, even though the model group name matches."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = "gemini-2.5-flash"
router.get_model_ids.return_value = ["deploy-openai"]
non_vertex_deployment = MagicMock()
non_vertex_deployment.litellm_params.custom_llm_provider = "openai"
non_vertex_deployment.litellm_params.model = "gpt-4o"
router.get_deployment = MagicMock(return_value=non_vertex_deployment)
instance = self._instance(track_unmanaged=True, router=router)
prom = MagicMock()
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), prom)
assert result is None
prom.record_check_batch_cost_error.assert_called_once_with(
"unmanaged_no_matching_deployment"
)
def test_flag_on_uses_later_vertex_deployment_with_matching_suffix(self):
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = "azure-gemini"
router.get_model_ids.return_value = ["deploy-azure"]
non_vertex_deployment = MagicMock()
non_vertex_deployment.litellm_params.custom_llm_provider = "azure"
non_vertex_deployment.litellm_params.model = "azure/gemini-2.5-flash"
router.get_deployment = MagicMock(return_value=non_vertex_deployment)
router.get_model_list.return_value = [
{
"model_name": "azure-gemini",
"litellm_params": {
"model": "azure/gemini-2.5-flash",
"custom_llm_provider": "azure",
},
"model_info": {"id": "deploy-azure"},
},
{
"model_name": "vertex-gemini",
"litellm_params": {
"model": "vertex_ai/gemini-2.5-flash",
"custom_llm_provider": "vertex_ai",
},
"model_info": {"id": "deploy-vertex"},
},
]
instance = self._instance(track_unmanaged=True, router=router)
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), MagicMock())
assert result == ("deploy-vertex", "8823717160934178816")
router.get_model_ids.assert_called_once_with(model_name="azure-gemini")
def test_flag_on_no_matching_deployment_records_metric(self):
"""Flag on but no vertex_ai deployment for the model: skip with a distinct metric."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = None
router.get_model_ids.return_value = []
instance = self._instance(track_unmanaged=True, router=router)
prom = MagicMock()
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), prom)
assert result is None
prom.record_check_batch_cost_error.assert_called_once_with(
"unmanaged_no_matching_deployment"
)
def test_flag_on_non_gcs_input_is_not_unmanaged_vertex(self):
"""Flag on, but input_file_id is not a gs:// publishers path: treat as unroutable,
do not attempt model derivation."""
router = MagicMock()
instance = self._instance(track_unmanaged=True, router=router)
prom = MagicMock()
job = self._job(
file_object=_unmanaged_vertex_file_object(input_file_id="file-abc-123")
)
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(job, prom)
assert result is None
prom.record_check_batch_cost_error.assert_called_once_with("invalid_unified_id")
router.resolve_model_name_from_model_id.assert_not_called()
@pytest.mark.asyncio
async def test_end_to_end_costs_unmanaged_batch(self):
"""Flag on, completed unmanaged batch: the poller polls Vertex with the raw job id,
computes cost, and marks batch_processed=True. Fails before this change (the row is
skipped at the unified-id gate)."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = "gemini-2.5-flash"
router.get_model_ids.return_value = ["deploy-1"]
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "gs://bucket/out/predictions.jsonl"
mock_response.error_file_id = None
mock_response.completed_at = None
mock_response.created_at = None
mock_response.model_dump_json.return_value = (
'{"id":"8823717160934178816","status":"completed"}'
)
router.aretrieve_batch = AsyncMock(return_value=mock_response)
router.get_deployment_credentials_with_provider = MagicMock(
return_value={"vertex_project": "p", "vertex_location": "us-central1"}
)
deployment = MagicMock()
deployment.litellm_params.custom_llm_provider = "vertex_ai"
deployment.litellm_params.model = "vertex_ai/gemini-2.5-flash"
deployment.model_name = "gemini-2.5-flash"
deployment.model_info.model_dump.return_value = {}
router.get_deployment = MagicMock(return_value=deployment)
instance = self._instance(track_unmanaged=True, router=router)
instance.proxy_logging_obj.get_proxy_hook.return_value = None
instance._has_batch_processed_column = True
prisma = instance.prisma_client
prisma.db = MagicMock()
prisma.db.litellm_managedobjecttable = MagicMock()
prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0)
prisma.db.litellm_managedobjecttable.update = AsyncMock()
prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[self._job()]
)
prisma.db.litellm_usertable = MagicMock()
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1"}'
with (
patch(_IS_B64, side_effect=[False, None]),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(
0.01,
{"prompt_tokens": 10, "completion_tokens": 5},
["gemini-2.5-flash"],
),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("gemini-2.5-flash", "vertex_ai", None, None),
),
patch(
"litellm.litellm_core_utils.litellm_logging.Logging"
) as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await instance.check_batch_cost()
router.aretrieve_batch.assert_awaited_once()
assert router.aretrieve_batch.call_args[1]["model"] == "deploy-1"
assert router.aretrieve_batch.call_args[1]["batch_id"] == "8823717160934178816"
mock_logging_obj.async_success_handler.assert_awaited_once()
assert mock_logging_obj.async_success_handler.call_args[1]["batch_cost"] == 0.01
assert prisma.db.litellm_managedobjecttable.update.call_count == 1
update_data = prisma.db.litellm_managedobjecttable.update.call_args[1]["data"]
assert update_data["batch_processed"] is True
assert update_data["status"] == "complete"
class TestUnmanagedBedrockRouting:
"""Routing of unmanaged Bedrock batches whose unified_object_id is a raw model-invocation-job ARN."""
_ARN = "arn:aws:bedrock:us-east-1:298249409318:model-invocation-job/1ofb47x17jua"
def _instance(self, track_unmanaged, router):
from litellm_enterprise.proxy.common_utils.check_batch_cost import (
CheckBatchCost,
)
return CheckBatchCost(
proxy_logging_obj=MagicMock(),
prisma_client=MagicMock(),
llm_router=router,
track_unmanaged_batch_cost=track_unmanaged,
)
def _job(self, file_object=None):
job = MagicMock()
job.unified_object_id = self._ARN
job.file_object = (
file_object if file_object is not None else _unmanaged_bedrock_file_object()
)
return job
def _bedrock_deployment(self):
deployment = MagicMock()
deployment.litellm_params.custom_llm_provider = "bedrock"
deployment.litellm_params.model = "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
return deployment
def test_flag_off_skips_arn_unified_id_unchanged(self):
"""Default (flag off): a raw ARN unified_object_id is skipped exactly as before; no
model derivation or router lookup happens."""
router = MagicMock()
instance = self._instance(track_unmanaged=False, router=router)
prom = MagicMock()
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), prom)
assert result is None
prom.record_check_batch_cost_error.assert_called_once_with("invalid_unified_id")
router.resolve_model_name_from_model_id.assert_not_called()
router.get_model_ids.assert_not_called()
def test_flag_on_routes_to_bedrock_deployment(self):
"""Flag on: derive the bare model from the s3:// object key (":" restored to "-" is
matched fuzzily), resolve it to a deployment id, and use the raw ARN as the batch id."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = "claude-sonnet-4"
router.get_model_ids.return_value = ["deploy-1"]
router.get_deployment = MagicMock(return_value=self._bedrock_deployment())
instance = self._instance(track_unmanaged=True, router=router)
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), MagicMock())
assert result == ("deploy-1", self._ARN)
def test_flag_on_skips_non_bedrock_deployment_sharing_model_group(self):
"""Flag on, but the only deployment for the model group is a non-bedrock provider:
must not be selected, even though the model group name matches."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = "claude-sonnet-4"
router.get_model_ids.return_value = ["deploy-anthropic"]
non_bedrock_deployment = MagicMock()
non_bedrock_deployment.litellm_params.custom_llm_provider = "anthropic"
non_bedrock_deployment.litellm_params.model = "claude-sonnet-4-20250514"
router.get_deployment = MagicMock(return_value=non_bedrock_deployment)
instance = self._instance(track_unmanaged=True, router=router)
prom = MagicMock()
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), prom)
assert result is None
prom.record_check_batch_cost_error.assert_called_once_with(
"unmanaged_no_matching_deployment"
)
def test_flag_on_matches_deployment_despite_colon_dash_mismatch(self):
"""The S3 object key has ':' replaced with '-' (e.g. 'v1-0'), but the configured
deployment's actual bedrock model id uses ':' (e.g. 'v1:0'). Routing must still match."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = None
router.get_model_ids.return_value = []
router.get_model_list.return_value = [
{
"model_name": "claude-sonnet-4",
"litellm_params": {
"model": "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0",
"custom_llm_provider": "bedrock",
},
"model_info": {"id": "deploy-bedrock"},
}
]
instance = self._instance(track_unmanaged=True, router=router)
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), MagicMock())
assert result == ("deploy-bedrock", self._ARN)
def test_flag_on_no_matching_deployment_records_metric(self):
"""Flag on but no bedrock deployment for the model: skip with a distinct metric."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = None
router.get_model_ids.return_value = []
router.get_model_list.return_value = []
instance = self._instance(track_unmanaged=True, router=router)
prom = MagicMock()
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(self._job(), prom)
assert result is None
prom.record_check_batch_cost_error.assert_called_once_with(
"unmanaged_no_matching_deployment"
)
def test_flag_on_non_s3_input_is_not_unmanaged_bedrock(self):
"""Flag on, but input_file_id is not a litellm-bedrock-files- s3:// key: treat as
unroutable, do not attempt model derivation."""
router = MagicMock()
instance = self._instance(track_unmanaged=True, router=router)
prom = MagicMock()
job = self._job(
file_object=_unmanaged_bedrock_file_object(input_file_id="file-abc-123")
)
with patch(_IS_B64, return_value=False):
result = instance._resolve_job_routing(job, prom)
assert result is None
prom.record_check_batch_cost_error.assert_called_once_with("invalid_unified_id")
router.resolve_model_name_from_model_id.assert_not_called()
@pytest.mark.asyncio
async def test_end_to_end_costs_unmanaged_batch(self):
"""Flag on, completed unmanaged batch: the poller polls Bedrock with the raw ARN,
computes cost, and marks batch_processed=True."""
router = MagicMock()
router.resolve_model_name_from_model_id.return_value = "claude-sonnet-4"
router.get_model_ids.return_value = ["deploy-1"]
mock_response = MagicMock()
mock_response.status = "completed"
mock_response.output_file_id = "s3://bucket/out/predictions.jsonl"
mock_response.error_file_id = None
mock_response.completed_at = None
mock_response.created_at = None
mock_response.model_dump_json.return_value = (
f'{{"id":"{self._ARN}","status":"completed"}}'
)
router.aretrieve_batch = AsyncMock(return_value=mock_response)
router.get_deployment_credentials_with_provider = MagicMock(
return_value={"aws_region_name": "us-east-1"}
)
deployment = self._bedrock_deployment()
deployment.model_name = "claude-sonnet-4"
deployment.model_info.model_dump.return_value = {}
router.get_deployment = MagicMock(return_value=deployment)
instance = self._instance(track_unmanaged=True, router=router)
instance.proxy_logging_obj.get_proxy_hook.return_value = None
instance._has_batch_processed_column = True
prisma = instance.prisma_client
prisma.db = MagicMock()
prisma.db.litellm_managedobjecttable = MagicMock()
prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0)
prisma.db.litellm_managedobjecttable.update = AsyncMock()
prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[self._job()]
)
prisma.db.litellm_usertable = MagicMock()
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
mock_file_content = MagicMock()
mock_file_content.content = b'{"id":"req-1"}'
with (
patch(_IS_B64, side_effect=[False, None]),
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=mock_file_content,
),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(
0.02,
{"prompt_tokens": 10, "completion_tokens": 5},
["claude-sonnet-4"],
),
),
patch(
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
return_value=("claude-sonnet-4", "bedrock", None, None),
),
patch(
"litellm.litellm_core_utils.litellm_logging.Logging"
) as mock_logging_cls,
):
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_cls.return_value = mock_logging_obj
await instance.check_batch_cost()
router.aretrieve_batch.assert_awaited_once()
assert router.aretrieve_batch.call_args[1]["model"] == "deploy-1"
assert router.aretrieve_batch.call_args[1]["batch_id"] == self._ARN
mock_logging_obj.async_success_handler.assert_awaited_once()
assert mock_logging_obj.async_success_handler.call_args[1]["batch_cost"] == 0.02
assert prisma.db.litellm_managedobjecttable.update.call_count == 1
update_data = prisma.db.litellm_managedobjecttable.update.call_args[1]["data"]
assert update_data["batch_processed"] is True
assert update_data["status"] == "complete"
class TestUnmanagedBatchCostFlagIsGeneralized:
"""The single track_unmanaged_batch_cost flag must cover both Vertex and Bedrock, not
just the provider it was originally added for."""
def test_one_flag_routes_both_vertex_and_bedrock_jobs(self):
from litellm_enterprise.proxy.common_utils.check_batch_cost import (
CheckBatchCost,
)
router = MagicMock()
router.resolve_model_name_from_model_id.side_effect = [
"gemini-2.5-flash",
"claude-sonnet-4",
]
router.get_model_ids.side_effect = [["deploy-vertex"], ["deploy-bedrock"]]
def _get_deployment(model_id):
if model_id == "deploy-vertex":
deployment = MagicMock()
deployment.litellm_params.custom_llm_provider = "vertex_ai"
deployment.litellm_params.model = "vertex_ai/gemini-2.5-flash"
return deployment
deployment = MagicMock()
deployment.litellm_params.custom_llm_provider = "bedrock"
deployment.litellm_params.model = "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
return deployment
router.get_deployment = MagicMock(side_effect=_get_deployment)
instance = CheckBatchCost(
proxy_logging_obj=MagicMock(),
prisma_client=MagicMock(),
llm_router=router,
track_unmanaged_batch_cost=True,
)
vertex_job = MagicMock()
vertex_job.unified_object_id = "8823717160934178816"
vertex_job.file_object = _unmanaged_vertex_file_object()
bedrock_job = MagicMock()
bedrock_job.unified_object_id = TestUnmanagedBedrockRouting._ARN
bedrock_job.file_object = _unmanaged_bedrock_file_object()
with patch(_IS_B64, return_value=False):
vertex_result = instance._resolve_job_routing(vertex_job, MagicMock())
bedrock_result = instance._resolve_job_routing(bedrock_job, MagicMock())
assert vertex_result == ("deploy-vertex", "8823717160934178816")
assert bedrock_result == ("deploy-bedrock", TestUnmanagedBedrockRouting._ARN)
class TestManagedOutputFileIdEncodesPublicModelGroup:
"""LIT-4964 regression: the unified output file id created by the background poller must
encode the public model group as ``target_model_names``, not the provider model.
Key model-access checks resolve a managed file id back to a model via ``target_model_names``,
so encoding the provider model (e.g. ``gpt-5.5``) makes
``GET /v1/files/{output_file_id}/content`` fail for every key.
"""
_PUBLIC_MODEL_GROUP = "gpt-5-batch"
_RAW_OUTPUT_FILE_ID = "file-batch-output-abc123"
@staticmethod
def _managed_input_file_id(model_group: str) -> str:
import base64
unified_id = (
"litellm_proxy:application/octet-stream;unified_id,c4843482-b176-4901-8292-7523fd0f2c6e;"
f"target_model_names,{model_group}"
)
return base64.urlsafe_b64encode(unified_id.encode()).decode().rstrip("=")
def _job(self, input_file_id: str) -> MagicMock:
from litellm.types.utils import LiteLLMBatch
job = MagicMock()
job.id = "job-lit-4964"
job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
job.created_by = "user-1"
job.team_id = None
job.file_object = LiteLLMBatch(
id="batch-456",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id=input_file_id,
object="batch",
status="completed",
).model_dump_json()
return job
async def _run(self, job: MagicMock) -> str:
from litellm_enterprise.proxy.common_utils.check_batch_cost import (
CheckBatchCost,
)
from litellm.types.utils import LiteLLMBatch
from enterprise.litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles,
)
router = MagicMock()
router.get_deployment_credentials_with_provider = MagicMock(
return_value={"api_key": "sk-test"}
)
deployment = MagicMock()
deployment.litellm_params.custom_llm_provider = "azure"
deployment.litellm_params.model = "azure/gpt-5.5"
deployment.model_name = self._PUBLIC_MODEL_GROUP
deployment.model_info.model_dump.return_value = {}
router.get_deployment = MagicMock(return_value=deployment)
hook = MagicMock()
hook.get_unified_output_file_id = (
lambda output_file_id, model_id, model_name: _PROXY_LiteLLMManagedFiles.get_unified_output_file_id(
None, output_file_id=output_file_id, model_id=model_id, model_name=model_name
)
)
hook.store_unified_file_id = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.get_proxy_hook.return_value = hook
prisma_client = MagicMock()
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
instance = CheckBatchCost(
proxy_logging_obj=proxy_logging_obj,
prisma_client=prisma_client,
llm_router=router,
)
response = LiteLLMBatch(
id="batch-456",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id=job.file_object,
object="batch",
status="completed",
)
response.output_file_id = self._RAW_OUTPUT_FILE_ID
file_content = MagicMock()
file_content.content = b'{"id":"req-1"}'
with (
patch(
"litellm.files.main.afile_content",
new_callable=AsyncMock,
return_value=file_content,
),
patch(
"litellm.batches.batch_utils._get_file_content_as_dictionary",
return_value=[{"id": "req-1"}],
),
patch(
"litellm.batches.batch_utils.calculate_batch_cost_and_usage",
new_callable=AsyncMock,
return_value=(0.01, {"prompt_tokens": 10}, ["gpt-5.5"]),
),
patch("litellm.litellm_core_utils.litellm_logging.Logging") as logging_cls,
):
logging_obj = MagicMock()
logging_obj.async_success_handler = AsyncMock()
logging_cls.return_value = logging_obj
await instance._track_completed_batch_cost(
job=job,
response=response,
model_id="model-123",
batch_id="batch-456",
prom_logger=None,
)
return response.output_file_id
@pytest.mark.asyncio
async def test_target_model_names_comes_from_input_file_not_provider_model(self):
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
get_models_from_unified_file_id,
)
output_file_id = await self._run(
self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP))
)
decoded = _is_base64_encoded_unified_file_id(output_file_id)
assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP]
assert "gpt-5.5" not in decoded
@pytest.mark.asyncio
async def test_key_scoped_to_model_group_can_read_the_output_file(self):
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_model
from litellm.proxy.auth.auth_utils import (
_extract_models_from_managed_resource_id,
)
output_file_id = await self._run(
self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP))
)
models = _extract_models_from_managed_resource_id(output_file_id, "file_id", None)
assert models == [self._PUBLIC_MODEL_GROUP]
assert (
await can_key_call_model(
model=models[0],
llm_model_list=None,
valid_token=UserAPIKeyAuth(
api_key="sk-test", models=[self._PUBLIC_MODEL_GROUP]
),
llm_router=None,
)
is True
)
@pytest.mark.asyncio
async def test_falls_back_to_deployment_model_group_without_managed_input_file(self):
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
get_models_from_unified_file_id,
)
output_file_id = await self._run(self._job("file-raw-provider-input"))
decoded = _is_base64_encoded_unified_file_id(output_file_id)
assert get_models_from_unified_file_id(decoded) == [self._PUBLIC_MODEL_GROUP]
class TestBatchCostAttribution:
"""CheckBatchCost rebuilds the creator's spend metadata from the managed-object row so
the batch-cost log is attributed like a non-batch request."""
def _instance(self, key_row=None, team_row=None, user_row=None):
from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost
prisma = MagicMock()
prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=key_row)
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row)
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
return CheckBatchCost(
proxy_logging_obj=MagicMock(),
prisma_client=prisma,
llm_router=MagicMock(),
)
def _job(self, **overrides):
from types import SimpleNamespace
fields = {
"created_by": "alice",
"team_id": "team-alpha",
"api_key": "hash-alice",
"request_tags": ["env:prod"],
}
fields.update(overrides)
return SimpleNamespace(unified_object_id="uoi", **fields)
@pytest.mark.asyncio
async def test_metadata_carries_key_team_and_tags(self):
"""The spend row names the creating key, its team, both aliases, and the tags."""
from types import SimpleNamespace
instance = self._instance(
key_row=SimpleNamespace(key_alias="prod-key"),
team_row=SimpleNamespace(team_alias="Team Alpha"),
user_row=SimpleNamespace(user_email="alice@example.com", user_alias=None),
)
metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1")
assert metadata["user_api_key"] == "hash-alice"
assert metadata["user_api_key_user_id"] == "alice"
assert metadata["user_api_key_team_id"] == "team-alpha"
assert metadata["user_api_key_alias"] == "prod-key"
assert metadata["user_api_key_team_alias"] == "Team Alpha"
assert metadata["tags"] == ["env:prod"]
@pytest.mark.asyncio
async def test_metadata_tolerates_legacy_row_without_columns(self):
"""Rows created before the columns existed carry only created_by/team_id and must
still produce an attributed row rather than raising."""
instance = self._instance()
job = self._job(api_key=None, request_tags=None)
metadata = await instance._build_creator_attribution_metadata(job, "batch-1")
assert metadata["user_api_key"] is None
assert metadata["user_api_key_user_id"] == "alice"
assert metadata["user_api_key_team_id"] == "team-alpha"
assert "tags" not in metadata
@pytest.mark.asyncio
async def test_metadata_keeps_key_when_team_key_has_no_user(self):
"""A team-scoped key carries no user id. The user lookup is skipped (prisma rejects
a None user_id) and the key hash still drives key-level attribution."""
from types import SimpleNamespace
instance = self._instance(key_row=SimpleNamespace(key_alias="svc-key"))
job = self._job(created_by=None)
metadata = await instance._build_creator_attribution_metadata(job, "batch-1")
assert metadata["user_api_key"] == "hash-alice"
assert metadata["user_api_key_user_id"] is None
assert metadata["user_api_key_alias"] == "svc-key"
instance.prisma_client.db.litellm_usertable.find_unique.assert_not_called()
@pytest.mark.asyncio
async def test_metadata_drops_non_string_tags(self):
"""Non-string tags are dropped so a malformed stored value cannot slip past the
tag-budget checks that consume this metadata."""
instance = self._instance()
job = self._job(request_tags=["env:prod", 7, None, "team:ml"])
metadata = await instance._build_creator_attribution_metadata(job, "batch-1")
assert metadata["tags"] == ["env:prod", "team:ml"]
@pytest.mark.asyncio
async def test_key_alias_lookup_failure_does_not_break_attribution(self):
"""An alias lookup failure must not lose the spend row; the key hash and team still
attribute it."""
instance = self._instance()
instance.prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
side_effect=Exception("db down")
)
metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1")
assert metadata["user_api_key"] == "hash-alice"
assert metadata.get("user_api_key_alias") is None
@pytest.mark.asyncio
async def test_unnamed_key_keeps_the_creating_user_alias(self):
"""Regression: a key generated without key_alias resolves to no alias, and the
overwrite must not null out the creating user's alias that _get_user_info supplied.
Most keys carry no alias, so this is the common batch, not an edge case."""
from types import SimpleNamespace
instance = self._instance(
key_row=SimpleNamespace(key_alias=None),
user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"),
)
metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1")
assert metadata["user_api_key_alias"] == "Alice Chen"
assert metadata["user_api_key"] == "hash-alice"
@pytest.mark.asyncio
async def test_rotated_key_keeps_the_creating_user_alias(self):
"""Batches outlive keys. When the creating key has been rotated or deleted the
lookup returns no row, and the spend log keeps a resolvable name instead of null."""
from types import SimpleNamespace
instance = self._instance(
key_row=None,
user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"),
)
metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1")
assert metadata["user_api_key_alias"] == "Alice Chen"
@pytest.mark.asyncio
async def test_named_key_still_owns_the_alias(self):
"""The fallback must not weaken the intended precedence: a key that has its own
alias still overrides the creating user's."""
from types import SimpleNamespace
instance = self._instance(
key_row=SimpleNamespace(key_alias="prod-key"),
user_row=SimpleNamespace(user_email="alice@example.com", user_alias="Alice Chen"),
)
metadata = await instance._build_creator_attribution_metadata(self._job(), "batch-1")
assert metadata["user_api_key_alias"] == "prod-key"
class TestPollPageStarvation:
"""LIT-5462 regression: a row that can never be costed used to keep its slot in the
MAX_OBJECTS_PER_POLL_CYCLE page forever, so once enough of them accumulated no newer
batch was ever polled or costed."""
def _instance(self, prisma, llm_router):
from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost
proxy_logging_obj = MagicMock()
proxy_logging_obj.get_proxy_hook.return_value = None
return CheckBatchCost(
proxy_logging_obj=proxy_logging_obj,
prisma_client=prisma,
llm_router=llm_router,
)
def _prisma(self, jobs):
prisma = MagicMock()
prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0)
prisma.db.litellm_managedobjecttable.update = AsyncMock()
prisma.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=jobs)
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
return prisma
def _job(self, job_id, unified_object_id):
job = MagicMock()
job.id = job_id
job.unified_object_id = unified_object_id
job.created_by = "user-1"
return job
@staticmethod
def _encode(unified_id: str) -> str:
import base64
return base64.urlsafe_b64encode(unified_id.encode()).decode().rstrip("=")
@pytest.mark.asyncio
async def test_unified_id_without_model_id_is_retired(self):
"""A unified id that decodes but carries no model_id is unroutable no matter what
the config says, so it must leave the poll page instead of being retried forever."""
prisma = self._prisma(
[self._job("job-no-model", self._encode("litellm_proxy;llm_batch_id:poison-no-model"))]
)
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock()
await self._instance(prisma, llm_router).check_batch_cost()
llm_router.aretrieve_batch.assert_not_awaited()
prisma.db.litellm_managedobjecttable.update.assert_awaited_once()
call = prisma.db.litellm_managedobjecttable.update.call_args[1]
assert call["where"] == {"id": "job-no-model"}
assert call["data"] == {"batch_processed": True}
@pytest.mark.asyncio
async def test_provider_404_retires_job(self):
"""The provider dropping its record of the batch is permanent: no later retrieve
can succeed, so the row must stop occupying a slot."""
import litellm
prisma = self._prisma(
[
self._job(
"job-gone",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_deadbeef"),
)
]
)
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock(
side_effect=litellm.NotFoundError(
message="No batch found with id 'batch_deadbeef'.",
model="model-123",
llm_provider="openai",
)
)
await self._instance(prisma, llm_router).check_batch_cost()
prisma.db.litellm_managedobjecttable.update.assert_awaited_once()
assert prisma.db.litellm_managedobjecttable.update.call_args[1]["data"] == {
"batch_processed": True
}
@pytest.mark.asyncio
async def test_provider_404_with_deployment_gone_keeps_job(self):
"""With the batch's own deployment removed from the router, default fallbacks can
send the retrieve to a provider that never saw the batch. That 404 proves nothing,
so the row must stay unprocessed instead of losing its spend forever."""
import litellm
prisma = self._prisma(
[
self._job(
"job-misrouted",
self._encode("litellm_proxy;model_id:model-gone;llm_batch_id:batch_alive"),
)
]
)
llm_router = MagicMock()
llm_router.get_deployment = MagicMock(return_value=None)
llm_router.aretrieve_batch = AsyncMock(
side_effect=litellm.NotFoundError(
message="No batch found with id 'batch_alive'.",
model="model-gone",
llm_provider="openai",
)
)
await self._instance(prisma, llm_router).check_batch_cost()
prisma.db.litellm_managedobjecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_transient_provider_error_keeps_job_for_retry(self):
"""A failure that may clear up (timeout, 5xx) must still leave the row unprocessed."""
prisma = self._prisma(
[
self._job(
"job-flaky",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_flaky"),
)
]
)
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock(side_effect=Exception("connection reset"))
await self._instance(prisma, llm_router).check_batch_cost()
prisma.db.litellm_managedobjecttable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_retirement_falls_back_to_status_without_batch_processed_column(self):
"""Older schemas have no batch_processed column, so the only way to stop selecting
the row is the status filter the poll query already applies."""
prisma = self._prisma(
[self._job("job-legacy", self._encode("litellm_proxy;llm_batch_id:poison-no-model"))]
)
instance = self._instance(prisma, MagicMock())
instance._has_batch_processed_column = False
await instance.check_batch_cost()
assert prisma.db.litellm_managedobjecttable.update.call_args[1]["data"] == {
"status": "stale_expired"
}
@pytest.mark.asyncio
async def test_stale_cleanup_gives_up_on_never_costed_completed_rows(self):
"""A row already in a terminal status is never rewritten by the staleness sweep, so
it needs its own bound or it starves newer batches indefinitely."""
prisma = self._prisma([])
await self._instance(prisma, MagicMock()).check_batch_cost()
calls = prisma.db.litellm_managedobjecttable.update_many.call_args_list
assert len(calls) == 2, "expected the staleness sweep plus the never-costed sweep"
where = calls[1][1]["where"]
assert where["file_purpose"] == "batch"
assert where["batch_processed"] is False
assert where["status"] == {"in": ["complete", "completed"]}
assert "created_at" in where
assert calls[1][1]["data"] == {"batch_processed": True}
@pytest.mark.asyncio
async def test_newer_batch_is_polled_once_dead_rows_are_retired(self):
"""The end state the customer cares about: dead rows retire on the cycle they are
first seen, and the healthy batch behind them keeps getting polled."""
dead_rows = [
self._job("job-no-model", self._encode("litellm_proxy;llm_batch_id:poison-no-model")),
self._job(
"job-gone",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_deadbeef"),
),
]
live_row = self._job(
"job-live",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_live"),
)
prisma = self._prisma(dead_rows + [live_row])
import litellm
in_progress = MagicMock()
in_progress.status = "in_progress"
async def _retrieve(model, batch_id, litellm_metadata):
if batch_id == "batch_deadbeef":
raise litellm.NotFoundError(
message=f"No batch found with id '{batch_id}'.",
model=model,
llm_provider="openai",
)
return in_progress
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock(side_effect=_retrieve)
await self._instance(prisma, llm_router).check_batch_cost()
retired = [
call[1]["where"]["id"]
for call in prisma.db.litellm_managedobjecttable.update.call_args_list
]
assert retired == ["job-no-model", "job-gone"]
assert (
llm_router.aretrieve_batch.await_args_list[-1][1]["batch_id"] == "batch_live"
), "the newer healthy batch must still be polled in the same cycle"
@pytest.mark.asyncio
async def test_404_that_does_not_name_the_batch_keeps_job_for_retry(self):
"""A 404 about something other than the batch, e.g. a renamed Azure deployment, is
fixable in config, so the row must survive to be costed after the fix."""
import litellm
prisma = self._prisma(
[
self._job(
"job-bad-deployment",
self._encode("litellm_proxy;model_id:model-123;llm_batch_id:batch_real"),
)
]
)
llm_router = MagicMock()
llm_router.aretrieve_batch = AsyncMock(
side_effect=litellm.NotFoundError(
message="Error code: 404 - DeploymentNotFound",
model="model-123",
llm_provider="azure",
)
)
await self._instance(prisma, llm_router).check_batch_cost()
prisma.db.litellm_managedobjecttable.update.assert_not_awaited()