mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
51ba6e39cd
commit
6748607d91
2 changed files with 88 additions and 20 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue