From 13f4b383f1d07d4e0ba57d95f61b67d0ccf5b299 Mon Sep 17 00:00:00 2001 From: Hermes Agent Date: Fri, 21 Aug 2026 05:26:32 +0800 Subject: [PATCH] fix(proxy): rewrap chained batch output file IDs --- .../openai_files_endpoints/common_utils.py | 11 +++- ..._batch_update_db_managed_output_file_id.py | 62 ++++++++++++++++++- 2 files changed, 71 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index b9af01e9aea..13d8bb6dd85 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -1160,8 +1160,17 @@ async def ensure_batch_response_managed_file_ids( 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): + if not raw_file_id: continue + if _is_base64_encoded_unified_file_id(raw_file_id): + try: + managed_file = await ManagedFileRepository(prisma_client).table.find_first( + where={"unified_file_id": raw_file_id} + ) + if managed_file: + continue + except Exception: + continue try: new_unified_file_id = managed_files_obj.get_unified_output_file_id( output_file_id=raw_file_id, 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 index ebd33aa2e53..794199fff1f 100644 --- 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 @@ -1,11 +1,13 @@ """Regression: update_batch_in_database must not persist raw provider output_file_id.""" +import base64 import json from types import SimpleNamespace from typing import Optional -import pytest from unittest.mock import AsyncMock, MagicMock +import pytest + from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.openai_files_endpoints.common_utils import ( ensure_batch_response_managed_file_ids, @@ -213,6 +215,64 @@ async def test_ensure_batch_response_normalizes_error_file_id(): assert mock_managed_files.get_unified_output_file_id.call_count == 2 +@pytest.mark.asyncio +async def test_ensure_batch_response_rewraps_chained_proxy_output_file_id(): + upstream_unified_file_id = base64.urlsafe_b64encode( + b"litellm_proxy:application/json;unified_id,upstream-uuid;target_model_names,upstream-model;llm_output_file_id,upstream-file;llm_output_file_model_id,upstream-model-id" + ).decode().rstrip("=") + front_proxy_unified_file_id = "file-front-proxy-output" + response = _build_batch_response( + output_file_id=upstream_unified_file_id, + hidden_params={"model_id": "front-proxy-model-id", "model_name": "litellm_proxy/upstream-model"}, + ) + mock_managed_files = _build_managed_files_mock(unified_id=front_proxy_unified_file_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 response.output_file_id == front_proxy_unified_file_id + mock_managed_files.get_unified_output_file_id.assert_called_once_with( + output_file_id=upstream_unified_file_id, + model_id="front-proxy-model-id", + model_name="litellm_proxy/upstream-model", + ) + assert mock_managed_files.store_unified_file_id.call_args.kwargs["model_mappings"] == { + "front-proxy-model-id": upstream_unified_file_id + } + + +@pytest.mark.asyncio +async def test_ensure_batch_response_preserves_registered_output_file_id(): + registered_unified_file_id = base64.urlsafe_b64encode( + b"litellm_proxy:application/json;unified_id,front-uuid;target_model_names,front-model;llm_output_file_id,front-file;llm_output_file_model_id,front-model-id" + ).decode().rstrip("=") + response = _build_batch_response( + output_file_id=registered_unified_file_id, + hidden_params={"model_id": "front-proxy-model-id", "model_name": "litellm_proxy/upstream-model"}, + ) + mock_managed_files = _build_managed_files_mock() + mock_prisma = _build_prisma_mock() + mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock( + return_value=SimpleNamespace(unified_file_id=registered_unified_file_id) + ) + + 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 == registered_unified_file_id + mock_managed_files.get_unified_output_file_id.assert_not_called() + + @pytest.mark.asyncio async def test_ensure_batch_response_swallows_conversion_errors(): """When the managed-files conversion raises, the failure is logged, not propagated."""