mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(proxy): register managed-file ids on terminal batch retrieve (#33989)
The terminal-state short-circuit in retrieve_batch only ran the lookup-only resolve_input_file_id_to_unified / resolve_output_file_ids_to_unified helpers. When the batch poller (CheckBatchCost failed/expired/cancelled branch) wrote a terminal row from provider output without registering managed-file rows, those lookups found nothing and GET /v1/batches/{id} returned the raw provider output_file_id/error_file_id. /v1/files/{id}/content only applies the managed-file ownership ACL to base64 unified ids, so a raw file-* id skipped the ownership check entirely, letting any authenticated key read another user's batch files.
Route the terminal short-circuit through ensure_batch_response_managed_file_ids, which registers the missing managed-file rows owned by the batch creator and rewrites the raw ids to unified ids. ensure_batch_response_managed_file_ids gains a unified_batch_id param so it can derive model_id when the rehydrated DB row has empty _hidden_params.
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
23de7a15d9
commit
6924f90133
4 changed files with 63 additions and 7 deletions
|
|
@ -27,6 +27,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
|
|||
decode_model_from_file_id,
|
||||
encode_batch_response_ids,
|
||||
encode_file_id_with_model,
|
||||
ensure_batch_response_managed_file_ids,
|
||||
get_batch_id_from_unified_batch_id,
|
||||
get_batch_from_database,
|
||||
get_credentials_for_model,
|
||||
|
|
@ -440,10 +441,19 @@ async def retrieve_batch(
|
|||
)
|
||||
|
||||
# The DB may store raw provider file IDs (before hooks translate them).
|
||||
# Resolve any raw input/output/error file IDs to unified IDs.
|
||||
# Register any missing managed-file rows and rewrite raw IDs to unified
|
||||
# IDs so terminal rows written by the poller don't leak raw provider
|
||||
# file IDs that bypass the /v1/files ownership check.
|
||||
if unified_batch_id:
|
||||
await resolve_input_file_id_to_unified(response, prisma_client)
|
||||
await resolve_output_file_ids_to_unified(response, prisma_client)
|
||||
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,
|
||||
unified_batch_id=unified_batch_id,
|
||||
)
|
||||
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.update_request_status(
|
||||
|
|
|
|||
|
|
@ -933,6 +933,7 @@ async def ensure_batch_response_managed_file_ids(
|
|||
verbose_proxy_logger,
|
||||
user_api_key_dict=None,
|
||||
db_batch_object=None,
|
||||
unified_batch_id: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Normalize batch file IDs to managed unified IDs before DB persistence."""
|
||||
await resolve_input_file_id_to_unified(response, prisma_client)
|
||||
|
|
@ -943,6 +944,8 @@ async def ensure_batch_response_managed_file_ids(
|
|||
|
||||
hidden_params = getattr(response, "_hidden_params", None) or {}
|
||||
model_id = hidden_params.get("model_id")
|
||||
if not model_id and unified_batch_id:
|
||||
model_id = get_model_id_from_unified_batch_id(unified_batch_id)
|
||||
if not model_id:
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -200,6 +200,40 @@ async def test_ensure_batch_response_resolves_model_name_from_unified_file_id():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_derives_model_id_from_unified_batch_id():
|
||||
"""Terminal DB rows have empty _hidden_params (rehydrated from stored JSON).
|
||||
|
||||
Regression for GH #33989: the helper must still register the raw output/error
|
||||
ids as managed files by deriving model_id from the unified batch id, otherwise
|
||||
the terminal-retrieve short-circuit leaks raw provider file ids.
|
||||
"""
|
||||
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
|
||||
response = _build_batch_response(
|
||||
output_file_id="file-raw-output",
|
||||
error_file_id="file-raw-error",
|
||||
hidden_params={},
|
||||
)
|
||||
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(),
|
||||
db_batch_object=SimpleNamespace(created_by="user-from-db", team_id=None, status="failed"),
|
||||
unified_batch_id="litellm_proxy;model_id:deployment-42;llm_batch_id:batch_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
|
||||
assert (
|
||||
mock_managed_files.get_unified_output_file_id.call_args.kwargs["model_id"]
|
||||
== "deployment-42"
|
||||
)
|
||||
|
||||
|
||||
@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)."""
|
||||
|
|
|
|||
|
|
@ -937,6 +937,7 @@ class RetrieveHarness:
|
|||
update_batch_in_db: AsyncMock
|
||||
resolve_input: AsyncMock
|
||||
resolve_output: AsyncMock
|
||||
ensure_managed: AsyncMock
|
||||
|
||||
@property
|
||||
def router_aretrieve(self) -> AsyncMock:
|
||||
|
|
@ -975,6 +976,7 @@ def retrieve_harness():
|
|||
update_batch_in_db = AsyncMock(return_value=None)
|
||||
resolve_input = AsyncMock(return_value=None)
|
||||
resolve_output = AsyncMock(return_value=None)
|
||||
ensure_managed = AsyncMock(return_value=None)
|
||||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(
|
||||
|
|
@ -1003,6 +1005,7 @@ def retrieve_harness():
|
|||
stack.enter_context(patch.object(endpoints, "update_batch_in_database", update_batch_in_db))
|
||||
stack.enter_context(patch.object(endpoints, "resolve_input_file_id_to_unified", resolve_input))
|
||||
stack.enter_context(patch.object(endpoints, "resolve_output_file_ids_to_unified", resolve_output))
|
||||
stack.enter_context(patch.object(endpoints, "ensure_batch_response_managed_file_ids", ensure_managed))
|
||||
stack.enter_context(patch.object(litellm, "aretrieve_batch", litellm_aretrieve))
|
||||
stack.enter_context(patch.object(litellm, "enable_loadbalancing_on_batch_endpoints", False))
|
||||
stack.enter_context(patch.object(proxy_server, "llm_router", router))
|
||||
|
|
@ -1026,6 +1029,7 @@ def retrieve_harness():
|
|||
update_batch_in_db=update_batch_in_db,
|
||||
resolve_input=resolve_input,
|
||||
resolve_output=resolve_output,
|
||||
ensure_managed=ensure_managed,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1260,15 +1264,20 @@ async def test_retrieve__db_terminal_state_short_circuits(retrieve_harness, stat
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve__db_terminal_unified_resolves_file_ids(retrieve_harness):
|
||||
db_batch_object = MagicMock()
|
||||
db_response = make_batch(id="batch-from-db", status="completed")
|
||||
retrieve_harness.get_batch_from_db.return_value = (MagicMock(), db_response)
|
||||
retrieve_harness.get_batch_from_db.return_value = (db_batch_object, db_response)
|
||||
|
||||
with patch.object(endpoints, "_is_base64_encoded_unified_file_id", return_value=UNIFIED_BATCH_ID):
|
||||
await call_retrieve(retrieve_harness, "batch-unified-blob")
|
||||
|
||||
# Terminal short-circuit still resolves raw provider file ids to unified.
|
||||
retrieve_harness.resolve_input.assert_called_once()
|
||||
retrieve_harness.resolve_output.assert_called_once()
|
||||
# Regression GH #33989: the terminal short-circuit must register missing
|
||||
# managed-file rows (owned by the batch creator) and rewrite raw provider
|
||||
# file ids, not merely look them up, so raw ids never leak past the ACL.
|
||||
retrieve_harness.ensure_managed.assert_called_once()
|
||||
ensure_kwargs = retrieve_harness.ensure_managed.call_args.kwargs
|
||||
assert ensure_kwargs["unified_batch_id"] == UNIFIED_BATCH_ID
|
||||
assert ensure_kwargs["db_batch_object"] is db_batch_object
|
||||
retrieve_harness.litellm_aretrieve.assert_not_called()
|
||||
retrieve_harness.router_aretrieve.assert_not_called()
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue