mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
edd5727f3c
commit
bae731ddfc
3 changed files with 75 additions and 6 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue