diff --git a/litellm/__init__.py b/litellm/__init__.py index 6c780231d2a..5305edc9be6 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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] = ( diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 24aff5548e3..c175c41c3c9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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", diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 9159a8ff9da..295dfeb788e 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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. diff --git a/litellm/proxy/common_utils/debug_utils.py b/litellm/proxy/common_utils/debug_utils.py index 99eeeda1c86..9b2c3ddce46 100644 --- a/litellm/proxy/common_utils/debug_utils.py +++ b/litellm/proxy/common_utils/debug_utils.py @@ -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 diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 35c9edb937d..096e23e673d 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -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(): """ diff --git a/litellm/proxy/middleware/prometheus_auth_middleware.py b/litellm/proxy/middleware/prometheus_auth_middleware.py index 3b30fd3d63c..cfc4cbd64b2 100644 --- a/litellm/proxy/middleware/prometheus_auth_middleware.py +++ b/litellm/proxy/middleware/prometheus_auth_middleware.py @@ -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)) diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index d6ce454682e..7d9da543c75 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -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(): diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 4b8b341a4bf..d030fabe8b5 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -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 diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index eed4f076467..db3ae9ad942 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -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", diff --git a/tests/spend_tracking_tests/test_spend_accuracy_tests.py b/tests/spend_tracking_tests/test_spend_accuracy_tests.py index 15e00d93356..b50afeb843a 100644 --- a/tests/spend_tracking_tests/test_spend_accuracy_tests.py +++ b/tests/spend_tracking_tests/test_spend_accuracy_tests.py @@ -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}") diff --git a/tests/test_health.py b/tests/test_health.py index 00f095022b9..15dc2330ffb 100644 --- a/tests/test_health.py +++ b/tests/test_health.py @@ -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(): """ diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index dd24ac87495..398084fdc5e 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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: 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 353e67c9f77..2edcb00c967 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -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(): diff --git a/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py b/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py index 310ee11573b..6cab5baee9a 100644 --- a/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py +++ b/tests/test_litellm/proxy/middleware/test_prometheus_auth_middleware.py @@ -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") diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py index a65462d3f9b..f82da59899b 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -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(): diff --git a/tests/test_litellm/proxy/test_sensitive_route_auth.py b/tests/test_litellm/proxy/test_sensitive_route_auth.py new file mode 100644 index 00000000000..19998e52779 --- /dev/null +++ b/tests/test_litellm/proxy/test_sensitive_route_auth.py @@ -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" + ) diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 5fdcb9456df..c27d7eedcdb 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -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 ):