diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 935b96a0e39..166ef7a66d0 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -523,6 +523,10 @@ async def retrieve_batch( # noqa: PLR0915 custom_llm_provider=custom_llm_provider, **data # type: ignore ) + response = await proxy_logging_obj.post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + # FIX: Update the database with the latest state from provider await update_batch_in_database( batch_id=batch_id, @@ -533,19 +537,9 @@ async def retrieve_batch( # noqa: PLR0915 verbose_proxy_logger=verbose_proxy_logger, db_batch_object=db_batch_object, operation="retrieve", + user_api_key_dict=user_api_key_dict, ) - ### CALL HOOKS ### - modify outgoing data - response = await proxy_logging_obj.post_call_success_hook( - data=data, user_api_key_dict=user_api_key_dict, response=response - ) - - # Fix: bug_feb14_batch_retrieve_returns_raw_input_file_id - # Resolve raw provider file IDs (input, output, error) to unified IDs. - if unified_batch_id: - await resolve_input_file_id_to_unified(response, prisma_client) - await resolve_output_file_ids_to_unified(response, prisma_client) - ### ALERTING ### asyncio.create_task( proxy_logging_obj.update_request_status( @@ -917,10 +911,14 @@ async def cancel_batch( **_cancel_batch_data, ) - # FIX: Update the database with the new cancelled state managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") from litellm.proxy.proxy_server import prisma_client + response = await proxy_logging_obj.post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + + # FIX: Update the database with the new cancelled state await update_batch_in_database( batch_id=batch_id, unified_batch_id=unified_batch_id, @@ -929,11 +927,7 @@ async def cancel_batch( prisma_client=prisma_client, verbose_proxy_logger=verbose_proxy_logger, operation="cancel", - ) - - ### CALL HOOKS ### - modify outgoing data - response = await proxy_logging_obj.post_call_success_hook( - data=data, user_api_key_dict=user_api_key_dict, response=response + user_api_key_dict=user_api_key_dict, ) ### ALERTING ### diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 30c78ed5ba7..0415bb456ec 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -727,6 +727,76 @@ async def resolve_output_file_ids_to_unified(response, prisma_client) -> None: pass +async def ensure_batch_response_managed_file_ids( + response, + managed_files_obj, + prisma_client, + verbose_proxy_logger, + user_api_key_dict=None, + db_batch_object=None, +) -> None: + """Normalize batch file IDs to managed unified IDs before DB persistence.""" + await resolve_input_file_id_to_unified(response, prisma_client) + await resolve_output_file_ids_to_unified(response, prisma_client) + + if managed_files_obj is None: + return + + hidden_params = getattr(response, "_hidden_params", None) or {} + model_id = hidden_params.get("model_id") + if not model_id: + return + + model_name = hidden_params.get("model_name") + unified_file_id = hidden_params.get("unified_file_id") + if not model_name and isinstance(unified_file_id, str): + decoded_unified_file_id = ( + _is_base64_encoded_unified_file_id(unified_file_id) or unified_file_id + ) + target_model_names = get_models_from_unified_file_id(decoded_unified_file_id) + if target_model_names: + model_name = ",".join(target_model_names) + + if user_api_key_dict is None and db_batch_object is not None: + from litellm.proxy._types import UserAPIKeyAuth + + user_api_key_dict = UserAPIKeyAuth( + user_id=getattr(db_batch_object, "created_by", None) or "default-user-id", + team_id=getattr(db_batch_object, "team_id", None), + ) + if user_api_key_dict is None: + return + + for file_attr in ("output_file_id", "error_file_id"): + raw_file_id = getattr(response, file_attr, None) + if not raw_file_id or _is_base64_encoded_unified_file_id(raw_file_id): + continue + try: + new_unified_file_id = managed_files_obj.get_unified_output_file_id( + output_file_id=raw_file_id, + model_id=model_id, + model_name=model_name, + ) + await managed_files_obj.store_unified_file_id( + file_id=new_unified_file_id, + file_object=None, + litellm_parent_otel_span=getattr( + user_api_key_dict, "parent_otel_span", None + ), + model_mappings={model_id: raw_file_id}, + user_api_key_dict=user_api_key_dict, + ) + setattr(response, file_attr, new_unified_file_id) + verbose_proxy_logger.debug( + f"Converted batch {file_attr} {raw_file_id!r} to managed ID before DB write" + ) + except Exception as e: + verbose_proxy_logger.warning( + f"Failed to convert batch {file_attr}={raw_file_id!r} to managed ID " + f"before DB write: {e}" + ) + + async def get_batch_from_database( batch_id: str, unified_batch_id: Union[str, Literal[False]], @@ -800,6 +870,7 @@ async def update_batch_in_database( verbose_proxy_logger, db_batch_object=None, operation: str = "update", + user_api_key_dict=None, ): """ Update batch status and object in ManagedObjectTable. @@ -813,6 +884,7 @@ async def update_batch_in_database( verbose_proxy_logger: Logger instance db_batch_object: Optional existing database object (for comparison) operation: Description of operation ("update", "cancel", etc.) + user_api_key_dict: Optional auth context for creating managed file IDs """ import litellm.utils @@ -823,6 +895,18 @@ async def update_batch_in_database( if not prisma_client: return + # Always normalize the response's file IDs to unified managed IDs + # (mutates in place) so the caller returns unified IDs to the user + # even when we skip the DB update below for an unchanged status. + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=managed_files_obj, + prisma_client=prisma_client, + verbose_proxy_logger=verbose_proxy_logger, + user_api_key_dict=user_api_key_dict, + db_batch_object=db_batch_object, + ) + # Only update if status has changed (when db_batch_object is provided) if db_batch_object and response.status == db_batch_object.status: return diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index 06637b844b1..acb79a7577d 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -31,10 +31,8 @@ import litellm # the cassette state the branch is being tested with. from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap -_local_cost_map = GetModelCostMap.load_local_model_cost_map() -for _k, _v in _local_cost_map.items(): +for _k, _v in GetModelCostMap.load_local_model_cost_map().items(): litellm.model_cost.setdefault(_k, _v) -del _local_cost_map from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, diff --git a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py new file mode 100644 index 00000000000..d8669960674 --- /dev/null +++ b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py @@ -0,0 +1,260 @@ +"""Regression: update_batch_in_database must not persist raw provider output_file_id.""" + +import json +from types import SimpleNamespace +from typing import Optional +import pytest +from unittest.mock import AsyncMock, MagicMock + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.openai_files_endpoints.common_utils import ( + ensure_batch_response_managed_file_ids, + update_batch_in_database, +) +from litellm.types.utils import LiteLLMBatch + + +def _build_batch_response( + *, + batch_id: str = "batch_managed_ids_test", + status: str = "completed", + output_file_id: Optional[str] = "file-rawoutput789", + error_file_id: Optional[str] = None, + hidden_params: Optional[dict] = None, +) -> LiteLLMBatch: + batch = LiteLLMBatch( + id=batch_id, + object="batch", + status=status, + endpoint="/v1/chat/completions", + input_file_id="file-input123", + output_file_id=output_file_id, + error_file_id=error_file_id, + completion_window="24h", + created_at=1234567890, + ) + if hidden_params is not None: + batch._hidden_params = hidden_params # type: ignore[attr-defined] + return batch + + +def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ="): + mock = MagicMock() + mock.get_unified_output_file_id = MagicMock(return_value=unified_id) + mock.store_unified_file_id = AsyncMock() + return mock + + +def _build_prisma_mock(): + mock = MagicMock() + mock.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) + mock.db.litellm_managedobjecttable.update = AsyncMock() + return mock + + +@pytest.mark.asyncio +async def test_update_batch_in_database_stores_unified_output_file_id(): + raw_output_file_id = "file-rawoutput789" + unified_output_file_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ=" + batch_id = "batch_managed_ids_test" + unified_batch_id = ( + "litellm_proxy;model_id:my-model;llm_batch_id:batch_managed_ids_test" + ) + + response = _build_batch_response( + batch_id=batch_id, + output_file_id=raw_output_file_id, + hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"}, + ) + + mock_managed_files = _build_managed_files_mock(unified_id=unified_output_file_id) + mock_prisma = _build_prisma_mock() + + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + response=response, + managed_files_obj=mock_managed_files, + prisma_client=mock_prisma, + verbose_proxy_logger=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"), + ) + + stored = json.loads( + mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs["data"][ + "file_object" + ] + ) + assert stored["output_file_id"] == unified_output_file_id + assert stored["output_file_id"] != raw_output_file_id + + +@pytest.mark.asyncio +async def test_ensure_batch_response_normalizes_error_file_id(): + """Both output_file_id and error_file_id must be normalized to managed IDs.""" + unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ=" + response = _build_batch_response( + output_file_id="file-raw-output", + error_file_id="file-raw-error", + hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"}, + ) + + mock_managed_files = _build_managed_files_mock(unified_id=unified_id) + mock_prisma = _build_prisma_mock() + + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=mock_managed_files, + prisma_client=mock_prisma, + verbose_proxy_logger=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"), + ) + + assert response.output_file_id == unified_id + assert response.error_file_id == unified_id + assert mock_managed_files.get_unified_output_file_id.call_count == 2 + + +@pytest.mark.asyncio +async def test_ensure_batch_response_swallows_conversion_errors(): + """When the managed-files conversion raises, the failure is logged, not propagated.""" + raw_output_file_id = "file-raw-output" + response = _build_batch_response( + output_file_id=raw_output_file_id, + hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"}, + ) + + mock_managed_files = MagicMock() + mock_managed_files.get_unified_output_file_id = MagicMock( + side_effect=RuntimeError("boom") + ) + mock_managed_files.store_unified_file_id = AsyncMock() + + mock_logger = MagicMock() + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=mock_managed_files, + prisma_client=_build_prisma_mock(), + verbose_proxy_logger=mock_logger, + user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"), + ) + + assert response.output_file_id == raw_output_file_id + mock_logger.warning.assert_called() + + +@pytest.mark.asyncio +async def test_ensure_batch_response_builds_auth_from_db_batch_object(): + """If user_api_key_dict is omitted, fall back to created_by/team_id on db_batch_object.""" + unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ=" + response = _build_batch_response( + output_file_id="file-raw-output", + hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"}, + ) + + mock_managed_files = _build_managed_files_mock(unified_id=unified_id) + db_batch_object = SimpleNamespace( + created_by="user-from-db", team_id="team-from-db", status="completed" + ) + + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=mock_managed_files, + prisma_client=_build_prisma_mock(), + verbose_proxy_logger=MagicMock(), + db_batch_object=db_batch_object, + ) + + forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[ + "user_api_key_dict" + ] + assert forwarded_auth.user_id == "user-from-db" + assert forwarded_auth.team_id == "team-from-db" + + +@pytest.mark.asyncio +async def test_ensure_batch_response_resolves_model_name_from_unified_file_id(): + """When hidden_params lacks model_name, derive it from unified_file_id.""" + unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ=" + response = _build_batch_response( + output_file_id="file-raw-output", + hidden_params={ + "model_id": "my-model", + "unified_file_id": "litellm_proxy:application/octet-stream;unified_id,abc;target_model_names,gpt-4o-mini,gemini-2.0-flash", + }, + ) + + mock_managed_files = _build_managed_files_mock(unified_id=unified_id) + + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=mock_managed_files, + prisma_client=_build_prisma_mock(), + verbose_proxy_logger=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"), + ) + + assert ( + mock_managed_files.get_unified_output_file_id.call_args.kwargs["model_name"] + == "gpt-4o-mini,gemini-2.0-flash" + ) + + +@pytest.mark.asyncio +async def test_ensure_batch_response_returns_early_without_managed_files_obj(): + """Without managed_files_obj, the helper is a no-op (no conversion attempted).""" + response = _build_batch_response( + output_file_id="file-raw-output", + hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"}, + ) + + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=None, + prisma_client=_build_prisma_mock(), + verbose_proxy_logger=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"), + ) + + assert response.output_file_id == "file-raw-output" + + +@pytest.mark.asyncio +async def test_ensure_batch_response_returns_early_without_model_id(): + """Without model_id in hidden_params, the helper cannot create managed IDs.""" + response = _build_batch_response( + output_file_id="file-raw-output", + hidden_params={"model_name": "openai/gpt-4o"}, + ) + mock_managed_files = _build_managed_files_mock() + + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=mock_managed_files, + prisma_client=_build_prisma_mock(), + verbose_proxy_logger=MagicMock(), + user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"), + ) + + assert response.output_file_id == "file-raw-output" + mock_managed_files.get_unified_output_file_id.assert_not_called() + + +@pytest.mark.asyncio +async def test_ensure_batch_response_returns_early_without_auth(): + """Without user_api_key_dict or db_batch_object, no conversion is attempted.""" + response = _build_batch_response( + output_file_id="file-raw-output", + hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"}, + ) + mock_managed_files = _build_managed_files_mock() + + await ensure_batch_response_managed_file_ids( + response=response, + managed_files_obj=mock_managed_files, + prisma_client=_build_prisma_mock(), + verbose_proxy_logger=MagicMock(), + ) + + assert response.output_file_id == "file-raw-output" + mock_managed_files.get_unified_output_file_id.assert_not_called()