Merge pull request #26912 from stuxf/codex/auth-sensitive-routes

chore(proxy): guard sensitive public endpoints
This commit is contained in:
yuneng-jiang 2026-05-04 17:04:10 -07:00 • committed by GitHub
commit de7175d6ab
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 419 additions and 114 deletions

View file

@ -166,7 +166,7 @@ langfuse_default_tags: Optional[List[str]] = None
langsmith_batch_size: Optional[int] = None
prometheus_initialize_budget_metrics: Optional[bool] = False
prometheus_latency_buckets: Optional[List[float]] = None
require_auth_for_metrics_endpoint: Optional[bool] = False
require_auth_for_metrics_endpoint: Optional[bool] = True
argilla_batch_size: Optional[int] = None
datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload.
gcs_pub_sub_use_v1: Optional[bool] = (

View file

@ -617,13 +617,12 @@ class LiteLLMRoutes(enum.Enum):
"/",
"/health/liveliness",
"/health/liveness",
"/health/readiness",
"/test",
"/config/yaml",
"/metrics",
"/litellm/.well-known/litellm-ui-config",
"/.well-known/litellm-ui-config",
"/public/model_hub",
"/public/model_hub/info",
"/public/agent_hub",
"/public/mcp_hub",
"/public/skill_hub",

View file

@ -87,6 +87,23 @@ except ImportError as e:
user_api_key_service_logger_obj = ServiceLogging() # used for tracking latency on OTEL
def _normalize_public_auth_route(route: str) -> str:
if route != "/" and route.endswith("/"):
return route.rstrip("/")
return route
def _route_requires_auth_despite_public(
route: str, general_settings: Optional[dict]
) -> bool:
normalized_route = _normalize_public_auth_route(route)
if normalized_route == "/metrics":
return litellm.require_auth_for_metrics_endpoint is not False
return False
custom_litellm_key_header = APIKeyHeader(
name=SpecialHeaders.custom_litellm_api_key.value,
auto_error=False,
@ -714,7 +731,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
"""
######## Route Checks Before Reading DB / Cache for "token" ################
if (
if not _route_requires_auth_despite_public(
route=route, general_settings=general_settings
) and (
route in LiteLLMRoutes.public_routes.value # type: ignore
or route_in_additonal_public_routes(current_route=route)
):
@ -1698,7 +1717,7 @@ async def _run_centralized_common_checks(
user_custom_auth,
)
# Public routes (e.g. /health/readiness, /metrics) are exempt from
# Public routes (e.g. /health/liveness) are exempt from
# auth in the builder — the wrapper must not retroactively apply
# authz on top, or k8s readiness probes and other unauthenticated
# callers get 401.

View file

@ -50,7 +50,10 @@ def configure_gc_thresholds():
configure_gc_thresholds()
@router.get("/debug/asyncio-tasks")
@router.get(
"/debug/asyncio-tasks",
dependencies=[Depends(user_api_key_auth)],
)
async def get_active_tasks_stats():
"""
Returns:
@ -103,7 +106,11 @@ if os.environ.get("LITELLM_PROFILE", "false").lower() == "true":
tracemalloc.start(10)
@router.get("/memory-usage", include_in_schema=False)
@router.get(
"/memory-usage",
dependencies=[Depends(user_api_key_auth)],
include_in_schema=False,
)
async def memory_usage():
# Take a snapshot of the current memory usage
snapshot = tracemalloc.take_snapshot()
@ -711,7 +718,11 @@ async def configure_gc_thresholds_endpoint(
}
@router.get("/otel-spans", include_in_schema=False)
@router.get(
"/otel-spans",
dependencies=[Depends(user_api_key_auth)],
include_in_schema=False,
)
async def get_otel_spans():
from litellm.proxy.proxy_server import open_telemetry_logger

View file

@ -1447,14 +1447,11 @@ def callback_name(callback):
return str(callback)
@router.get(
"/health/readiness",
tags=["health"],
dependencies=[Depends(user_api_key_auth)],
)
async def health_readiness(response: Response):
async def _get_health_readiness_details(
response: Optional[Response] = None,
) -> Dict[str, Any]:
"""
Unprotected endpoint for checking if worker can receive requests
Detailed health payload for authenticated diagnostics.
"""
from litellm.proxy.proxy_server import prisma_client, version
@ -1473,7 +1470,7 @@ async def health_readiness(response: Response):
success_callback_names = litellm.success_callback
# check Cache
cache_type = None
cache_type: Any = None
if litellm.cache is not None:
from litellm.caching.caching import RedisSemanticCache
@ -1482,6 +1479,7 @@ async def health_readiness(response: Response):
if isinstance(litellm.cache.cache, RedisSemanticCache):
# ping the cache
# TODO: @ishaan-jaff - we should probably not ping the cache on every /health/readiness check
index_info: Any
try:
index_info = await litellm.cache.cache._index_info()
except Exception as e:
@ -1499,7 +1497,7 @@ async def health_readiness(response: Response):
# serve requests that depend on persisted state (keys, budgets,
# spend logs). Return 503 so orchestrators take this pod out of
# rotation; "Not connected" (no DB configured at all) stays 200.
if db_health_status["status"] != "connected":
if response is not None and db_health_status["status"] != "connected":
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return {
"status": "healthy",
@ -1526,6 +1524,52 @@ async def health_readiness(response: Response):
raise HTTPException(status_code=503, detail=f"Service Unhealthy ({str(e)})")
def _allow_public_health_readiness_details() -> bool:
from litellm.proxy.proxy_server import general_settings
return general_settings.get("allow_public_health_readiness_details") is True
async def _set_public_readiness_status(response: Response) -> None:
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return
db_health_status = await _db_health_readiness_check()
if db_health_status["status"] != "connected":
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
@router.get(
"/health/readiness",
tags=["health"],
)
async def health_readiness(response: Response):
"""
Public readiness probe. Keep this low-detail for unauthenticated load
balancers by default. Admins can opt into the legacy detailed public
payload with general_settings.allow_public_health_readiness_details.
"""
if _allow_public_health_readiness_details():
return await _get_health_readiness_details(response=response)
await _set_public_readiness_status(response=response)
return {"status": "healthy"}
@router.get(
"/health/readiness/details",
tags=["health"],
dependencies=[Depends(user_api_key_auth)],
)
async def health_readiness_details(response: Response):
"""
Authenticated readiness diagnostics with DB/cache/callback metadata.
"""
return await _get_health_readiness_details(response=response)
@router.get(
"/health/backlog",
tags=["health"],
@ -1561,7 +1605,6 @@ async def health_liveliness():
@router.options(
"/health/readiness",
tags=["health"],
dependencies=[Depends(user_api_key_auth)],
)
async def health_readiness_options():
"""

View file

@ -20,13 +20,13 @@ class PrometheusAuthMiddleware:
"""
Middleware to authenticate requests to the metrics endpoint.
By default, auth is not run on the metrics endpoint.
By default, auth is run on the metrics endpoint.
Enabled by setting the following in proxy_config.yaml:
To allow unauthenticated metrics in proxy_config.yaml:
```yaml
litellm_settings:
require_auth_for_metrics_endpoint: true
require_auth_for_metrics_endpoint: false
```
"""
@ -39,8 +39,8 @@ class PrometheusAuthMiddleware:
await self.app(scope, receive, send)
return
# Only run auth if configured to do so
if litellm.require_auth_for_metrics_endpoint is True:
# Run auth by default; allow legacy public metrics only when explicitly disabled.
if litellm.require_auth_for_metrics_endpoint is not False:
# user_api_key_auth reads the request body, which consumes ASGI `receive`.
# Buffer those messages and replay them for the inner app; otherwise a
# successful auth would forward an exhausted receive and /metrics hangs.
@ -52,10 +52,29 @@ class PrometheusAuthMiddleware:
return message
request = Request(scope, receive_for_auth)
api_key = request.headers.get(_AUTHORIZATION_HEADER) or ""
try:
await user_api_key_auth(request=request, api_key=api_key)
await user_api_key_auth(
request=request,
api_key=request.headers.get(_AUTHORIZATION_HEADER) or "",
azure_api_key_header=request.headers.get(
SpecialHeaders.azure_authorization.value
)
or "",
anthropic_api_key_header=request.headers.get(
SpecialHeaders.anthropic_authorization.value
),
google_ai_studio_api_key_header=request.headers.get(
SpecialHeaders.google_ai_studio_authorization.value
),
azure_apim_header=request.headers.get(
SpecialHeaders.azure_apim_authorization.value
)
or "",
custom_litellm_key_header=request.headers.get(
SpecialHeaders.custom_litellm_api_key.value
),
)
except Exception as e:
# Send 401 response directly via ASGI protocol
error_message = getattr(e, "message", str(e))

View file

@ -5,7 +5,7 @@ from importlib.resources import files
from typing import Any, Dict, List, Optional
import litellm
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, HTTPException
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.get_blog_posts import (
@ -14,8 +14,9 @@ from litellm.litellm_core_utils.get_blog_posts import (
GetBlogPosts,
get_blog_posts,
)
from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy._types import (
CommonProxyErrors,
)
from litellm.types.agents import AgentCard
from litellm.types.mcp import MCPPublicServer
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
@ -31,6 +32,7 @@ from litellm.types.utils import LlmProviders
router = APIRouter()
# ---------------------------------------------------------------------------
# /public/endpoints — helpers
# ---------------------------------------------------------------------------
@ -153,7 +155,6 @@ def _load_endpoints() -> List[Dict[str, Any]]:
@router.get(
"/public/model_hub",
tags=["public", "model management"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[ModelGroupInfoProxy],
)
async def public_model_hub():
@ -208,7 +209,6 @@ async def public_model_hub():
@router.get(
"/public/agent_hub",
tags=["[beta] Agents", "public"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[AgentCard],
)
async def get_agents():
@ -230,7 +230,6 @@ async def get_agents():
@router.get(
"/public/mcp_hub",
tags=["[beta] MCP", "public"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[MCPPublicServer],
)
async def get_mcp_servers():

View file

@ -3079,7 +3079,11 @@ async def global_spend_models(
return response
@router.get("/provider/budgets", response_model=ProviderBudgetResponse)
@router.get(
"/provider/budgets",
dependencies=[Depends(user_api_key_auth)],
response_model=ProviderBudgetResponse,
)
async def provider_budgets() -> ProviderBudgetResponse:
"""
Provider Budget Routing - Get Budget, Spend Details https://docs.litellm.ai/docs/proxy/provider_budget_routing

View file

@ -99,6 +99,11 @@ class UISettings(BaseModel):
description="If true, requires authentication for accessing the public AI Hub.",
)
allow_public_health_readiness_details: bool = Field(
default=False,
description="If true, returns the legacy detailed payload from the unauthenticated /health/readiness endpoint.",
)
forward_client_headers_to_llm_api: bool = Field(
default=False,
description=(
@ -169,6 +174,7 @@ ALLOWED_UI_SETTINGS_FIELDS = {
"disable_team_admin_delete_team_user",
"enabled_ui_pages_internal_users",
"require_auth_for_public_ai_hub",
"allow_public_health_readiness_details",
"forward_client_headers_to_llm_api",
"forward_llm_provider_auth_headers",
"disable_agents_for_internal_users",
@ -183,6 +189,7 @@ ALLOWED_UI_SETTINGS_FIELDS = {
# Flags that must be synced from the persisted UISettings into
# general_settings at runtime (on both read and write).
_RUNTIME_GENERAL_SETTINGS_FLAGS = [
"allow_public_health_readiness_details",
"forward_client_headers_to_llm_api",
"forward_llm_provider_auth_headers",
"disable_agents_for_internal_users",

View file

@ -128,8 +128,8 @@ async def get_spend_info(session, entity_type: str, entity_id: str):
async def get_proxy_readiness(session):
"""Fetch /health/readiness. Used both as a fail-fast gate and as a diagnostic on poll timeout."""
url = "http://0.0.0.0:4000/health/readiness"
"""Fetch authenticated readiness details. Used both as a fail-fast gate and as a diagnostic on poll timeout."""
url = "http://0.0.0.0:4000/health/readiness/details"
headers = {"Authorization": "Bearer sk-1234"}
async with session.get(url, headers=headers) as response:
return response.status, await response.json()
@ -140,7 +140,7 @@ async def assert_proxy_healthy(session):
status, body = await get_proxy_readiness(session)
if status != 200 or body.get("db") != "connected":
pytest.fail(
f"Proxy /health/readiness unhealthy (status={status}). "
f"Proxy /health/readiness/details unhealthy (status={status}). "
f"Cannot run spend accuracy test. Response: {body}"
)
print(f"Proxy readiness OK: {body}")

View file

@ -73,13 +73,32 @@ async def test_health_readiness():
response_json = await response.json()
print(response_json)
assert "litellm_version" in response_json
assert "status" in response_json
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
@pytest.mark.asyncio
async def test_health_readiness_details():
"""
Check if authenticated readiness diagnostics expose version metadata.
"""
async with aiohttp.ClientSession() as session:
url = "http://0.0.0.0:4000/health/readiness/details"
headers = {"Authorization": "Bearer sk-1234"}
async with session.get(url, headers=headers) as response:
status = response.status
response_json = await response.json()
print(response_json)
assert "status" in response_json
assert "litellm_version" in response_json
if status != 200:
raise Exception(f"Request did not return a 200 status code: {status}")
@pytest.mark.asyncio
async def test_health_liveliness():
"""

View file

@ -10,9 +10,11 @@ sys.path.insert(
import pytest
import litellm
import litellm.proxy.proxy_server
from litellm.caching.dual_cache import DualCache
from litellm.proxy._types import (
LiteLLMRoutes,
LiteLLM_JWTAuth,
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
@ -27,6 +29,7 @@ from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import (
_route_requires_auth_despite_public,
_reserve_budget_after_common_checks,
_run_centralized_common_checks,
_run_post_custom_auth_checks,
@ -59,6 +62,29 @@ def test_get_api_key():
) == (api_key, passed_in_key)
def test_route_requires_auth_despite_public_for_metrics(monkeypatch):
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
assert _route_requires_auth_despite_public("/metrics", {}) is True
assert _route_requires_auth_despite_public("/metrics/", {}) is True
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", False)
assert _route_requires_auth_despite_public("/metrics", {}) is False
def test_public_ai_hub_routes_remain_public():
for route in (
"/public/model_hub",
"/public/model_hub/info",
"/public/agent_hub",
"/public/mcp_hub",
"/public/skill_hub",
):
assert route in LiteLLMRoutes.public_routes.value
assert _route_requires_auth_despite_public(route, {}) is False
@pytest.mark.asyncio
async def test_should_clear_stale_budget_reservation_when_budget_checks_skip():
user_api_key_auth_obj = UserAPIKeyAuth(
@ -2352,18 +2378,18 @@ async def test_centralized_common_checks_short_circuits_when_master_key_unset():
@pytest.mark.asyncio
async def test_centralized_common_checks_skips_public_routes():
"""Regression: public routes (e.g. /health/readiness) are exempted
"""Regression: public routes (e.g. /health/liveness) are exempted
by the builder fast-path. The wrapper must not retroactively run
common_checks on top — the synthetic INTERNAL_USER_VIEW_ONLY token
has no user_id, so common_checks would reject the request as
admin-only. Breaks k8s readiness probes when master_key is set."""
admin-only."""
import litellm.proxy.proxy_server as _proxy_server_mod
from fastapi import Request
from starlette.datastructures import URL
token = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER_VIEW_ONLY)
request = Request(scope={"type": "http"})
request._url = URL(url="/health/readiness")
request._url = URL(url="/health/liveness")
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None)
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
@ -2378,7 +2404,7 @@ async def test_centralized_common_checks_skips_public_routes():
user_api_key_auth_obj=token,
request=request,
request_data={},
route="/health/readiness",
route="/health/liveness",
)
mock_checks.assert_not_awaited()
finally:

View file

@ -11,10 +11,14 @@ sys.path.insert(
import httpx
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
from prisma.errors import ClientNotConnectedError, HTTPClientClosedError, PrismaError
import litellm.proxy.health_endpoints._health_endpoints as _health_endpoints_module
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.health_endpoints._health_endpoints import (
_db_health_readiness_check,
get_callback_identifier,
@ -512,7 +516,7 @@ def proxy_client(monkeypatch):
Redis cache:
- If REDIS_HOST is set in environment, Redis cache will be automatically configured
- Cache configuration is included in /health/readiness endpoint response
- Cache diagnostics are included in the authenticated /health/readiness/details response
"""
client = create_proxy_test_client(monkeypatch)
with client:
@ -588,11 +592,7 @@ def test_health_liveness_endpoint(proxy_client):
def test_health_readiness(proxy_client):
"""
Test /health/readiness endpoint.
Database and Redis are optional - the endpoint should work whether they're available or not.
If DATABASE_URL is set, the endpoint will check database connectivity.
If REDIS_HOST is set, the endpoint will report cache status.
If neither is set, the endpoint should still return a valid health status.
Database and Redis are optional - the public endpoint should work whether they're available or not.
"""
# Measure the time taken for the health check call
start_time = time.perf_counter()
@ -614,40 +614,57 @@ def test_health_readiness(proxy_client):
duration_ms < 500
), f"Health check took {duration_ms:.2f}ms, expected < 500ms for readiness endpoint"
# Assert response contains expected fields
# Assert response contains only low-detail public probe fields
response_data = response.json()
assert "status" in response_data, "Response should contain 'status' field"
assert (
"litellm_version" in response_data
), "Response should contain 'litellm_version' field"
# Display all health endpoint response fields (matches what /health/readiness returns)
print("\n" + "-" * 60)
print("HEALTH ENDPOINT RESPONSE")
print("-" * 60)
print(f"Status: {response_data.get('status', 'unknown')}")
print(f"Database: {response_data.get('db', 'not reported')}")
print(f"LiteLLM Version: {response_data.get('litellm_version', 'unknown')}")
print(f"Success Callbacks: {response_data.get('success_callbacks', [])}")
print(f"Cache: {response_data.get('cache', 'none')}")
print(
f"Use AioHTTP Transport: {response_data.get('use_aiohttp_transport', 'unknown')}"
)
assert response_data == {"status": "healthy"}
print(f"Response time: {duration_ms:.2f}ms")
# If database status is reported, verify it's a valid status
# Database may be "connected", "disconnected", "unknown", or "Not connected" (when prisma_client is None)
if "db" in response_data:
db_status = response_data["db"]
# Database status can be any of these valid states
assert db_status in [
"connected",
"disconnected",
"unknown",
"Not connected",
], f"Unexpected db status: {db_status}"
print("=" * 60 + "\n")
def test_health_readiness_details_returns_diagnostic_fields(monkeypatch):
"""
Detailed readiness diagnostics stay available behind the auth dependency.
"""
app = FastAPI()
app.include_router(_health_endpoints_module.router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN
)
client = TestClient(app)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
response = client.get("/health/readiness/details")
assert response.status_code == 200, response.text
response_data = response.json()
assert response_data["status"] == "healthy"
assert "litellm_version" in response_data
assert "success_callbacks" in response_data
assert "cache" in response_data
def test_health_readiness_allows_explicit_legacy_public_details(monkeypatch):
"""
Operators can explicitly preserve the legacy public readiness payload.
"""
app = FastAPI()
app.include_router(_health_endpoints_module.router)
client = TestClient(app)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"allow_public_health_readiness_details": True},
)
response = client.get("/health/readiness")
assert response.status_code == 200, response.text
response_data = response.json()
assert response_data["status"] == "healthy"
assert "litellm_version" in response_data
assert "success_callbacks" in response_data
assert "cache" in response_data
def test_get_callback_identifier_string_and_object_with_callback_name():
@ -1503,8 +1520,7 @@ async def test_health_readiness_returns_503_when_db_disconnected():
result = await health_readiness(response=response)
assert response.status_code == 503
assert result["db"] == "disconnected"
assert result["status"] == "healthy" # body shape unchanged for back-compat
assert result == {"status": "healthy"}
@pytest.mark.asyncio
@ -1527,7 +1543,7 @@ async def test_health_readiness_returns_200_when_db_connected():
result = await health_readiness(response=response)
assert response.status_code == 200
assert result["db"] == "connected"
assert result == {"status": "healthy"}
@pytest.mark.asyncio
@ -1546,7 +1562,7 @@ async def test_health_readiness_returns_200_when_no_db_configured():
result = await health_readiness(response=response)
assert response.status_code == 200
assert result["db"] == "Not connected"
assert result == {"status": "healthy"}
def test_clean_endpoint_data_strips_credentials_keeps_routing_fields():

View file

@ -1,18 +1,5 @@
import json
import os
import sys
import pytest
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import pytest
from fastapi import FastAPI
from fastapi.responses import JSONResponse
from fastapi.testclient import TestClient
import litellm
@ -21,7 +8,7 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi
# Fake auth functions to simulate valid and invalid auth behavior.
async def fake_valid_auth(request, api_key):
async def fake_valid_auth(request, api_key, **kwargs):
# Simulate valid authentication: do nothing (i.e. pass)
return
@ -35,15 +22,11 @@ async def fake_valid_auth_reads_body(request, api_key, **kwargs):
return
async def fake_invalid_auth(request, api_key):
print("running fake invalid auth", request, api_key)
async def fake_invalid_auth(request, api_key, **kwargs):
# Simulate invalid auth by raising an exception.
raise Exception("Invalid API key")
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@pytest.fixture
def app_with_middleware():
"""Create a FastAPI app with the PrometheusAuthMiddleware and dummy endpoints."""
@ -98,7 +81,7 @@ def test_valid_auth_metrics(app_with_middleware, monkeypatch):
Test that a request to /metrics (and /metrics/) with valid auth headers passes.
"""
# Enable auth on metrics endpoints.
litellm.require_auth_for_metrics_endpoint = True
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
# Patch the auth function to simulate a valid authentication.
monkeypatch.setattr(
"litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth",
@ -123,7 +106,7 @@ def test_invalid_auth_metrics(app_with_middleware, monkeypatch):
"""
Test that a request to /metrics with invalid auth headers fails with a 401.
"""
litellm.require_auth_for_metrics_endpoint = True
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
# Patch the auth function to simulate a failed authentication.
monkeypatch.setattr(
"litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth",
@ -138,12 +121,48 @@ def test_invalid_auth_metrics(app_with_middleware, monkeypatch):
assert "Unauthorized access to metrics endpoint" in response.text
def test_metrics_auth_uses_real_auth_when_route_is_public(
app_with_middleware, monkeypatch
):
"""
Regression: /metrics is statically public, but require_auth_for_metrics_endpoint
must still force the real auth path.
"""
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-master")
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
client = TestClient(app_with_middleware)
response = client.get("/metrics")
assert response.status_code == 401, response.text
assert "Unauthorized access to metrics endpoint" in response.text
def test_metrics_auth_is_required_by_default(app_with_middleware, monkeypatch):
"""
Metrics should require auth unless explicitly configured as public.
"""
monkeypatch.setattr(
"litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth",
fake_invalid_auth,
)
client = TestClient(app_with_middleware)
response = client.get("/metrics")
assert response.status_code == 401, response.text
assert "Unauthorized access to metrics endpoint" in response.text
def test_no_auth_metrics_when_disabled(app_with_middleware, monkeypatch):
"""
Test that when require_auth_for_metrics_endpoint is False, requests to /metrics
bypass the auth check.
"""
litellm.require_auth_for_metrics_endpoint = False
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", False)
# To ensure auth is not run, patch the auth function with one that will raise if called.
def should_not_be_called(*args, **kwargs):
@ -160,11 +179,11 @@ def test_no_auth_metrics_when_disabled(app_with_middleware, monkeypatch):
assert response.json() == {"msg": "metrics OK"}
def test_non_metrics_requests_pass_through(app_with_middleware):
def test_non_metrics_requests_pass_through(app_with_middleware, monkeypatch):
"""
Test that non-metrics endpoints pass through the middleware unaffected.
"""
litellm.require_auth_for_metrics_endpoint = True
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
client = TestClient(app_with_middleware)
@ -182,7 +201,7 @@ def test_non_metrics_requests_dont_trigger_auth(app_with_middleware, monkeypatch
Test that non-metrics requests never trigger auth, even when auth is enabled
and the auth function would reject the request.
"""
litellm.require_auth_for_metrics_endpoint = True
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
def should_not_be_called(*args, **kwargs):
raise Exception("Auth should not be called for non-metrics requests")

View file

@ -91,6 +91,19 @@ def test_get_litellm_model_cost_map_returns_cost_map():
)
def test_public_ai_hub_info_is_public_by_default(monkeypatch):
app = FastAPI()
app.include_router(router)
client = TestClient(app)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-master")
response = client.get("/public/model_hub/info")
assert response.status_code == 200, response.text
def test_watsonx_provider_fields():
"""Test that Watsonx provider has all required credential fields including multiple auth options."""
app = FastAPI()
@ -166,9 +179,9 @@ def test_anthropic_provider_fields_support_byok():
"Anthropic api_key must be optional so admins can configure BYOK models "
"without entering a key. See BYOK tutorial."
)
assert fields_by_key["api_key"].get("tooltip"), (
"Anthropic api_key must have a tooltip explaining the BYOK use case."
)
assert fields_by_key["api_key"].get(
"tooltip"
), "Anthropic api_key must have a tooltip explaining the BYOK use case."
assert "api_base" in fields_by_key, (
"Anthropic provider form must expose api_base so cloud customers "
"can override the upstream URL without env var access."
@ -176,16 +189,16 @@ def test_anthropic_provider_fields_support_byok():
api_base_field = fields_by_key["api_base"]
assert api_base_field["required"] is False
assert api_base_field["field_type"] == "text"
assert api_base_field.get("tooltip"), (
"api_base should have a tooltip explaining it is optional."
)
assert api_base_field.get(
"tooltip"
), "api_base should have a tooltip explaining it is optional."
# UI forms render fields in credential_fields order; api_base should come first
# so an admin sees the URL override before the key field.
field_order = [f["key"] for f in anthropic["credential_fields"]]
assert field_order.index("api_base") < field_order.index("api_key"), (
"api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)."
)
assert field_order.index("api_base") < field_order.index(
"api_key"
), "api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)."
def test_public_model_hub_with_healthy_model():

View file

@ -0,0 +1,34 @@
from fastapi.routing import APIRoute
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.debug_utils import router as debug_router
from litellm.proxy.spend_tracking.spend_management_endpoints import (
router as spend_router,
)
def _get_route_dependency_calls(router, path: str, method: str):
for route in router.routes:
if (
isinstance(route, APIRoute)
and route.path == path
and method in route.methods
):
return [dependency.call for dependency in route.dependant.dependencies]
raise AssertionError(f"Route {method} {path} not found")
def test_sensitive_debug_routes_require_auth_dependency():
for path, method in (
("/debug/asyncio-tasks", "GET"),
("/otel-spans", "GET"),
):
assert user_api_key_auth in _get_route_dependency_calls(
debug_router, path, method
)
def test_provider_budgets_requires_auth_dependency():
assert user_api_key_auth in _get_route_dependency_calls(
spend_router, "/provider/budgets", "GET"
)

View file

@ -868,6 +868,7 @@ class TestProxySettingEndpoints:
mock_db_record = MagicMock()
mock_db_record.ui_settings = {
"disable_model_add_for_internal_users": True,
"require_auth_for_public_ai_hub": True,
"unexpected_flag": True,
}
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(
@ -880,10 +881,12 @@ class TestProxySettingEndpoints:
assert response.status_code == 200
data = response.json()
assert data["values"]["disable_model_add_for_internal_users"] is True
assert data["values"]["require_auth_for_public_ai_hub"] is True
assert "unexpected_flag" not in data["values"]
assert (
"disable_model_add_for_internal_users" in data["field_schema"]["properties"]
)
assert "require_auth_for_public_ai_hub" in data["field_schema"]["properties"]
mock_prisma.db.litellm_uisettings.find_unique.assert_called_once_with(
where={"id": "ui_settings"}
)
@ -1070,6 +1073,43 @@ class TestProxySettingEndpoints:
assert "unsupported_flag" not in stored_settings
assert stored_settings["disable_model_add_for_internal_users"] is False
def test_update_ui_settings_preserves_public_ai_hub_auth_flag(
self, mock_auth, monkeypatch
):
"""Public AI Hub auth is an existing UI setting and must remain writable."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
mock_prisma = MagicMock()
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
payload = {"require_auth_for_public_ai_hub": True}
try:
response = client.patch("/update/ui_settings", json=payload)
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
data = response.json()
assert data["status"] == "success"
assert data["settings"]["require_auth_for_public_ai_hub"] is True
call_args = mock_prisma.db.litellm_uisettings.upsert.call_args
stored_settings = json.loads(call_args.kwargs["data"]["create"]["ui_settings"])
assert stored_settings["require_auth_for_public_ai_hub"] is True
def test_update_ui_settings_persists_forward_llm_provider_auth_headers(
self, mock_auth, monkeypatch
):
@ -1147,6 +1187,43 @@ class TestProxySettingEndpoints:
assert response.status_code == 200
assert general_settings.get("forward_llm_provider_auth_headers") is True
def test_update_ui_settings_syncs_public_health_readiness_details_to_general_settings(
self, mock_auth, monkeypatch
):
"""Public readiness details flag must be synced so the health route sees it."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
general_settings: dict = {}
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings", general_settings
)
mock_prisma = MagicMock()
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
payload = {"allow_public_health_readiness_details": True}
try:
response = client.patch("/update/ui_settings", json=payload)
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
assert general_settings.get("allow_public_health_readiness_details") is True
def test_update_ui_settings_persists_and_syncs_disable_key_generate_for_org_admin(
self, mock_auth, monkeypatch
):