mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(proxy): let internal users test the connection of models they can call
/health/test_connection authorized every request through ModelManagementAuthChecks.can_user_make_model_call, the check used for adding, updating and deleting models. For a model without a team_id that means proxy admins only, so the "Test Connection" button the UI shows to every user failed with 403 for internal users. Testing a connection is a read against a configured deployment, so use the same permission as calling the model: after the management check denies a non-team model, allow the request when the deployment was resolved from the router, the request does not override any connection field (a request that sets api_base, api_key, ... describes a different endpoint and stays a management operation), and the caller's key and user may call the deployment's model_name. Fixes #40265
This commit is contained in:
parent
c2c2a623c0
commit
232581c6c0
2 changed files with 196 additions and 4 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue