diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index ae5905f9cdf..a1f63f388b4 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -504,7 +504,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if retrieve_file_id else False ) - if potential_file_id: + if potential_file_id and "llm_output_file_id," in potential_file_id: model_id = self.get_model_id_from_unified_file_id(potential_file_id) if model_id: data["model"] = model_id @@ -1058,7 +1058,12 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): return file_id.split("llm_output_file_model_id,")[1].split(";")[0] def get_output_file_id_from_unified_file_id(self, file_id: str) -> str: - return file_id.split("llm_output_file_id,")[1].split(";")[0] + marker = "llm_output_file_id," + if marker not in file_id: + raise ValueError( + f"Unified id does not contain {marker!r}: {file_id[:80]!r}" + ) + return file_id.split(marker, 1)[1].split(";")[0] async def async_post_call_success_hook( self, data: Dict, user_api_key_dict: UserAPIKeyAuth, response: LLMResponseTypes @@ -1099,13 +1104,33 @@ 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_output_file_id = _is_base64_encoded_unified_file_id( + file_id_value ) - setattr(response, file_attr, unified_file_id) + if ( + decoded_output_file_id + and "llm_output_file_id," in decoded_output_file_id + ): + provider_file_id = ( + self.get_output_file_id_from_unified_file_id( + decoded_output_file_id + ) + ) + unified_file_id = file_id_value + elif decoded_output_file_id: + verbose_logger.warning( + f"Skipping {file_attr}={file_id_value!r}: " + "unified id is not a managed file output id" + ) + continue + 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 +1150,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..1336490a344 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,93 @@ 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 + + +@pytest.mark.asyncio +async def test_should_skip_non_file_unified_id_on_output_file_id(): + """Batch-style unified ids lack llm_output_file_id; must not IndexError or re-wrap.""" + import base64 + + managed_files = _make_managed_files_instance() + batch_unified = ( + base64.urlsafe_b64encode( + b"litellm_proxy;model_id:openai/openai/gpt-5.5-batch;llm_batch_id:batch_abc" + ) + .decode() + .rstrip("=") + ) + + batch_response = _make_batch_response( + model_id="openai/openai/gpt-5.5-batch", + model_name="openai/openai/gpt-5.5-batch", + output_file_id=batch_unified, + ) + user_api_key_dict = _make_user_api_key_dict() + + with patch("litellm.afile_retrieve", AsyncMock()) as mock_afile_retrieve: + 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 == batch_unified + mock_afile_retrieve.assert_not_called() + managed_files.store_unified_file_id.assert_not_awaited()