diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 64fd59bbe44..cfa2ea71dc3 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -8,7 +8,7 @@ import time import traceback from collections.abc import Iterable, Mapping from datetime import datetime, timedelta, timezone -from typing import Any, Final, Literal, TypedDict, cast +from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, cast import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, Response, status @@ -67,6 +67,10 @@ from litellm.proxy.middleware.in_flight_requests_middleware import ( ) from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager from litellm.router import Router + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + from litellm.types.router import Deployment from litellm.router_utils.clientside_credential_handler import ( _ADMIN_CONFIG_FIELDS_TO_CLEAR_ON_BASE_OVERRIDE, # pyright: ignore[reportPrivateUsage] # one canonical list, shared with the router path clientside_credential_keys, @@ -2000,6 +2004,130 @@ async def health_liveliness_options(): return Response(headers=response_headers, status_code=200) +_NON_ADMIN_TEST_CONNECTION_RESULT_KEYS: Final[frozenset[str]] = frozenset(("model", "error", "mode_error")) + + +def _test_connection_result_for_display( + litellm_params: Mapping[str, object], result: Mapping[str, object], *, outcome_only: bool +) -> Mapping[str, object]: + cleaned: Final = _clean_endpoint_data({**litellm_params, **result}, details=True) + if not outcome_only: + return cleaned + return { # mutable-ok: fresh filtered copy handed to the caller + k: v for k, v in cleaned.items() if k in _NON_ADMIN_TEST_CONNECTION_RESULT_KEYS + } + + +def _configured_probe_mode(model_info: Mapping[str, object] | None, model: object) -> str | None: + configured: Final = model_info.get("mode") if model_info else None + if configured is not None: + return str(configured) + cost_entry: Final = litellm.model_cost.get(model) if isinstance(model, str) else None + return cost_entry.get("mode") if cost_entry else None + + +def _probe_is_configured_deployment( + *, + model_params: "Deployment", + configured_litellm_params: Mapping[str, object], + configured_mode: str | None, + request_litellm_params: Mapping[str, object], + requested_mode: str | None, +) -> bool: + if getattr(model_params.model_info, "team_id", None) is not None: + return False + if any(configured_litellm_params.get(key) != value for key, value in request_litellm_params.items()): + return False + return requested_mode is None or requested_mode == configured_mode + + +async def _assert_caller_can_call_model( + *, + model: str, + user_api_key_dict: UserAPIKeyAuth, + prisma_client: "PrismaClient", + llm_router: "Router | None", +) -> None: + from litellm.proxy.auth.auth_checks import ( + UserNotFoundError, + can_key_call_model, + can_user_call_model, + get_user_object, + ) + from litellm.proxy.proxy_server import llm_model_list, user_api_key_cache + + try: + await can_key_call_model( + model=model, + llm_model_list=llm_model_list, + valid_token=user_api_key_dict, + llm_router=llm_router, + ) + try: + 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, + ) + except UserNotFoundError: + user_object = None + await can_user_call_model(model=model, llm_router=llm_router, user_object=user_object) + except ProxyException as e: + raise HTTPException( + status_code=403, + detail={"error": str(e.message)}, # mutable-ok: same 403 payload shape as the rest of this endpoint + ) from e + + +async def _authorize_test_connection( + *, + model_params: "Deployment", + user_api_key_dict: UserAPIKeyAuth, + prisma_client: "PrismaClient", + premium_user: bool, + llm_router: "Router | None", + configured_model_name: str | None, + configured_litellm_params: Mapping[str, object], + configured_mode: str | None, + request_litellm_params: Mapping[str, object], + requested_mode: str | None, +) -> bool: + """Returns True when a non-admin was admitted to probe a team-less deployment as configured.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelManagementAuthChecks, + ) + + 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 False + except HTTPException as management_denial: + if ( + management_denial.status_code != 403 + or configured_model_name is None + or not _probe_is_configured_deployment( + model_params=model_params, + configured_litellm_params=configured_litellm_params, + configured_mode=configured_mode, + request_litellm_params=request_litellm_params, + requested_mode=requested_mode, + ) + ): + raise + await _assert_caller_can_call_model( + model=configured_model_name, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + llm_router=llm_router, + ) + return True + + @router.post( "/health/test_connection", tags=["health"], @@ -2086,9 +2214,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, @@ -2118,6 +2243,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 @@ -2137,6 +2263,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. @@ -2159,6 +2286,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 @@ -2181,7 +2309,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( + admitted_as_caller: Final = await _authorize_test_connection( model_params=Deployment( model_name="test_model", litellm_params=LiteLLM_Params(**litellm_params), @@ -2190,6 +2318,12 @@ 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, + configured_litellm_params=config_litellm_params, + configured_mode=_configured_probe_mode(loaded_model_info, config_litellm_params.get("model")), + request_litellm_params=request_litellm_params, + requested_mode=mode or request_litellm_params.get("mode"), ) mode = mode or litellm_params.pop("mode", None) @@ -2204,7 +2338,9 @@ async def test_model_connection( ) # Clean the result for display - cleaned_result: Final = _clean_endpoint_data({**litellm_params, **result}, details=True) + cleaned_result: Final = _test_connection_result_for_display( + litellm_params, result, outcome_only=admitted_as_caller + ) return { "status": "error" if "error" in result else "success", 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..fd37cfcff2f 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -3,7 +3,7 @@ import copy import json import time from collections.abc import Iterator, Mapping, Sequence -from contextlib import contextmanager +from contextlib import ExitStack, contextmanager from datetime import datetime, timedelta from types import SimpleNamespace from typing import Final @@ -33,6 +33,7 @@ from litellm.proxy.health_endpoints._health_endpoints import ( from litellm.proxy.health_endpoints._health_endpoints import ( test_model_connection as health_test_model_connection, ) +from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo # Import shared proxy test helpers from conftest from tests.test_litellm.proxy.conftest import create_proxy_test_client @@ -4179,3 +4180,242 @@ 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(mode: str | None = "chat") -> Deployment: + 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", mode=mode), + ) + + +def _internal_user() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + token="internal-user-token", + user_id="internal-user", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + +@contextmanager +def _probe_environment( + deployment: Deployment, + *, + key_check: AsyncMock | None = None, + user_lookup: AsyncMock | None = None, + user_check: AsyncMock | None = None, + health_check: AsyncMock | None = None, +) -> Iterator[None]: + mock_router = MagicMock() + mock_router.get_deployment.return_value = deployment + with ExitStack() as stack: + stack.enter_context( + patch( # test-quality-ok: the endpoint reads proxy_server globals, stubbed like the rest of this file + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ) + ) + stack.enter_context( + patch( # test-quality-ok: the endpoint reads proxy_server globals, stubbed like the rest of this file + "litellm.proxy.proxy_server.llm_router", mock_router + ) + ) + stack.enter_context( + patch( # test-quality-ok: the endpoint reads proxy_server globals, stubbed like the rest of this file + "litellm.proxy.proxy_server.premium_user", True + ) + ) + if key_check is not None: + stack.enter_context( + patch( # test-quality-ok: the access check needs a live router and DB, stubbed like the rest of this file + "litellm.proxy.auth.auth_checks.can_key_call_model", key_check + ) + ) + if user_lookup is not None: + stack.enter_context( + patch( # test-quality-ok: the user lookup needs a live DB, stubbed like the rest of this file + "litellm.proxy.auth.auth_checks.get_user_object", user_lookup + ) + ) + if user_check is not None: + stack.enter_context( + patch( # test-quality-ok: the access check needs a live router and DB, stubbed like the rest of this file + "litellm.proxy.auth.auth_checks.can_user_call_model", user_check + ) + ) + if health_check is not None: + stack.enter_context( + patch( # test-quality-ok: the probe is a real provider call, stubbed like the rest of this file + "litellm.ahealth_check", health_check + ) + ) + yield + + +async def _probe_as_internal_user( + litellm_params: dict[str, object], mode: str | None = "chat", deployment_id: str = "non-team-deployment-id" +) -> dict[str, object]: + return await health_test_model_connection( + request=MagicMock(), + mode=mode, + litellm_params=litellm_params, + model_info={"id": deployment_id}, + user_api_key_dict=_internal_user(), + ) + + +@pytest.mark.asyncio +async def test_test_model_connection_allows_internal_user_to_probe_configured_model_they_can_call(): + key_check = AsyncMock(return_value=True) + user_check = AsyncMock(return_value=True) + health_check = AsyncMock(return_value={"status": "healthy"}) + + with _probe_environment( + _configured_non_team_deployment(), + key_check=key_check, + user_lookup=AsyncMock(return_value=None), + user_check=user_check, + health_check=health_check, + ): + result = await _probe_as_internal_user({"model": "openai/gpt-4o"}) + + 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" + assert "api_base" not in result["result"] + assert "api_key" not in result["result"] + + +@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 + + denied = ProxyException( + message="Key not allowed to access model. Tried to access gpt-4o", + type=ProxyErrorTypes.key_model_access_denied, + param="model", + code=403, + ) + health_check = AsyncMock() + + with ( + _probe_environment( + _configured_non_team_deployment(), key_check=AsyncMock(side_effect=denied), health_check=health_check + ), + pytest.raises(HTTPException) as exc_info, + ): + await _probe_as_internal_user({"model": "openai/gpt-4o"}) + + 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 +@pytest.mark.parametrize( + "litellm_params, mode", + [ + ({"model": "openai/gpt-4o", "api_base": "https://somewhere-else.invalid/v1"}, "chat"), + ({"model": "openai/gpt-4o-mini"}, "chat"), + ({"model": "openai/gpt-4o"}, "image_generation"), + ], + ids=["api_base", "model", "mode"], +) +async def test_test_model_connection_keeps_overrides_admin_only(litellm_params: dict[str, object], mode: str) -> None: + from fastapi import HTTPException + + key_check = AsyncMock(return_value=True) + health_check = AsyncMock() + + with ( + _probe_environment(_configured_non_team_deployment(), key_check=key_check, health_check=health_check), + pytest.raises(HTTPException) as exc_info, + ): + await _probe_as_internal_user(litellm_params, mode=mode) + + assert exc_info.value.status_code == 403 + key_check.assert_not_awaited() + health_check.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_test_model_connection_keeps_team_deployments_admin_only_for_non_admins(): + from fastapi import HTTPException + + from litellm.proxy._types import LiteLLM_TeamTable + + team_deployment = Deployment( + model_name="gpt-4o", + litellm_params=LiteLLM_Params(model="openai/gpt-4o", api_key="TEAM-A-API-KEY"), + model_info=ModelInfo(id="team-a-deployment-id", team_id="team-a"), + ) + team_row = SimpleNamespace( + model_dump=lambda: LiteLLM_TeamTable(team_id="team-a", members_with_roles=[]).model_dump() + ) + key_check = AsyncMock(return_value=True) + health_check = AsyncMock() + + with ( + _probe_environment(team_deployment, key_check=key_check, health_check=health_check), + patch( # test-quality-ok: the team lookup needs a live DB, stubbed like the surrounding team tests + "litellm.proxy.management_endpoints.model_management_endpoints.TeamRepository" + ) as MockTeamRepo, + ): + repo = MagicMock() + repo.table.find_unique = AsyncMock(return_value=team_row) + MockTeamRepo.return_value = repo + with pytest.raises(HTTPException) as exc_info: + await _probe_as_internal_user({"model": "openai/gpt-4o"}, deployment_id="team-a-deployment-id") + + assert exc_info.value.status_code == 403 + key_check.assert_not_awaited() + health_check.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_test_model_connection_tolerates_missing_user_record_for_non_admin(): + from litellm.proxy.auth.auth_checks import UserNotFoundError + + user_check = AsyncMock(return_value=True) + + with _probe_environment( + _configured_non_team_deployment(), + key_check=AsyncMock(return_value=True), + user_lookup=AsyncMock(side_effect=UserNotFoundError("gone")), + user_check=user_check, + health_check=AsyncMock(return_value={"status": "healthy"}), + ): + result = await _probe_as_internal_user({"model": "openai/gpt-4o"}) + + assert result["status"] == "success" + assert user_check.await_args.kwargs["user_object"] is None + + +@pytest.mark.asyncio +async def test_test_model_connection_accepts_mode_the_probe_would_infer() -> None: + + deployment = Deployment( + model_name="gpt-4o", + litellm_params=LiteLLM_Params(model="gpt-4o", api_key="CONFIGURED-API-KEY"), + model_info=ModelInfo(id="non-team-deployment-id"), + ) + health_check = AsyncMock(return_value={"status": "healthy"}) + + with _probe_environment( + deployment, + key_check=AsyncMock(return_value=True), + user_lookup=AsyncMock(return_value=None), + user_check=AsyncMock(return_value=True), + health_check=health_check, + ): + result = await _probe_as_internal_user({"model": "gpt-4o"}) + + assert result["status"] == "success" + assert health_check.await_args.kwargs["mode"] == "chat"