mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-05 08:07:05 +00:00
test: fake the provider file boundary in the poller error-file regression test
This commit is contained in:
parent
4ef012627d
commit
99884f0eaa
1 changed files with 122 additions and 216 deletions
|
|
@ -1382,37 +1382,31 @@ class TestCheckBatchCost:
|
|||
the error file's lines, batch_failed_requests on the spend log undercounts:
|
||||
regression test for the poller path merging error-file failures.
|
||||
"""
|
||||
import base64
|
||||
from unittest.mock import patch
|
||||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
|
||||
return_value=1
|
||||
)
|
||||
import httpx
|
||||
import respx
|
||||
|
||||
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_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
mock_job = MagicMock()
|
||||
mock_job.id = "job-error-file-1"
|
||||
mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA=="
|
||||
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"
|
||||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[mock_job]
|
||||
)
|
||||
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.error_file_id = "file-error-456"
|
||||
mock_response.model_dump_json.return_value = (
|
||||
'{"id":"batch-1","status":"completed"}'
|
||||
)
|
||||
|
||||
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_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"
|
||||
|
|
@ -1420,83 +1414,73 @@ class TestCheckBatchCost:
|
|||
mock_deployment.model_info.model_dump.return_value = {}
|
||||
mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment)
|
||||
|
||||
output_file_content = MagicMock()
|
||||
output_file_content.content = b'{"id":"req-1"}'
|
||||
error_file_content = MagicMock()
|
||||
error_file_content.content = (
|
||||
b'{"id":"err-1","error":{"message":"rejected"}}\n'
|
||||
b'{"id":"err-2","error":{"message":"rejected"}}\n\n'
|
||||
succeeded_line = json.dumps(
|
||||
{
|
||||
"custom_id": "req-1",
|
||||
"response": {
|
||||
"status_code": 200,
|
||||
"body": {
|
||||
"id": "chatcmpl-1",
|
||||
"object": "chat.completion",
|
||||
"model": "gpt-4",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "hi"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 5,
|
||||
"total_tokens": 15,
|
||||
},
|
||||
},
|
||||
},
|
||||
"error": None,
|
||||
}
|
||||
)
|
||||
rejected_line = json.dumps(
|
||||
{
|
||||
"custom_id": "req-2",
|
||||
"response": {
|
||||
"status_code": 400,
|
||||
"body": {"error": {"message": "bad request"}},
|
||||
},
|
||||
"error": None,
|
||||
}
|
||||
)
|
||||
error_file_lines = "\n".join(
|
||||
json.dumps({"custom_id": custom_id, "error": {"message": "rejected"}}) for custom_id in ("req-3", "req-4")
|
||||
)
|
||||
|
||||
def _file_content_for(**kwargs):
|
||||
if kwargs["file_id"] == "file-error-456":
|
||||
return error_file_content
|
||||
return output_file_content
|
||||
|
||||
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, 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,
|
||||
side_effect=_file_content_for,
|
||||
) 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=_batch_cost_result(
|
||||
0.01,
|
||||
{"prompt_tokens": 10, "completion_tokens": 5},
|
||||
["gpt-4"],
|
||||
successful_requests=3,
|
||||
failed_requests=1,
|
||||
),
|
||||
),
|
||||
patch(
|
||||
"litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider",
|
||||
return_value=("gpt-4", "openai", None, None),
|
||||
),
|
||||
patch(
|
||||
respx.mock(assert_all_called=True) as provider,
|
||||
patch( # test-quality-ok: the poller builds Logging inline, the only seam to its handler kwargs
|
||||
"litellm.litellm_core_utils.litellm_logging.Logging"
|
||||
) as mock_logging_cls,
|
||||
):
|
||||
provider.get("https://api.openai.com/v1/files/file-output-123/content").mock(
|
||||
return_value=httpx.Response(200, content=f"{succeeded_line}\n{rejected_line}\n".encode())
|
||||
)
|
||||
provider.get("https://api.openai.com/v1/files/file-error-456/content").mock(
|
||||
return_value=httpx.Response(200, content=f"{error_file_lines}\n\n".encode())
|
||||
)
|
||||
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 == 2, (
|
||||
"the poller must fetch the error file in addition to the output file"
|
||||
)
|
||||
fetched_file_ids = {
|
||||
call.kwargs["file_id"] for call in mock_afile_content.await_args_list
|
||||
}
|
||||
assert fetched_file_ids == {"file-output-123", "file-error-456"}
|
||||
|
||||
mock_logging_obj.async_success_handler.assert_awaited_once()
|
||||
handler_kwargs = mock_logging_obj.async_success_handler.await_args.kwargs
|
||||
assert handler_kwargs["batch_successful_requests"] == 3
|
||||
assert handler_kwargs["batch_successful_requests"] == 1
|
||||
assert handler_kwargs["batch_failed_requests"] == 3, (
|
||||
"2 error-file lines must add to the output file's 1 failed request"
|
||||
"2 error-file lines must add to the output file's 1 rejected request"
|
||||
)
|
||||
assert handler_kwargs["batch_cost"] == 0.01
|
||||
assert handler_kwargs["batch_models"] == ["gpt-4"]
|
||||
assert handler_kwargs["batch_usage"].total_tokens == 15
|
||||
assert handler_kwargs["batch_cost"] > 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_terminal_batch_with_missing_output_file_is_retired_unbilled(
|
||||
|
|
@ -1513,13 +1497,9 @@ class TestCheckBatchCost:
|
|||
|
||||
from litellm.exceptions import NotFoundError
|
||||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
|
||||
return_value=1
|
||||
)
|
||||
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_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
mock_job = MagicMock()
|
||||
mock_job.id = "job-output-gone-1"
|
||||
|
|
@ -1529,23 +1509,17 @@ class TestCheckBatchCost:
|
|||
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_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_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"}
|
||||
)
|
||||
mock_llm_router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"})
|
||||
|
||||
with (
|
||||
patch(
|
||||
|
|
@ -1566,12 +1540,10 @@ class TestCheckBatchCost:
|
|||
|
||||
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 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
|
||||
|
||||
|
|
@ -1584,13 +1556,9 @@ class TestCheckBatchCost:
|
|||
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=1
|
||||
)
|
||||
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_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
mock_job = MagicMock()
|
||||
mock_job.id = "job-raw-file-1"
|
||||
|
|
@ -1599,9 +1567,7 @@ class TestCheckBatchCost:
|
|||
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]
|
||||
)
|
||||
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"
|
||||
|
|
@ -1612,14 +1578,10 @@ class TestCheckBatchCost:
|
|||
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_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_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"
|
||||
|
|
@ -1634,9 +1596,7 @@ class TestCheckBatchCost:
|
|||
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
|
||||
)
|
||||
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"}'
|
||||
|
|
@ -1679,9 +1639,7 @@ class TestCheckBatchCost:
|
|||
"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,
|
||||
patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls,
|
||||
):
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
|
|
@ -1732,9 +1690,7 @@ class TestUnmanagedVertexRouting:
|
|||
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()
|
||||
)
|
||||
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):
|
||||
|
|
@ -1772,9 +1728,7 @@ class TestUnmanagedVertexRouting:
|
|||
|
||||
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.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):
|
||||
|
|
@ -1794,9 +1748,7 @@ class TestUnmanagedVertexRouting:
|
|||
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"
|
||||
)
|
||||
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()
|
||||
|
|
@ -1844,9 +1796,7 @@ class TestUnmanagedVertexRouting:
|
|||
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"
|
||||
)
|
||||
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,
|
||||
|
|
@ -1854,9 +1804,7 @@ class TestUnmanagedVertexRouting:
|
|||
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")
|
||||
)
|
||||
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)
|
||||
|
|
@ -1880,9 +1828,7 @@ class TestUnmanagedVertexRouting:
|
|||
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"}'
|
||||
)
|
||||
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"}
|
||||
|
|
@ -1904,9 +1850,7 @@ class TestUnmanagedVertexRouting:
|
|||
prisma.db.litellm_managedobjecttable = MagicMock()
|
||||
prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
|
||||
prisma.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[self._job()]
|
||||
)
|
||||
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)
|
||||
|
||||
|
|
@ -1937,9 +1881,7 @@ class TestUnmanagedVertexRouting:
|
|||
"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,
|
||||
patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls,
|
||||
):
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
|
|
@ -1980,9 +1922,7 @@ class TestUnmanagedBedrockRouting:
|
|||
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()
|
||||
)
|
||||
job.file_object = file_object if file_object is not None else _unmanaged_bedrock_file_object()
|
||||
return job
|
||||
|
||||
def _bedrock_deployment(self):
|
||||
|
|
@ -2037,9 +1977,7 @@ class TestUnmanagedBedrockRouting:
|
|||
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"
|
||||
)
|
||||
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
|
||||
|
|
@ -2077,9 +2015,7 @@ class TestUnmanagedBedrockRouting:
|
|||
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"
|
||||
)
|
||||
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
|
||||
|
|
@ -2087,9 +2023,7 @@ class TestUnmanagedBedrockRouting:
|
|||
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")
|
||||
)
|
||||
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)
|
||||
|
|
@ -2112,13 +2046,9 @@ class TestUnmanagedBedrockRouting:
|
|||
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"}}'
|
||||
)
|
||||
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"}
|
||||
)
|
||||
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"
|
||||
|
|
@ -2134,9 +2064,7 @@ class TestUnmanagedBedrockRouting:
|
|||
prisma.db.litellm_managedobjecttable = MagicMock()
|
||||
prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
|
||||
prisma.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
prisma.db.litellm_managedobjecttable.find_many = AsyncMock(
|
||||
return_value=[self._job()]
|
||||
)
|
||||
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)
|
||||
|
||||
|
|
@ -2167,9 +2095,7 @@ class TestUnmanagedBedrockRouting:
|
|||
"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,
|
||||
patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls,
|
||||
):
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
|
|
@ -2293,9 +2219,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup:
|
|||
)
|
||||
|
||||
router = MagicMock()
|
||||
router.get_deployment_credentials_with_provider = MagicMock(
|
||||
return_value={"api_key": "sk-test"}
|
||||
)
|
||||
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"
|
||||
|
|
@ -2304,8 +2228,8 @@ class TestManagedOutputFileIdEncodesPublicModelGroup:
|
|||
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(
|
||||
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
|
||||
)
|
||||
)
|
||||
|
|
@ -2374,9 +2298,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup:
|
|||
get_models_from_unified_file_id,
|
||||
)
|
||||
|
||||
output_file_id = await self._run(
|
||||
self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP))
|
||||
)
|
||||
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]
|
||||
|
|
@ -2390,9 +2312,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup:
|
|||
_extract_models_from_managed_resource_id,
|
||||
)
|
||||
|
||||
output_file_id = await self._run(
|
||||
self._job(self._managed_input_file_id(self._PUBLIC_MODEL_GROUP))
|
||||
)
|
||||
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]
|
||||
|
|
@ -2400,9 +2320,7 @@ class TestManagedOutputFileIdEncodesPublicModelGroup:
|
|||
await can_key_call_model(
|
||||
model=models[0],
|
||||
llm_model_list=None,
|
||||
valid_token=UserAPIKeyAuth(
|
||||
api_key="sk-test", models=[self._PUBLIC_MODEL_GROUP]
|
||||
),
|
||||
valid_token=UserAPIKeyAuth(api_key="sk-test", models=[self._PUBLIC_MODEL_GROUP]),
|
||||
llm_router=None,
|
||||
)
|
||||
is True
|
||||
|
|
@ -2419,6 +2337,8 @@ class TestManagedOutputFileIdEncodesPublicModelGroup:
|
|||
|
||||
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."""
|
||||
|
|
@ -2514,9 +2434,7 @@ class TestBatchCostAttribution:
|
|||
"""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")
|
||||
)
|
||||
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")
|
||||
|
||||
|
|
@ -2612,9 +2530,7 @@ class TestPollPageStarvation:
|
|||
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"))]
|
||||
)
|
||||
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()
|
||||
|
||||
|
|
@ -2652,9 +2568,7 @@ class TestPollPageStarvation:
|
|||
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
|
||||
}
|
||||
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):
|
||||
|
|
@ -2707,17 +2621,13 @@ class TestPollPageStarvation:
|
|||
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"))]
|
||||
)
|
||||
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"
|
||||
}
|
||||
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):
|
||||
|
|
@ -2772,14 +2682,11 @@ class TestPollPageStarvation:
|
|||
|
||||
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
|
||||
]
|
||||
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"
|
||||
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):
|
||||
|
|
@ -2808,6 +2715,7 @@ class TestPollPageStarvation:
|
|||
|
||||
prisma.db.litellm_managedobjecttable.update.assert_not_awaited()
|
||||
|
||||
|
||||
class _FakeManagedObjectRow:
|
||||
"""One managed batch row the provider has finished but nothing has costed yet."""
|
||||
|
||||
|
|
@ -2824,8 +2732,12 @@ class _FakeManagedObjectRow:
|
|||
self.request_tags = None
|
||||
self.created_at = 1700000000
|
||||
self.file_object = json.dumps(
|
||||
{"id": "batch-456", "status": "in_progress", "input_file_id": "file-input-1",
|
||||
"output_file_id": _CLAIM_OUTPUT_FILE_ID}
|
||||
{
|
||||
"id": "batch-456",
|
||||
"status": "in_progress",
|
||||
"input_file_id": "file-input-1",
|
||||
"output_file_id": _CLAIM_OUTPUT_FILE_ID,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -2925,9 +2837,7 @@ class TestMultiPodBatchCostClaim:
|
|||
|
||||
router = MagicMock()
|
||||
router.aretrieve_batch = AsyncMock(return_value=response)
|
||||
router.get_deployment_credentials_with_provider = MagicMock(
|
||||
return_value={"api_key": "sk-test"}
|
||||
)
|
||||
router.get_deployment_credentials_with_provider = MagicMock(return_value={"api_key": "sk-test"})
|
||||
router.get_deployment = MagicMock(return_value=deployment)
|
||||
return router
|
||||
|
||||
|
|
@ -3087,9 +2997,7 @@ class TestMultiPodBatchCostClaim:
|
|||
await asyncio.Event().wait()
|
||||
|
||||
with self._billing_patches(journal, during_fetch=_never_returns) as logging_obj:
|
||||
interrupted = asyncio.create_task(
|
||||
self._instance(prisma, self._router()).check_batch_cost()
|
||||
)
|
||||
interrupted = asyncio.create_task(self._instance(prisma, self._router()).check_batch_cost())
|
||||
await asyncio.wait_for(reached_fetch.wait(), timeout=5)
|
||||
assert row.batch_processed is False, "an in-flight costing must not mark the row processed"
|
||||
interrupted.cancel()
|
||||
|
|
@ -3124,9 +3032,7 @@ class TestMultiPodBatchCostClaim:
|
|||
await finish_fetch.wait()
|
||||
|
||||
with self._billing_patches(journal, during_fetch=_wait_for_the_delete_attempt):
|
||||
costing = asyncio.create_task(
|
||||
self._instance(prisma, self._router()).check_batch_cost()
|
||||
)
|
||||
costing = asyncio.create_task(self._instance(prisma, self._router()).check_batch_cost())
|
||||
await asyncio.wait_for(reached_fetch.wait(), timeout=5)
|
||||
|
||||
with pytest.raises(HTTPException) as blocked:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue