mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): normalize batch file IDs before ManagedObjectTable write (#28339)
* fix(proxy): normalize batch file IDs before ManagedObjectTable write Run post_call_success_hook before update_batch_in_database on retrieve/cancel, and ensure_batch_response_managed_file_ids so file_object never stores raw provider output_file_id or error_file_id. Co-authored-by: Cursor <cursoragent@cursor.com> * fix(proxy): address Greptile review on batch file ID normalization Remove redundant resolve_* calls after update_batch_in_database and rename loop variable to avoid shadowing hidden_params unified_file_id. Co-authored-by: Cursor <cursoragent@cursor.com> * fix(tests): add mistral/ministral-8b-2512 to cost map and backfill in conftest Mistral rotated the 'mistral/mistral-tiny' alias to return 'ministral-8b-2512' as the response model, which was missing from the cost map. This caused test_completion_mistral_api and test_completion_mistral_api_modified_input to fail in litellm.completion_cost lookup. - Add mistral/ministral-8b-2512 entry to both the in-tree model_prices_and_context_window.json and the bundled litellm/model_prices_and_context_window_backup.json (mirrors the existing openrouter/mistralai/ministral-8b-2512 pricing). - litellm.model_cost is loaded at import time from the URL pinned to main, so the new backup entry isn't visible at test runtime until it also lands on main. Backfill any entries missing from the remote-fetched map into litellm.model_cost in the local_testing conftest so cost-calculator lookups succeed on this branch. * fix(tests): drop unnecessary del of conftest backfill loop vars * fix: resolve batch response file IDs even when status unchanged The status-unchanged early return in update_batch_in_database was skipping ensure_batch_response_managed_file_ids, leaving raw provider input_file_id (and other raw IDs) in the user-facing response when polling an in-progress batch. Move the in-place file ID normalization above the early return so the response always carries unified managed IDs while still skipping the DB write when nothing changed. Co-authored-by: Yassin Kortam <yassin@berri.ai> * test(batches): cover ensure_batch_response_managed_file_ids branches Add tests for the previously-uncovered paths in ensure_batch_response_managed_file_ids: error_file_id normalization, swallowed conversion errors, UserAPIKeyAuth fallback from db_batch_object, model_name resolution from unified_file_id, and early returns when managed_files_obj, model_id, or auth context are missing. --------- Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Co-authored-by: Claude <claude@anthropic.com> Co-authored-by: Yassin Kortam <yassin@berri.ai> Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
parent
1b733b966d
commit
679e6e346f
4 changed files with 356 additions and 20 deletions
|
|
@ -523,6 +523,10 @@ async def retrieve_batch( # noqa: PLR0915
|
|||
custom_llm_provider=custom_llm_provider, **data # type: ignore
|
||||
)
|
||||
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
# FIX: Update the database with the latest state from provider
|
||||
await update_batch_in_database(
|
||||
batch_id=batch_id,
|
||||
|
|
@ -533,19 +537,9 @@ async def retrieve_batch( # noqa: PLR0915
|
|||
verbose_proxy_logger=verbose_proxy_logger,
|
||||
db_batch_object=db_batch_object,
|
||||
operation="retrieve",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
### CALL HOOKS ### - modify outgoing data
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
# Fix: bug_feb14_batch_retrieve_returns_raw_input_file_id
|
||||
# Resolve raw provider file IDs (input, output, error) to unified IDs.
|
||||
if unified_batch_id:
|
||||
await resolve_input_file_id_to_unified(response, prisma_client)
|
||||
await resolve_output_file_ids_to_unified(response, prisma_client)
|
||||
|
||||
### ALERTING ###
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(
|
||||
|
|
@ -917,10 +911,14 @@ async def cancel_batch(
|
|||
**_cancel_batch_data,
|
||||
)
|
||||
|
||||
# FIX: Update the database with the new cancelled state
|
||||
managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files")
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
)
|
||||
|
||||
# FIX: Update the database with the new cancelled state
|
||||
await update_batch_in_database(
|
||||
batch_id=batch_id,
|
||||
unified_batch_id=unified_batch_id,
|
||||
|
|
@ -929,11 +927,7 @@ async def cancel_batch(
|
|||
prisma_client=prisma_client,
|
||||
verbose_proxy_logger=verbose_proxy_logger,
|
||||
operation="cancel",
|
||||
)
|
||||
|
||||
### CALL HOOKS ### - modify outgoing data
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
data=data, user_api_key_dict=user_api_key_dict, response=response
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
### ALERTING ###
|
||||
|
|
|
|||
|
|
@ -727,6 +727,76 @@ async def resolve_output_file_ids_to_unified(response, prisma_client) -> None:
|
|||
pass
|
||||
|
||||
|
||||
async def ensure_batch_response_managed_file_ids(
|
||||
response,
|
||||
managed_files_obj,
|
||||
prisma_client,
|
||||
verbose_proxy_logger,
|
||||
user_api_key_dict=None,
|
||||
db_batch_object=None,
|
||||
) -> None:
|
||||
"""Normalize batch file IDs to managed unified IDs before DB persistence."""
|
||||
await resolve_input_file_id_to_unified(response, prisma_client)
|
||||
await resolve_output_file_ids_to_unified(response, prisma_client)
|
||||
|
||||
if managed_files_obj is None:
|
||||
return
|
||||
|
||||
hidden_params = getattr(response, "_hidden_params", None) or {}
|
||||
model_id = hidden_params.get("model_id")
|
||||
if not model_id:
|
||||
return
|
||||
|
||||
model_name = hidden_params.get("model_name")
|
||||
unified_file_id = hidden_params.get("unified_file_id")
|
||||
if not model_name and isinstance(unified_file_id, str):
|
||||
decoded_unified_file_id = (
|
||||
_is_base64_encoded_unified_file_id(unified_file_id) or unified_file_id
|
||||
)
|
||||
target_model_names = get_models_from_unified_file_id(decoded_unified_file_id)
|
||||
if target_model_names:
|
||||
model_name = ",".join(target_model_names)
|
||||
|
||||
if user_api_key_dict is None and db_batch_object is not None:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_id=getattr(db_batch_object, "created_by", None) or "default-user-id",
|
||||
team_id=getattr(db_batch_object, "team_id", None),
|
||||
)
|
||||
if user_api_key_dict is None:
|
||||
return
|
||||
|
||||
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):
|
||||
continue
|
||||
try:
|
||||
new_unified_file_id = managed_files_obj.get_unified_output_file_id(
|
||||
output_file_id=raw_file_id,
|
||||
model_id=model_id,
|
||||
model_name=model_name,
|
||||
)
|
||||
await managed_files_obj.store_unified_file_id(
|
||||
file_id=new_unified_file_id,
|
||||
file_object=None,
|
||||
litellm_parent_otel_span=getattr(
|
||||
user_api_key_dict, "parent_otel_span", None
|
||||
),
|
||||
model_mappings={model_id: raw_file_id},
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
setattr(response, file_attr, new_unified_file_id)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Converted batch {file_attr} {raw_file_id!r} to managed ID before DB write"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
f"Failed to convert batch {file_attr}={raw_file_id!r} to managed ID "
|
||||
f"before DB write: {e}"
|
||||
)
|
||||
|
||||
|
||||
async def get_batch_from_database(
|
||||
batch_id: str,
|
||||
unified_batch_id: Union[str, Literal[False]],
|
||||
|
|
@ -800,6 +870,7 @@ async def update_batch_in_database(
|
|||
verbose_proxy_logger,
|
||||
db_batch_object=None,
|
||||
operation: str = "update",
|
||||
user_api_key_dict=None,
|
||||
):
|
||||
"""
|
||||
Update batch status and object in ManagedObjectTable.
|
||||
|
|
@ -813,6 +884,7 @@ async def update_batch_in_database(
|
|||
verbose_proxy_logger: Logger instance
|
||||
db_batch_object: Optional existing database object (for comparison)
|
||||
operation: Description of operation ("update", "cancel", etc.)
|
||||
user_api_key_dict: Optional auth context for creating managed file IDs
|
||||
"""
|
||||
import litellm.utils
|
||||
|
||||
|
|
@ -823,6 +895,18 @@ async def update_batch_in_database(
|
|||
if not prisma_client:
|
||||
return
|
||||
|
||||
# Always normalize the response's file IDs to unified managed IDs
|
||||
# (mutates in place) so the caller returns unified IDs to the user
|
||||
# even when we skip the DB update below for an unchanged status.
|
||||
await ensure_batch_response_managed_file_ids(
|
||||
response=response,
|
||||
managed_files_obj=managed_files_obj,
|
||||
prisma_client=prisma_client,
|
||||
verbose_proxy_logger=verbose_proxy_logger,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
db_batch_object=db_batch_object,
|
||||
)
|
||||
|
||||
# Only update if status has changed (when db_batch_object is provided)
|
||||
if db_batch_object and response.status == db_batch_object.status:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -31,10 +31,8 @@ import litellm
|
|||
# the cassette state the branch is being tested with.
|
||||
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
|
||||
|
||||
_local_cost_map = GetModelCostMap.load_local_model_cost_map()
|
||||
for _k, _v in _local_cost_map.items():
|
||||
for _k, _v in GetModelCostMap.load_local_model_cost_map().items():
|
||||
litellm.model_cost.setdefault(_k, _v)
|
||||
del _local_cost_map
|
||||
|
||||
from tests._vcr_conftest_common import ( # noqa: E402,F401
|
||||
VerboseReporterState,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,260 @@
|
|||
"""Regression: update_batch_in_database must not persist raw provider output_file_id."""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from typing import Optional
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
ensure_batch_response_managed_file_ids,
|
||||
update_batch_in_database,
|
||||
)
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
|
||||
def _build_batch_response(
|
||||
*,
|
||||
batch_id: str = "batch_managed_ids_test",
|
||||
status: str = "completed",
|
||||
output_file_id: Optional[str] = "file-rawoutput789",
|
||||
error_file_id: Optional[str] = None,
|
||||
hidden_params: Optional[dict] = None,
|
||||
) -> LiteLLMBatch:
|
||||
batch = LiteLLMBatch(
|
||||
id=batch_id,
|
||||
object="batch",
|
||||
status=status,
|
||||
endpoint="/v1/chat/completions",
|
||||
input_file_id="file-input123",
|
||||
output_file_id=output_file_id,
|
||||
error_file_id=error_file_id,
|
||||
completion_window="24h",
|
||||
created_at=1234567890,
|
||||
)
|
||||
if hidden_params is not None:
|
||||
batch._hidden_params = hidden_params # type: ignore[attr-defined]
|
||||
return batch
|
||||
|
||||
|
||||
def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ="):
|
||||
mock = MagicMock()
|
||||
mock.get_unified_output_file_id = MagicMock(return_value=unified_id)
|
||||
mock.store_unified_file_id = AsyncMock()
|
||||
return mock
|
||||
|
||||
|
||||
def _build_prisma_mock():
|
||||
mock = MagicMock()
|
||||
mock.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
|
||||
mock.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
return mock
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_batch_in_database_stores_unified_output_file_id():
|
||||
raw_output_file_id = "file-rawoutput789"
|
||||
unified_output_file_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
|
||||
batch_id = "batch_managed_ids_test"
|
||||
unified_batch_id = (
|
||||
"litellm_proxy;model_id:my-model;llm_batch_id:batch_managed_ids_test"
|
||||
)
|
||||
|
||||
response = _build_batch_response(
|
||||
batch_id=batch_id,
|
||||
output_file_id=raw_output_file_id,
|
||||
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
|
||||
)
|
||||
|
||||
mock_managed_files = _build_managed_files_mock(unified_id=unified_output_file_id)
|
||||
mock_prisma = _build_prisma_mock()
|
||||
|
||||
await update_batch_in_database(
|
||||
batch_id=batch_id,
|
||||
unified_batch_id=unified_batch_id,
|
||||
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"),
|
||||
)
|
||||
|
||||
stored = json.loads(
|
||||
mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs["data"][
|
||||
"file_object"
|
||||
]
|
||||
)
|
||||
assert stored["output_file_id"] == unified_output_file_id
|
||||
assert stored["output_file_id"] != raw_output_file_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_normalizes_error_file_id():
|
||||
"""Both output_file_id and error_file_id must be normalized to managed IDs."""
|
||||
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
|
||||
response = _build_batch_response(
|
||||
output_file_id="file-raw-output",
|
||||
error_file_id="file-raw-error",
|
||||
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
|
||||
)
|
||||
|
||||
mock_managed_files = _build_managed_files_mock(unified_id=unified_id)
|
||||
mock_prisma = _build_prisma_mock()
|
||||
|
||||
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 == unified_id
|
||||
assert response.error_file_id == unified_id
|
||||
assert mock_managed_files.get_unified_output_file_id.call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_swallows_conversion_errors():
|
||||
"""When the managed-files conversion raises, the failure is logged, not propagated."""
|
||||
raw_output_file_id = "file-raw-output"
|
||||
response = _build_batch_response(
|
||||
output_file_id=raw_output_file_id,
|
||||
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
|
||||
)
|
||||
|
||||
mock_managed_files = MagicMock()
|
||||
mock_managed_files.get_unified_output_file_id = MagicMock(
|
||||
side_effect=RuntimeError("boom")
|
||||
)
|
||||
mock_managed_files.store_unified_file_id = AsyncMock()
|
||||
|
||||
mock_logger = MagicMock()
|
||||
await ensure_batch_response_managed_file_ids(
|
||||
response=response,
|
||||
managed_files_obj=mock_managed_files,
|
||||
prisma_client=_build_prisma_mock(),
|
||||
verbose_proxy_logger=mock_logger,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
|
||||
)
|
||||
|
||||
assert response.output_file_id == raw_output_file_id
|
||||
mock_logger.warning.assert_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_builds_auth_from_db_batch_object():
|
||||
"""If user_api_key_dict is omitted, fall back to created_by/team_id on db_batch_object."""
|
||||
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
|
||||
response = _build_batch_response(
|
||||
output_file_id="file-raw-output",
|
||||
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
|
||||
)
|
||||
|
||||
mock_managed_files = _build_managed_files_mock(unified_id=unified_id)
|
||||
db_batch_object = SimpleNamespace(
|
||||
created_by="user-from-db", team_id="team-from-db", status="completed"
|
||||
)
|
||||
|
||||
await ensure_batch_response_managed_file_ids(
|
||||
response=response,
|
||||
managed_files_obj=mock_managed_files,
|
||||
prisma_client=_build_prisma_mock(),
|
||||
verbose_proxy_logger=MagicMock(),
|
||||
db_batch_object=db_batch_object,
|
||||
)
|
||||
|
||||
forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[
|
||||
"user_api_key_dict"
|
||||
]
|
||||
assert forwarded_auth.user_id == "user-from-db"
|
||||
assert forwarded_auth.team_id == "team-from-db"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_resolves_model_name_from_unified_file_id():
|
||||
"""When hidden_params lacks model_name, derive it from unified_file_id."""
|
||||
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
|
||||
response = _build_batch_response(
|
||||
output_file_id="file-raw-output",
|
||||
hidden_params={
|
||||
"model_id": "my-model",
|
||||
"unified_file_id": "litellm_proxy:application/octet-stream;unified_id,abc;target_model_names,gpt-4o-mini,gemini-2.0-flash",
|
||||
},
|
||||
)
|
||||
|
||||
mock_managed_files = _build_managed_files_mock(unified_id=unified_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 (
|
||||
mock_managed_files.get_unified_output_file_id.call_args.kwargs["model_name"]
|
||||
== "gpt-4o-mini,gemini-2.0-flash"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_returns_early_without_managed_files_obj():
|
||||
"""Without managed_files_obj, the helper is a no-op (no conversion attempted)."""
|
||||
response = _build_batch_response(
|
||||
output_file_id="file-raw-output",
|
||||
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
|
||||
)
|
||||
|
||||
await ensure_batch_response_managed_file_ids(
|
||||
response=response,
|
||||
managed_files_obj=None,
|
||||
prisma_client=_build_prisma_mock(),
|
||||
verbose_proxy_logger=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
|
||||
)
|
||||
|
||||
assert response.output_file_id == "file-raw-output"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_returns_early_without_model_id():
|
||||
"""Without model_id in hidden_params, the helper cannot create managed IDs."""
|
||||
response = _build_batch_response(
|
||||
output_file_id="file-raw-output",
|
||||
hidden_params={"model_name": "openai/gpt-4o"},
|
||||
)
|
||||
mock_managed_files = _build_managed_files_mock()
|
||||
|
||||
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 == "file-raw-output"
|
||||
mock_managed_files.get_unified_output_file_id.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_returns_early_without_auth():
|
||||
"""Without user_api_key_dict or db_batch_object, no conversion is attempted."""
|
||||
response = _build_batch_response(
|
||||
output_file_id="file-raw-output",
|
||||
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
|
||||
)
|
||||
mock_managed_files = _build_managed_files_mock()
|
||||
|
||||
await ensure_batch_response_managed_file_ids(
|
||||
response=response,
|
||||
managed_files_obj=mock_managed_files,
|
||||
prisma_client=_build_prisma_mock(),
|
||||
verbose_proxy_logger=MagicMock(),
|
||||
)
|
||||
|
||||
assert response.output_file_id == "file-raw-output"
|
||||
mock_managed_files.get_unified_output_file_id.assert_not_called()
|
||||
Loading…
Add table
Reference in a new issue