mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge 03ee545366 into 5028f9ec59
This commit is contained in:
commit
79239495af
2 changed files with 383 additions and 7 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue