diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index db6ec754c6e..f6174f62ab4 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -1997,6 +1997,69 @@ async def health_liveliness_options(): return Response(headers=response_headers, status_code=200) +async def _authorize_test_connection( + *, + model_params: Any, + user_api_key_dict: UserAPIKeyAuth, + prisma_client: Any, + premium_user: bool, + llm_router: Any, + configured_model_name: str | None, + request_litellm_params: Mapping[str, object], +) -> None: + """Decide whether the caller may probe this model. + + Proxy admins and team admins may probe any model they manage, as before. + Any other user may probe a configured model they are allowed to call, but + only as configured: a request that sets its own connection fields describes + a different endpoint, and probing that stays a management operation. + """ + from litellm.proxy.auth.auth_checks import ( + can_key_call_model, + can_user_call_model, + get_user_object, + ) + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelManagementAuthChecks, + ) + from litellm.proxy.proxy_server import llm_model_list, user_api_key_cache + + try: + await ModelManagementAuthChecks.can_user_make_model_call( + model_params=model_params, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + premium_user=premium_user, + ) + return + except HTTPException as management_denial: + if management_denial.status_code != 403: + raise + if configured_model_name is None or any(field in request_litellm_params for field in _CONFIG_CONNECTION_FIELDS): + raise + + try: + await can_key_call_model( + model=configured_model_name, + llm_model_list=llm_model_list, + valid_token=user_api_key_dict, + llm_router=llm_router, + ) + user_object = await get_user_object( + user_id=user_api_key_dict.user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + ) + await can_user_call_model( + model=configured_model_name, + llm_router=llm_router, + user_object=user_object, + ) + except ProxyException as e: + raise HTTPException(status_code=403, detail={"error": str(e.message)}) from e + + @router.post( "/health/test_connection", tags=["health"], @@ -2083,9 +2146,6 @@ async def test_model_connection( dict: A dictionary containing the health check result with either success information or error details. """ from litellm.proxy._types import CommonProxyErrors - from litellm.proxy.management_endpoints.model_management_endpoints import ( - ModelManagementAuthChecks, - ) from litellm.proxy.proxy_server import ( general_settings, llm_router, @@ -2115,6 +2175,7 @@ async def test_model_connection( # This gets the litellm_params from proxy config (with resolved env vars) config_litellm_params: dict = {} loaded_model_info: dict | None = None + configured_model_name: str | None = None if llm_router is not None: # Prefer disambiguation by deployment id (`model_info.id`) when # the caller supplies it. This is required when multiple @@ -2134,6 +2195,7 @@ async def test_model_connection( if deployment_by_id is not None: config_litellm_params = deployment_by_id.litellm_params.model_dump(exclude_none=True) loaded_model_info = deployment_by_id.model_info.model_dump(exclude_none=True) + configured_model_name = deployment_by_id.model_name elif model_name: # Fall back to model_name lookup for callers (e.g. the # "Add Model" wizard, or curl) that don't supply an id. @@ -2156,6 +2218,7 @@ async def test_model_connection( # variables from proxy config. config_litellm_params = dict(deployments[0].get("litellm_params", {})) loaded_model_info = dict(deployments[0].get("model_info") or {}) + configured_model_name = deployments[0].get("model_name") except Exception as e: verbose_proxy_logger.debug( "Could not find model %s in router: %s. Proceeding with request params only.", model_name, e @@ -2178,7 +2241,7 @@ async def test_model_connection( ) ## Auth check, on the final probe params so health_check_params cannot retarget it afterwards - await ModelManagementAuthChecks.can_user_make_model_call( + await _authorize_test_connection( model_params=Deployment( model_name="test_model", litellm_params=LiteLLM_Params(**litellm_params), @@ -2187,6 +2250,9 @@ async def test_model_connection( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, premium_user=premium_user, + llm_router=llm_router, + configured_model_name=configured_model_name, + request_litellm_params=request_litellm_params, ) mode = mode or litellm_params.pop("mode", None) diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index 761cd0685f2..ecf4523fb45 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -4179,3 +4179,129 @@ async def test_health_services_endpoint_pointfive_blocks_non_admin(monkeypatch, assert str(raised.value.code) == "403" logger_class.assert_not_called() + + +def _configured_non_team_deployment(): + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + return Deployment( + model_name="gpt-4o", + litellm_params=LiteLLM_Params( + model="openai/gpt-4o", + api_base="https://configured.invalid/v1", + api_key="CONFIGURED-API-KEY", + ), + model_info=ModelInfo(id="non-team-deployment-id"), + ) + + +def _internal_user(): + return UserAPIKeyAuth( + token="internal-user-token", + user_id="internal-user", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + +@pytest.mark.asyncio +async def test_test_model_connection_allows_internal_user_to_probe_configured_model_they_can_call(): + """ + An internal user may test a configured, non-team model they are allowed + to call, and the probe runs with the configured credentials. + """ + mock_router = MagicMock() + mock_router.get_deployment.return_value = _configured_non_team_deployment() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.auth.auth_checks.can_key_call_model", AsyncMock(return_value=True)) as key_check, + patch("litellm.proxy.auth.auth_checks.get_user_object", AsyncMock(return_value=None)), + patch("litellm.proxy.auth.auth_checks.can_user_call_model", AsyncMock(return_value=True)) as user_check, + patch("litellm.ahealth_check", AsyncMock(return_value={"status": "healthy"})) as health_check, + ): + result = await health_test_model_connection( + request=MagicMock(), + mode="chat", + litellm_params={"model": "openai/gpt-4o"}, + model_info={"id": "non-team-deployment-id"}, + user_api_key_dict=_internal_user(), + ) + + assert result["status"] == "success" + assert key_check.await_args.kwargs["model"] == "gpt-4o" + assert user_check.await_args.kwargs["model"] == "gpt-4o" + assert health_check.await_args.kwargs["model_params"]["api_key"] == "CONFIGURED-API-KEY" + + +@pytest.mark.asyncio +async def test_test_model_connection_denies_internal_user_without_model_access(): + from fastapi import HTTPException + + from litellm.proxy._types import ProxyErrorTypes, ProxyException + + mock_router = MagicMock() + mock_router.get_deployment.return_value = _configured_non_team_deployment() + denied = ProxyException( + message="Key not allowed to access model. Tried to access gpt-4o", + type=ProxyErrorTypes.key_model_access_denied, + param="model", + code=403, + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.auth.auth_checks.can_key_call_model", AsyncMock(side_effect=denied)), + patch("litellm.ahealth_check", AsyncMock()) as health_check, + pytest.raises(HTTPException) as exc_info, + ): + await health_test_model_connection( + request=MagicMock(), + mode="chat", + litellm_params={"model": "openai/gpt-4o"}, + model_info={"id": "non-team-deployment-id"}, + user_api_key_dict=_internal_user(), + ) + + assert exc_info.value.status_code == 403 + assert "not allowed to access model" in exc_info.value.detail["error"] + health_check.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_test_model_connection_keeps_connection_overrides_admin_only(): + from fastapi import HTTPException + + """ + A request that sets its own connection fields describes a different + endpoint than the configured one; probing that stays a management + operation, so an internal user is denied before any access check runs. + """ + mock_router = MagicMock() + mock_router.get_deployment.return_value = _configured_non_team_deployment() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.auth.auth_checks.can_key_call_model", AsyncMock(return_value=True)) as key_check, + patch("litellm.ahealth_check", AsyncMock()) as health_check, + pytest.raises(HTTPException) as exc_info, + ): + await health_test_model_connection( + request=MagicMock(), + mode="chat", + litellm_params={ + "model": "openai/gpt-4o", + "api_base": "https://somewhere-else.invalid/v1", + }, + model_info={"id": "non-team-deployment-id"}, + user_api_key_dict=_internal_user(), + ) + + assert exc_info.value.status_code == 403 + key_check.assert_not_awaited() + health_check.assert_not_awaited()