fix(proxy): skip double-wrapping unified batch output file ids on retrieve

After ensure_batch_response_managed_file_ids normalizes output_file_id, the managed files post-call hook was re-encoding the unified id and storing the nested id as the provider mapping. Use the decoded llm_output_file_id for retrieve and model_mappings instead.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-06-09 13:04:22 +05:30
parent 51ba6e39cd
commit 6748607d91
No known key found for this signature in database
2 changed files with 88 additions and 20 deletions

View file

@ -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(

View file

@ -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