mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(proxy): enforce the unified batch model grant before the DB shortcut and skip it for registry-routed vector store models
retrieve_batch returned a terminal batch from the DB before checking that the key may use the model encoded in a unified batch id; the grant check now runs right after pre-call processing. The vector store file list helper authorized data["model"] through handle_model_based_routing even when the vector store registry set it server-side and even with no caller, which crashed on a None key; it now authorizes only a caller-supplied hint and resolves credentials directly.
This commit is contained in:
parent
f3b198c1b7
commit
5783a38e27
4 changed files with 65 additions and 22 deletions
|
|
@ -470,6 +470,17 @@ async def retrieve_batch(
|
|||
route_type="aretrieve_batch",
|
||||
)
|
||||
|
||||
unified_model_id: Final = get_model_id_from_unified_batch_id(unified_batch_id) if unified_batch_id else None
|
||||
if unified_model_id is not None:
|
||||
resolved_unified_model: Final = (
|
||||
llm_router.resolve_model_name_from_model_id(unified_model_id) if llm_router is not None else None
|
||||
)
|
||||
await authorize_model_for_key(
|
||||
model_id=resolved_unified_model or unified_model_id,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# FIX: First, try to read from ManagedObjectTable for consistent state
|
||||
managed_files_obj: Final = proxy_logging_obj.get_proxy_hook("managed_files")
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
|
@ -590,13 +601,6 @@ 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,
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
get_custom_llm_provider_from_request_query,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
get_credentials_for_model,
|
||||
handle_model_based_routing,
|
||||
prepare_data_with_credentials,
|
||||
)
|
||||
|
|
@ -262,26 +263,15 @@ async def _update_request_data_with_model_routing_hint(
|
|||
model_id=model_hint, team_id=caller_team_id
|
||||
)
|
||||
should_route = credentials is not None
|
||||
else:
|
||||
if isinstance(model_hint, str) and should_authorize_model_hint:
|
||||
elif isinstance(model_hint, str):
|
||||
if should_authorize_model_hint:
|
||||
await _authorize_model_routing_hint(
|
||||
model=model_hint,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
(
|
||||
should_route,
|
||||
_model_used,
|
||||
_original_file_id,
|
||||
credentials,
|
||||
) = await handle_model_based_routing(
|
||||
file_id="",
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
check_file_id_encoding=False,
|
||||
)
|
||||
credentials = get_credentials_for_model(llm_router=llm_router, model_id=model_hint)
|
||||
should_route = True
|
||||
|
||||
if should_route and credentials is not None:
|
||||
prepare_data_with_credentials(
|
||||
|
|
|
|||
|
|
@ -2908,6 +2908,21 @@ async def test_retrieve__unified_batch_id_rejects_key_without_model_grant(retrie
|
|||
retrieve_harness.creds_resolver.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retrieve__unified_batch_id_rejects_key_without_model_grant_before_db_terminal_shortcut(
|
||||
retrieve_harness,
|
||||
):
|
||||
retrieve_harness.get_batch_from_db.return_value = (MagicMock(), make_batch(id="batch-from-db", status="completed"))
|
||||
|
||||
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.logging.post_call_success_hook.assert_not_called()
|
||||
retrieve_harness.ensure_managed_files.assert_not_called()
|
||||
retrieve_harness.router_aretrieve.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:
|
||||
|
|
|
|||
|
|
@ -306,6 +306,40 @@ async def test_vector_store_file_list_resolves_credentials_from_model_query_para
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_file_list_registry_routed_model_skips_key_model_grant():
|
||||
request = MagicMock(spec=Request)
|
||||
request.query_params = {}
|
||||
request.headers = {}
|
||||
|
||||
llm_router = MagicMock()
|
||||
llm_router.get_deployment_credentials_with_provider.return_value = {
|
||||
"api_key": "sk-team-openai",
|
||||
"api_base": "https://api.openai.com/v1",
|
||||
"custom_llm_provider": "openai",
|
||||
"model": "openai/gpt-4o-mini",
|
||||
}
|
||||
|
||||
data = {"vector_store_id": "vs_123", "model": "team-openai"}
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
models=["restricted-deployment"],
|
||||
team_models=["restricted-deployment"],
|
||||
)
|
||||
|
||||
result = await _update_request_data_with_model_routing_hint(
|
||||
data=data,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
assert result["api_key"] == "sk-team-openai"
|
||||
assert result["model"] == "openai/gpt-4o-mini"
|
||||
llm_router.get_deployment_credentials_with_provider.assert_called_once_with(
|
||||
model_id="team-openai"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vector_store_file_list_resolves_single_openai_team_deployment():
|
||||
request = MagicMock(spec=Request)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue