diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index 896b7ca33d7..b0996919fef 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -30,6 +30,29 @@ if TYPE_CHECKING: router: Final = APIRouter() +def _merge_credentials_keeping_requested_model( + data: dict, + credentials: dict, + requested_model: str | None, + file_id: str | None = None, +) -> None: + """ + Merge deployment credentials into the request while keeping the caller's + model group name in ``model``. + + Vector store file routes dispatch through the Router, which resolves the + deployment from ``data["model"]``. ``get_deployment_credentials_with_provider`` + returns the deployment's underlying ``litellm_params.model``, so merging it + verbatim erases the requested group: when several groups share one provider + model, the Router falls back to matching by litellm model and can serve the + request from a different group, with that group's access restrictions, + settings and spend attribution. + """ + prepare_data_with_credentials(data=data, credentials=credentials, file_id=file_id) + if requested_model is not None: + data["model"] = requested_model + + def _update_request_data_with_managed_file_id( data: dict, file_id: str, @@ -104,9 +127,10 @@ def _update_request_data_with_managed_file_id( if llm_router: credentials = llm_router.get_deployment_credentials_with_provider(model_id=routing_model) if credentials: - prepare_data_with_credentials( + _merge_credentials_keeping_requested_model( data=data, credentials=credentials, + requested_model=routing_model, file_id=llm_output_file_id, # Use the actual provider file ID ) verbose_logger.info( @@ -141,9 +165,10 @@ def _update_request_data_with_managed_file_id( if should_route: # Use model-based routing with credentials from config - prepare_data_with_credentials( + _merge_credentials_keeping_requested_model( data=data, credentials=credentials, + requested_model=model_used, file_id=original_file_id, # Use decoded file ID if from encoded ID ) @@ -236,6 +261,7 @@ async def _update_request_data_with_model_routing_hint( should_route = False credentials = None + routed_model: str | None = None if isinstance(model_hint, str) and "*" in model_hint: if llm_router is not None: if should_authorize_model_hint: @@ -248,6 +274,7 @@ async def _update_request_data_with_model_routing_hint( model_id=model_hint, team_id=caller_team_id ) should_route = credentials is not None + routed_model = model_hint else: if isinstance(model_hint, str) and should_authorize_model_hint: await _authorize_model_routing_hint( @@ -257,7 +284,7 @@ async def _update_request_data_with_model_routing_hint( ) ( should_route, - _model_used, + routed_model, _original_file_id, credentials, ) = handle_model_based_routing( @@ -269,9 +296,10 @@ async def _update_request_data_with_model_routing_hint( ) if should_route and credentials is not None: - prepare_data_with_credentials( + _merge_credentials_keeping_requested_model( data=data, credentials=credentials, + requested_model=routed_model, ) return data @@ -293,6 +321,7 @@ async def _update_request_data_with_model_routing_hint( model_names_to_check.append(model_name) openai_credentials = None + openai_credentials_model = None for model_name in model_names_to_check: credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_name, team_id=caller_team_id) if credentials is None: @@ -313,9 +342,14 @@ async def _update_request_data_with_model_routing_hint( if openai_credentials is not None: return data openai_credentials = credentials + openai_credentials_model = model_name if openai_credentials is not None: - prepare_data_with_credentials(data=data, credentials=openai_credentials) + _merge_credentials_keeping_requested_model( + data=data, + credentials=openai_credentials, + requested_model=openai_credentials_model, + ) elif len(model_names_to_check) == 1: await _authorize_model_routing_hint( model=model_names_to_check[0], diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index e7de8b54e4e..ceac1da8dcc 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -174,7 +174,7 @@ async def test_vector_store_file_list_resolves_credentials_from_model_query_para assert result["api_key"] == "sk-team-openai" assert result["api_base"] == "https://api.openai.com/v1" - assert result["model"] == "openai/gpt-4o-mini" + assert result["model"] == "team-openai" assert "custom_llm_provider" not in result llm_router.get_deployment_credentials_with_provider.assert_called_once_with( model_id="team-openai" @@ -207,7 +207,7 @@ async def test_vector_store_file_list_resolves_single_openai_team_deployment(): assert result["api_key"] == "sk-team-openai" assert result["api_base"] == "https://api.openai.com/v1" - assert result["model"] == "openai/gpt-4o-mini" + assert result["model"] == "team-openai" assert "custom_llm_provider" not in result llm_router.get_deployment_credentials_with_provider.assert_called_once_with( model_id="team-openai", team_id=None @@ -245,11 +245,60 @@ async def test_vector_store_file_list_wildcard_model_hint_falls_back_to_team_dep assert result["api_key"] == "sk-team-openai" assert result["api_base"] == "https://api.openai.com/v1" - assert result["model"] == "openai/gpt-4o-mini" + assert result["model"] == "team-openai" assert "custom_llm_provider" not in result assert llm_router.get_deployment_credentials_with_provider.call_count == 3 +@pytest.mark.asyncio +async def test_vector_store_file_list_keeps_requested_model_group_when_groups_share_a_model(): + """ + Two model groups can point at the same ``litellm_params.model``. Merging the + deployment credentials must not replace the requested group name with that + shared provider model, otherwise the Router resolves the deployment by + litellm model and can serve the request from the other group (#36103). + """ + llm_router = litellm.Router( + model_list=[ + { + "model_name": "vip-embeddings", + "litellm_params": { + "model": "openai/text-embedding-3-small", + "api_key": "sk-vip", + }, + "model_info": {"id": "vip-dep"}, + }, + { + "model_name": "public-embeddings", + "litellm_params": { + "model": "openai/text-embedding-3-small", + "api_key": "sk-public", + }, + "model_info": {"id": "public-dep"}, + }, + ] + ) + + request = MagicMock(spec=Request) + request.query_params = {"model": "public-embeddings"} + request.headers = {} + + result = await _update_request_data_with_model_routing_hint( + data={"vector_store_id": "vs_123"}, + request=request, + llm_router=llm_router, + ) + + assert result["api_key"] == "sk-public" + assert result["model"] == "public-embeddings" + + for _ in range(20): + deployment = await llm_router.async_get_available_deployment( + model=result["model"], request_kwargs={} + ) + assert deployment["model_info"]["id"] == "public-dep" + + @pytest.mark.asyncio async def test_vector_store_file_list_authorizes_wildcard_query_param_before_credentials(): from litellm.proxy.auth.auth_checks import ProxyException diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py index b1bd7ccbf0f..e9c434bdbb8 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -169,7 +169,7 @@ async def test_vector_store_file_list_resolves_managed_vector_store_before_team_ assert response == {"ok": True} assert captured_data["vector_store_id"] == "vs_provider_native" assert captured_data["api_key"] == "sk-managed-deployment" - assert captured_data["model"] == "openai/managed-deployment" + assert captured_data["model"] == "managed-deployment" llm_router.get_deployment_credentials_with_provider.assert_called_once_with( model_id="managed-deployment" )