fix(proxy): keep /health inside the caller's team and admin-gate the Bedrock runtime endpoint

This commit is contained in:
mateo-berri 2026-09-03 17:57:13 -07:00
parent e75dd56631
commit a84f4d6206
3 changed files with 169 additions and 18 deletions

View file

@ -35,7 +35,7 @@ from litellm.router_utils.auto_router_model_naming import (
# Provider routing fields. Allowed for proxy admins so they can see which
# region/version a deployment is checking; gated at the endpoint layer for
# non-admin callers (see _strip_admin_only_fields_from_health_result).
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS: Final = ("api_base", "api_version")
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS: Final = ("api_base", "api_version", "aws_bedrock_runtime_endpoint")
MINIMAL_DISPLAY_PARAMS: Final = frozenset({"model", "mode_error"})
@ -49,6 +49,7 @@ HEALTH_DISPLAY_PARAMS: Final = (
"base_model",
"aws_region_name",
"region_name",
"watsonx_region_name",
"vertex_project",
"vertex_location",
"tpm",

View file

@ -845,7 +845,7 @@ def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
def _strip_admin_only_fields_from_health_result(result: dict) -> dict:
"""
Return a copy of the /health response with provider routing fields
(``api_base``, ``api_version``) removed from each healthy/unhealthy
(``ADMIN_ONLY_HEALTH_DISPLAY_PARAMS``) removed from each healthy/unhealthy
endpoint entry. Used to hide those fields from non-admin callers while
still showing them which deployments they own and whether each one is
healthy. Proxy admins receive the unmodified result.
@ -877,6 +877,18 @@ def _health_accessible_model_names(
)
def _caller_may_probe_deployment(
deployment: Mapping[str, object], allowed_models: frozenset[str], llm_router: Router | None, team_id: str | None
) -> bool:
"""Same deployment visibility rule as request auth: another team's deployment is never in scope."""
if llm_router is None:
return deployment.get("model_name") in allowed_models
model: Final = dict(deployment)
return any(
llm_router.should_include_deployment(model_name=name, model=model, team_id=team_id) for name in allowed_models
)
def _resolve_targeted_model_ids(model_list: list, model: str | None, model_id: str | None) -> set | None:
"""
Resolve a ``/health`` ``model`` / ``model_id`` query param to the set of
@ -1079,7 +1091,9 @@ async def health_endpoint(
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
if is_admin:
return result
response.headers["Litellm-Health-Field-Notice"] = "api_base and api_version are admin-only on this endpoint"
response.headers["Litellm-Health-Field-Notice"] = (
f"{', '.join(ADMIN_ONLY_HEALTH_DISPLAY_PARAMS)} are admin-only on this endpoint"
)
return _strip_admin_only_fields_from_health_result(result)
try:
@ -1106,7 +1120,10 @@ async def health_endpoint(
allowed_models: Final = _health_accessible_model_names(user_api_key_dict, llm_router)
restrict_to_allowed_models: Final = allowed_models is not None
_llm_model_list: Final = [
m for m in copy.deepcopy(llm_model_list) if allowed_models is None or m.get("model_name") in allowed_models
m
for m in copy.deepcopy(llm_model_list)
if allowed_models is None
or _caller_may_probe_deployment(m, allowed_models, llm_router, user_api_key_dict.team_id)
]
if use_background_health_checks:
# The cached background result covers every model. When the

View file

@ -1,4 +1,5 @@
import asyncio
import copy
import json
import time
from collections.abc import Iterator, Mapping, Sequence
@ -19,6 +20,7 @@ import litellm.proxy.health_endpoints._health_endpoints as _health_endpoints_mod
from litellm.litellm_core_utils.health_check_helpers import TEST_IMAGE_BASE64
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.router import Router
from litellm.proxy.health_endpoints._health_endpoints import (
_db_health_readiness_check,
_show_no_redis_warning,
@ -1652,18 +1654,8 @@ async def test_health_endpoint_resolves_all_team_models_to_team_allowlist():
assert returned_names == {"model-b"}, f"all-team-models key should health-check the team's models: {returned_names}"
class _AccessGroupRouter:
"""Router stand-in exposing only the two lookups /health uses to expand a key's model grants."""
def __init__(self, access_groups: dict[str, list[str]], model_names: list[str]) -> None:
self._access_groups = access_groups
self._model_names = model_names
def get_model_access_groups(self, model_name=None, model_access_group=None, team_id=None):
return self._access_groups
def get_model_names(self, team_id=None):
return self._model_names
def _router_for(model_list: Sequence[Mapping[str, object]]) -> Router:
return Router(model_list=copy.deepcopy(list(model_list)))
_ACCESS_GROUP_MODEL_LIST = [
@ -1678,7 +1670,29 @@ _ACCESS_GROUP_MODEL_LIST = [
"model_info": {"id": "id-openai"},
},
]
_ACCESS_GROUP_ROUTER = _AccessGroupRouter({"bedrock-group": ["bedrock-nova"]}, ["bedrock-nova", "gpt-5.4-mini"])
_ACCESS_GROUP_ROUTER = _router_for(_ACCESS_GROUP_MODEL_LIST)
_TEAM_MODEL_LIST = [
_ACCESS_GROUP_MODEL_LIST[0],
{
"model_name": "bedrock-nova_team-b_9f2c",
"litellm_params": {"model": "bedrock/us.amazon.nova-2-lite-v1:0"},
"model_info": {
"id": "id-team-b",
"team_id": "team-b",
"team_public_model_name": "bedrock-nova",
"access_groups": ["bedrock-group"],
},
},
]
_TEAM_CACHED_RESULTS = {
"healthy_endpoints": [
{"model": "bedrock/us.amazon.nova-2-lite-v1:0", "model_id": "id-bedrock"},
{"model": "bedrock/us.amazon.nova-2-lite-v1:0", "model_id": "id-team-b"},
],
"unhealthy_endpoints": [],
"healthy_count": 2,
"unhealthy_count": 0,
}
_ACCESS_GROUP_CACHED_RESULTS = {
"healthy_endpoints": [
{"model": "bedrock/us.amazon.nova-2-lite-v1:0", "model_id": "id-bedrock"},
@ -2046,7 +2060,7 @@ async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not():
# withheld so clients that previously parsed them can detect the change.
assert (
non_admin_response.headers.get("Litellm-Health-Field-Notice")
== "api_base and api_version are admin-only on this endpoint"
== "api_base, api_version, aws_bedrock_runtime_endpoint are admin-only on this endpoint"
)
assert "Litellm-Health-Field-Notice" not in admin_response.headers
@ -2776,6 +2790,123 @@ def test_clean_endpoint_data_never_displays_credential_fields(credential_field,
assert canary not in str(cleaned)
async def _live_probed_model_ids(
model_list: Sequence[Mapping[str, object]], user_api_key_dict: UserAPIKeyAuth
) -> set[str]:
from fastapi import Response
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
captured: dict = {}
async def fake_perform(**kwargs):
captured["model_list"] = kwargs["model_list"]
return {"healthy_endpoints": [], "unhealthy_endpoints": [], "healthy_count": 0, "unhealthy_count": 0}
with (
_proxy_health_globals(model_list, _router_for(model_list)),
patch( # test-quality-ok: the model list handed to the probe is the assertion; no injection seam
"litellm.proxy.health_endpoints._health_endpoints._perform_health_check_and_save",
side_effect=fake_perform,
),
):
await health_endpoint(response=Response(), user_api_key_dict=user_api_key_dict)
return {m["model_info"]["id"] for m in captured["model_list"]}
@pytest.mark.asyncio
async def test_health_endpoint_hides_another_teams_deployment_behind_a_shared_access_group():
"""
Expanding an access group must not reach past the team boundary: a
team-a key holding the group name may not probe team-b's deployment even
though that deployment sits in the same group.
"""
probed = await _live_probed_model_ids(
_TEAM_MODEL_LIST,
UserAPIKeyAuth(api_key="hashed-test-key", models=["bedrock-group"], team_id="team-a"),
)
assert probed == {"id-bedrock"}
@pytest.mark.asyncio
async def test_health_endpoint_shows_a_teams_own_deployment_by_its_public_name():
"""
A team key names its team deployment by ``team_public_model_name``, while
the proxy model list carries the internal ``<name>_<team_id>_<uuid>``
name; the deployment must still be probed for its own team.
"""
probed = await _live_probed_model_ids(
_TEAM_MODEL_LIST,
UserAPIKeyAuth(api_key="hashed-test-key", models=["bedrock-nova"], team_id="team-b"),
)
assert probed == {"id-bedrock", "id-team-b"}
@pytest.mark.asyncio
async def test_health_endpoint_hides_another_teams_deployment_on_background_cache_path():
from fastapi import Response
from litellm.proxy.health_endpoints._health_endpoints import health_endpoint
with _proxy_health_globals(
_TEAM_MODEL_LIST,
_router_for(_TEAM_MODEL_LIST),
use_background_health_checks=True,
health_check_results=_TEAM_CACHED_RESULTS,
):
result = await health_endpoint(
response=Response(),
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-test-key", models=["bedrock-group"], team_id="team-a"),
model=None,
model_id=None,
)
assert [ep["model_id"] for ep in result["healthy_endpoints"]] == ["id-bedrock"]
assert result["healthy_count"] == 1
def test_health_test_connection_keeps_error_and_raw_request_through_the_allowlist(monkeypatch):
"""
The dashboard's Test Connect button reads ``result.error`` and
``result.raw_request_typed_dict`` from /health/test_connection, so the
allowlist must keep both while dropping the probe's own params.
"""
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
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)
with (
patch( # test-quality-ok: the endpoint reads the proxy-global DB client and 500s when it is None; it has no injection seam
"litellm.proxy.proxy_server.prisma_client", MagicMock()
),
respx.mock(assert_all_called=True) as respx_mock,
):
respx_mock.post(host="api.openai.com", path="/v1/chat/completions").respond(
status_code=401, json={"error": {"message": "Incorrect API key provided"}}
)
response = client.post(
"/health/test_connection",
json={
"mode": "chat",
"litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "sk-test", "timeout": 7},
},
)
assert response.status_code == 200, response.text
body = response.json()
assert body["status"] == "error"
assert "Incorrect API key provided" in body["result"]["error"]
assert "api.openai.com" in body["result"]["raw_request_typed_dict"]["raw_request_api_base"]
assert not {"api_key", "timeout", "exception"} & set(body["result"])
def test_clean_endpoint_data_keeps_only_json_safe_diagnostics():
"""
LIT-6907: _clean_endpoint_data used to copy every litellm_param not on a
@ -2796,6 +2927,7 @@ def test_clean_endpoint_data_keeps_only_json_safe_diagnostics():
"api_key": "CANARY-API-KEY",
"x-ratelimit-remaining-requests": 99,
"raw_request_typed_dict": {"raw_request_api_base": "https://example.test"},
"aws_bedrock_runtime_endpoint": "https://vpce-bedrock.example.test",
},
details=True,
)
@ -2806,6 +2938,7 @@ def test_clean_endpoint_data_keeps_only_json_safe_diagnostics():
"aws_region_name": "us-east-1",
"x-ratelimit-remaining-requests": 99,
"raw_request_typed_dict": {"raw_request_api_base": "https://example.test"},
"aws_bedrock_runtime_endpoint": "https://vpce-bedrock.example.test",
}
assert jsonable_encoder(cleaned) == cleaned