fix(proxy): apply model grants to unified file and batch ids on batch routes

Unified ids carry the deployment model inside the id, so a restricted key could
create, retrieve or cancel a batch on a deployment it is not granted. The model
parsed from a unified id now goes through the same grant check as header,
query and model-encoded id sources before the router is called.
This commit is contained in:
mubashir1osmani 2026-09-10 17:54:57 -04:00
parent edd5727f3c
commit bae731ddfc
3 changed files with 75 additions and 6 deletions

View file

@ -29,6 +29,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
add_internal_model_credentials,
apply_team_provider_credentials,
authorize_model_for_key,
batch_cost_poller_is_active,
decode_model_from_file_id,
encode_batch_response_ids,
@ -286,6 +287,7 @@ async def create_batch(
detail={"error": f"Expected 1 model, got {len(target_model_names)}"},
)
model: Final = target_model_names[0]
await authorize_model_for_key(model_id=model, llm_router=llm_router, user_api_key_dict=user_api_key_dict)
_create_batch_data["model"] = model
resolved_storage_url: Final = await _resolve_managed_input_file_storage_url(input_file_id)
@ -582,10 +584,17 @@ async def retrieve_batch(
)
if unified_batch_id:
unified_model_id: Final = get_model_id_from_unified_batch_id(unified_batch_id)
if unified_model_id is not None:
await authorize_model_for_key(
model_id=llm_router.resolve_model_name_from_model_id(unified_model_id) or unified_model_id,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
)
add_internal_model_credentials(
data=data,
llm_router=llm_router,
model_id=get_model_id_from_unified_batch_id(unified_batch_id),
model_id=unified_model_id,
)
response = await llm_router.aretrieve_batch(**data)
@ -998,6 +1007,11 @@ async def cancel_batch(
status_code=400,
detail={"error": "Invalid LiteLLM managed batch ID. Missing model_id."},
)
await authorize_model_for_key(
model_id=llm_router.resolve_model_name_from_model_id(model_id_from_batch) or model_id_from_batch,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
)
data["model"] = model_id_from_batch
data["batch_id"] = get_batch_id_from_unified_batch_id(unified_batch_id)
response = await llm_router.acancel_batch(**data)

View file

@ -179,6 +179,8 @@ def harness():
router = MagicMock(spec=Router)
router.model_group_alias = {}
router.get_model_access_groups = MagicMock(return_value={})
router.resolve_model_name_from_model_id = MagicMock(side_effect=lambda model_id: model_id)
router.model_list = []
router.acreate_batch = AsyncMock(return_value=make_batch())
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
@ -1165,6 +1167,8 @@ def retrieve_harness():
router = MagicMock(spec=Router)
router.model_group_alias = {}
router.get_model_access_groups = MagicMock(return_value={})
router.resolve_model_name_from_model_id = MagicMock(side_effect=lambda model_id: model_id)
router.model_list = []
router.aretrieve_batch = AsyncMock(return_value=make_batch())
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
@ -1622,6 +1626,8 @@ def list_harness():
router = MagicMock(spec=Router)
router.model_group_alias = {}
router.get_model_access_groups = MagicMock(return_value={})
router.resolve_model_name_from_model_id = MagicMock(side_effect=lambda model_id: model_id)
router.model_list = []
router.alist_batches = AsyncMock(return_value=FakeListPage([]))
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
@ -2020,6 +2026,8 @@ def cancel_harness():
router = MagicMock(spec=Router)
router.model_group_alias = {}
router.get_model_access_groups = MagicMock(return_value={})
router.resolve_model_name_from_model_id = MagicMock(side_effect=lambda model_id: model_id)
router.model_list = []
router.acancel_batch = AsyncMock(return_value=make_batch())
router.get_deployment_credentials_with_provider = MagicMock(side_effect=_creds_lookup)
@ -2816,3 +2824,53 @@ async def test_cancel__model_encoded_id_rejects_key_without_model_grant(cancel_h
assert exc_info.value.code == "403"
cancel_harness.creds_resolver.assert_not_called()
cancel_harness.litellm_acancel.assert_not_called()
def _b64_unified_id(decoded: str) -> str:
return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=")
UNIFIED_FILE_ID_FOR_GPT4O_MINI = _b64_unified_id(
"litellm_proxy:application/octet-stream;unified_id,c4843482-b176-4901-8292-7523fd0f2c6e;"
"target_model_names,gpt-4o-mini;llm_output_file_id,file-provider;llm_output_file_model_id,dep-1"
)
UNIFIED_BATCH_ID_FOR_GPT4O_MINI = _b64_unified_id(UNIFIED_BATCH_ID)
@pytest.mark.asyncio
async def test_create__unified_file_id_rejects_key_without_model_grant(harness):
"""The model carried inside a unified file id is caller-controlled too, so it is checked against the key's grants."""
set_body(
harness,
{
"input_file_id": UNIFIED_FILE_ID_FOR_GPT4O_MINI,
"endpoint": "/v1/chat/completions",
"completion_window": "24h",
},
)
with pytest.raises(ProxyException) as exc_info:
await call_create(harness, user=_key_restricted_to("vertex-model"))
assert exc_info.value.code == "403"
harness.router_acreate.assert_not_called()
harness.litellm_acreate.assert_not_called()
@pytest.mark.asyncio
async def test_retrieve__unified_batch_id_rejects_key_without_model_grant(retrieve_harness):
with pytest.raises(ProxyException) as exc_info:
await call_retrieve(retrieve_harness, UNIFIED_BATCH_ID_FOR_GPT4O_MINI, user=_key_restricted_to("vertex-model"))
assert exc_info.value.code == "403"
retrieve_harness.router_aretrieve.assert_not_called()
retrieve_harness.creds_resolver.assert_not_called()
@pytest.mark.asyncio
async def test_cancel__unified_batch_id_rejects_key_without_model_grant(cancel_harness):
with pytest.raises(ProxyException) as exc_info:
await call_cancel(cancel_harness, UNIFIED_BATCH_ID_FOR_GPT4O_MINI, user=_key_restricted_to("vertex-model"))
assert exc_info.value.code == "403"
cancel_harness.router_acancel.assert_not_called()

View file

@ -11,6 +11,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.openai_files_endpoints.common_utils import (
decode_model_from_file_id,
get_batch_id_from_unified_batch_id,
@ -412,11 +413,7 @@ async def test_cancel_batch_with_unified_id_routes_with_decoded_model_and_batch_
mock_request.url.path = f"/v1/batches/{unified_batch_id}/cancel"
mock_fastapi_response = MagicMock()
mock_fastapi_response.headers = {}
mock_user_api_key_dict = MagicMock()
mock_user_api_key_dict.parent_otel_span = None
mock_user_api_key_dict.user_id = "test_user"
mock_user_api_key_dict.allowed_model_region = None
mock_user_api_key_dict.team_metadata = {}
mock_user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="test_user", team_metadata={})
with (
patch("litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing") as mock_processor_cls,