mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge pull request #26912 from stuxf/codex/auth-sensitive-routes
chore(proxy): guard sensitive public endpoints
This commit is contained in:
commit
de7175d6ab
17 changed files with 419 additions and 114 deletions
|
|
@ -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] = (
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
34
tests/test_litellm/proxy/test_sensitive_route_auth.py
Normal file
34
tests/test_litellm/proxy/test_sensitive_route_auth.py
Normal 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"
|
||||
)
|
||||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue