diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index ae5905f9cdf..cbd17845c5f 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1099,13 +1099,24 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): for file_attr in ["output_file_id", "error_file_id"]: file_id_value = getattr(response, file_attr, None) if file_id_value and model_id: - original_file_id = file_id_value - unified_file_id = self.get_unified_output_file_id( - output_file_id=original_file_id, - model_id=model_id, - model_name=resolved_model_name, + decoded_unified_file_id = _is_base64_encoded_unified_file_id( + file_id_value ) - setattr(response, file_attr, unified_file_id) + if decoded_unified_file_id: + provider_file_id = ( + self.get_output_file_id_from_unified_file_id( + decoded_unified_file_id + ) + ) + unified_file_id = file_id_value + else: + provider_file_id = file_id_value + unified_file_id = self.get_unified_output_file_id( + output_file_id=provider_file_id, + model_id=model_id, + model_name=resolved_model_name, + ) + setattr(response, file_attr, unified_file_id) # Use llm_router credentials when available. Without credentials, # Azure and other auth-required providers return 500/401. @@ -1125,27 +1136,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): or {} ) file_object = await litellm.afile_retrieve( - file_id=original_file_id, + file_id=provider_file_id, **_creds, ) else: file_object = await litellm.afile_retrieve( custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai", # type: ignore[arg-type] - file_id=original_file_id, + file_id=provider_file_id, ) verbose_logger.debug( - f"Successfully retrieved file object for {file_attr}={original_file_id}" + f"Successfully retrieved file object for {file_attr}={provider_file_id}" ) except Exception as e: verbose_logger.warning( - f"Failed to retrieve file object for {file_attr}={original_file_id}: {str(e)}. Storing with None and will fetch on-demand." + f"Failed to retrieve file object for {file_attr}={provider_file_id}: {str(e)}. Storing with None and will fetch on-demand." ) await self.store_unified_file_id( file_id=unified_file_id, file_object=file_object, litellm_parent_otel_span=user_api_key_dict.parent_otel_span, - model_mappings={model_id: original_file_id}, + model_mappings={model_id: provider_file_id}, user_api_key_dict=user_api_key_dict, ) await self.store_unified_object_id( diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py index 9526304aff0..7ca86d24494 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_hook.py @@ -110,10 +110,9 @@ async def test_should_pass_credentials_to_afile_retrieve(): mock_afile_retrieve = AsyncMock(return_value=_make_file_object("file-output-abc")) - with patch( - "litellm.afile_retrieve", mock_afile_retrieve - ), patch( - "litellm.proxy.proxy_server.llm_router", mock_router + with ( + patch("litellm.afile_retrieve", mock_afile_retrieve), + patch("litellm.proxy.proxy_server.llm_router", mock_router), ): await managed_files.async_post_call_success_hook( data={}, @@ -128,7 +127,9 @@ async def test_should_pass_credentials_to_afile_retrieve(): f"afile_retrieve must receive api_key from router credentials. " f"Got kwargs: {call_kwargs.kwargs}" ) - assert call_kwargs.kwargs.get("api_base") == "https://my-azure.openai.azure.com/", ( + assert ( + call_kwargs.kwargs.get("api_base") == "https://my-azure.openai.azure.com/" + ), ( f"afile_retrieve must receive api_base from router credentials. " f"Got kwargs: {call_kwargs.kwargs}" ) @@ -150,10 +151,9 @@ async def test_should_fallback_when_no_router(): mock_afile_retrieve = AsyncMock(return_value=_make_file_object("file-output-abc")) - with patch( - "litellm.afile_retrieve", mock_afile_retrieve - ), patch( - "litellm.proxy.proxy_server.llm_router", None + with ( + patch("litellm.afile_retrieve", mock_afile_retrieve), + patch("litellm.proxy.proxy_server.llm_router", None), ): await managed_files.async_post_call_success_hook( data={}, @@ -165,3 +165,60 @@ async def test_should_fallback_when_no_router(): call_kwargs = mock_afile_retrieve.call_args assert call_kwargs.kwargs.get("custom_llm_provider") == "azure" assert call_kwargs.kwargs.get("file_id") == "file-output-abc" + + +@pytest.mark.asyncio +async def test_should_not_double_wrap_already_unified_output_file_id(): + """After ensure_batch_response_managed_file_ids, retrieve must not re-wrap + output_file_id or store a nested unified id as the provider mapping.""" + import base64 + + managed_files = _make_managed_files_instance() + provider_file_id = "file-WXWt9R4LzmU5WpeKzjCfLR" + model_id = "openai/openai/gpt-5.5-batch" + already_unified = managed_files.get_unified_output_file_id( + output_file_id=provider_file_id, + model_id=model_id, + model_name="openai/openai/gpt-5.5-batch", + ) + + batch_response = _make_batch_response( + model_id=model_id, + model_name="openai/openai/gpt-5.5-batch", + output_file_id=already_unified, + ) + user_api_key_dict = _make_user_api_key_dict() + + mock_credentials = { + "api_key": "test-key", + "api_base": "https://api.openai.com/v1", + "custom_llm_provider": "openai", + } + mock_router = MagicMock() + mock_router.get_deployment_credentials_with_provider = MagicMock( + return_value=mock_credentials + ) + mock_afile_retrieve = AsyncMock(return_value=_make_file_object(provider_file_id)) + + with ( + patch("litellm.afile_retrieve", mock_afile_retrieve), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + ): + await managed_files.async_post_call_success_hook( + data={}, + user_api_key_dict=user_api_key_dict, + response=batch_response, + ) + + assert batch_response.output_file_id == already_unified + mock_afile_retrieve.assert_called_once() + assert mock_afile_retrieve.call_args.kwargs["file_id"] == provider_file_id + managed_files.store_unified_file_id.assert_awaited_once() + assert managed_files.store_unified_file_id.await_args.kwargs["model_mappings"] == { + model_id: provider_file_id + } + + decoded = base64.urlsafe_b64decode( + already_unified + "=" * (-len(already_unified) % 4) + ).decode() + assert decoded.count(f"llm_output_file_id,{provider_file_id}") == 1