This commit is contained in:
IvanShang 2026-08-26 15:35:14 +08:00 committed by GitHub
commit 183cc9a677
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 71 additions and 2 deletions

View file

@ -1215,8 +1215,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,

View file

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