From cad49ee1718cabc7a133c1f3e195368007944265 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 18:03:02 -0700 Subject: [PATCH] fix(proxy): gate disable_global_guardrails on keys and teams to proxy admins (#42699) * fix(proxy): gate disable_global_guardrails on keys and teams to proxy admins Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: cover metadata smuggle with explicit false and UI toggle gating Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): satisfy PT017 in resend-stored guardrail flag test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): keep regenerate_key_fn under the C901 ceiling via a guardrail opt-out helper Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(ui): regenerate schema.d.ts for guardrail opt-out docstrings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): gate disable_global_guardrails on caller-sent metadata, not server defaults Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): audit cells for disable_global_guardrails admin gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): restore contracts.json formatting Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): share guardrail opt-out helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): hide the team disable_global_guardrails switch from non proxy admins Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): drop covers markers and bound the slow sink check to the sink delay Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../management_endpoints/common_utils.py | 28 ++ .../key_management_endpoints.py | 45 ++- .../management_endpoints/team_endpoints.py | 18 +- .../authorization/_guardrail_opt_out.py | 71 ++++ .../test_key_guardrail_opt_out.py | 371 ++++++++++++++++++ .../test_key_guardrail_opt_out_chaos.py | 319 +++++++++++++++ .../test_key_guardrail_opt_out_runtime.py | 230 +++++++++++ .../management_endpoints/test_common_utils.py | 129 ++++++ .../test_key_management_endpoints.py | 193 +++++++++ .../test_team_endpoints.py | 77 ++++ .../src/components/Teams.test.tsx | 47 +++ ui/litellm-dashboard/src/components/Teams.tsx | 48 +-- .../create_key_button.integration.test.tsx | 15 + .../organisms/create_key_button.tsx | 70 ++-- .../src/components/team/TeamInfo.test.tsx | 47 +++ .../src/components/team/TeamInfo.tsx | 40 +- .../key_edit_view.integration.test.tsx | 28 ++ .../components/templates/key_edit_view.tsx | 26 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 8 +- 19 files changed, 1714 insertions(+), 96 deletions(-) create mode 100644 tests/integration/authorization/_guardrail_opt_out.py create mode 100644 tests/integration/authorization/test_key_guardrail_opt_out.py create mode 100644 tests/integration/authorization/test_key_guardrail_opt_out_chaos.py create mode 100644 tests/integration/authorization/test_key_guardrail_opt_out_runtime.py diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 78e3ac7bd66..0bd4eb5a5d8 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -173,6 +173,34 @@ def _check_passthrough_routes_caller_permission( ) +def _check_disable_global_guardrails_caller_permission( + disable_global_guardrails: bool | None, + metadata: Mapping[str, object] | None, + user_api_key_dict: UserAPIKeyAuth, + *, + entity: str = "key", + existing_metadata: Mapping[str, object] | None = None, +) -> None: + """ + Only proxy admins may opt a key or team out of default-on guardrails, whether the + flag is top-level or under `metadata`. Re-sending a flag that is already stored is + not an opt-out, so non-admin edits of an already exempted object still go through. + """ + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: + return + requested: Final = bool(disable_global_guardrails) or ( + metadata is not None and bool(metadata.get("disable_global_guardrails")) + ) + if not requested: + return + if existing_metadata is not None and existing_metadata.get("disable_global_guardrails") is True: + return + raise HTTPException( + status_code=403, + detail={"error": f"Only proxy admins can set `disable_global_guardrails` on a {entity}."}, + ) + + def _is_user_team_admin(user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable) -> bool: for member in team_obj.members_with_roles: if (member.user_id is not None and member.user_id == user_api_key_dict.user_id) and member.role == "admin": diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 854b80ba2c3..306ea90d7f1 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -85,6 +85,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks from litellm.proxy.hooks.model_max_budget_limiter import build_model_max_budget_usage from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, _check_passthrough_routes_caller_permission, _is_user_org_admin_for_team, _is_user_team_admin, @@ -1224,6 +1225,7 @@ async def _common_key_generation_helper( # default_key_generate_params injected. _requested_max_budget: Final = data.max_budget _requested_team_id: Final = data.team_id + _requested_metadata: Final = data.metadata # pyright: ignore[reportUnknownMemberType] # request models declare `metadata` as bare dict # check if user set default key/generate params on config.yaml if litellm.default_key_generate_params is not None: @@ -1311,6 +1313,11 @@ async def _common_key_generation_helper( data=data, user_api_key_dict=user_api_key_dict, ) + _check_disable_global_guardrails_caller_permission( + data.disable_global_guardrails, + _requested_metadata, + user_api_key_dict, + ) # APPLY ENTERPRISE KEY MANAGEMENT PARAMS try: @@ -1966,7 +1973,7 @@ async def generate_key_fn( - metadata: Optional[dict] - Metadata for key, store information for key. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } - guardrails: Optional[List[str]] - List of active guardrails for the key - policies: Optional[List[str]] - List of policy names to apply to the key. Policies define guardrails, conditions, and inheritance rules. - - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. + - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. Proxy admin only. - throttle_on_budget_exceeded: Optional[bool] - When the key exceeds its max_budget, throttle its tpm/rpm to the global budget_exceeded_throttle_percentage instead of blocking the key entirely. - enable_prompt_caching: Optional[bool] - Auto-inject prompt caching breakpoints (Anthropic cache_control markers) on requests made with this key. Supported Claude models on Anthropic, Bedrock, Vertex AI, and Azure AI only. - permissions: Optional[dict] - key-specific permissions. Currently just used for turning off pii masking (if connected). Example - {"pii": false} @@ -2729,6 +2736,13 @@ async def _process_single_key_update( prisma_client=prisma_client, ) + _check_disable_global_guardrails_caller_permission( + update_key_request.disable_global_guardrails, + update_key_request.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict + user_api_key_dict, + existing_metadata=existing_key_row.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_VerificationToken.metadata is a bare dict + ) + enforce_batch_enqueued_token_limit_is_admin_only( data=update_key_request, existing_metadata=existing_key_row.metadata, @@ -3025,6 +3039,12 @@ async def _validate_update_key_data( data=data, user_api_key_dict=user_api_key_dict, ) + _check_disable_global_guardrails_caller_permission( + data.disable_global_guardrails, + data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict + user_api_key_dict, + existing_metadata=existing_key_row.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_VerificationToken.metadata is a bare dict + ) _validate_caller_can_change_key_ownership( data=data, @@ -3328,7 +3348,7 @@ async def update_key_fn( - send_invite_email: Optional[bool] - Send invite email to user_id - guardrails: Optional[List[str]] - List of active guardrails for the key - policies: Optional[List[str]] - List of policy names to apply to the key. Policies define guardrails, conditions, and inheritance rules. - - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. + - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. Proxy admin only. - throttle_on_budget_exceeded: Optional[bool] - When the key exceeds its max_budget, throttle its tpm/rpm to the global budget_exceeded_throttle_percentage instead of blocking the key entirely. - enable_prompt_caching: Optional[bool] - Auto-inject prompt caching breakpoints (Anthropic cache_control markers) on requests made with this key. Supported Claude models on Anthropic, Bedrock, Vertex AI, and Azure AI only. - prompts: Optional[List[str]] - List of prompts that the key is allowed to use. @@ -5622,6 +5642,21 @@ async def _execute_virtual_key_regeneration( return response +def _check_regenerate_guardrail_opt_out( + data: RegenerateKeyRequest | None, + existing_metadata: Mapping[str, object] | None, + user_api_key_dict: UserAPIKeyAuth, +) -> None: + if data is None: + return + _check_disable_global_guardrails_caller_permission( + data.disable_global_guardrails, + data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict + user_api_key_dict, + existing_metadata=existing_metadata, + ) + + @router.post( "/key/{key:path}/regenerate", tags=["key management"], @@ -5805,6 +5840,12 @@ async def regenerate_key_fn( detail={"error": f"Key {key} not found."}, ) + _check_regenerate_guardrail_opt_out( + data, + _key_in_db.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_VerificationToken.metadata is a bare dict + user_api_key_dict, + ) + # check if user has permission to regenerate key await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 09c6c05e22b..6cec3e714ec 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -127,6 +127,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( get_daily_activity_aggregated, ) from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, _check_passthrough_routes_caller_permission, _is_user_org_admin_for_team, _is_user_team_admin, @@ -1416,7 +1417,7 @@ async def new_team( - model_max_budget: Optional[dict] - Per-model max budget every key on the team inherits unless the key sets its own for that model. Example: {"gpt-4o": {"max_budget": 10, "budget_duration": "1d"}} - guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails) - policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies) - - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. + - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the team. Proxy admin only. - object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission. - team_member_budget: Optional[float] - The maximum budget allocated to an individual team member. - team_member_budget_duration: Optional[str] - The duration of the budget for the team member. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets) @@ -1638,6 +1639,12 @@ async def new_team( data.members_with_roles.append(Member(role="admin", user_id=user_api_key_dict.user_id)) _check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team") + _check_disable_global_guardrails_caller_permission( + data.disable_global_guardrails, + data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict + user_api_key_dict, + entity="team", + ) if isinstance(data.metadata, dict): TeamMemberBudgetHandler.strip_system_managed_metadata_keys(data.metadata) @@ -2172,7 +2179,7 @@ async def update_team( - model_max_budget: Optional[dict] - Per-model max budget every key on the team inherits unless the key sets its own for that model. Example: {"gpt-4o": {"max_budget": 10, "budget_duration": "1d"}} - guardrails: Optional[List[str]] - Guardrails for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails) - policies: Optional[List[str]] - Policies for the team. [Docs](https://docs.litellm.ai/docs/proxy/guardrails/guardrail_policies) - - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key. + - disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the team. Proxy admin only. - object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission. - team_member_budget: Optional[float] - The maximum budget allocated to an individual team member. - team_member_budget_duration: Optional[str] - The duration of the budget for the team member. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets) @@ -2313,6 +2320,13 @@ async def update_team( ) _check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team") + _check_disable_global_guardrails_caller_permission( + data.disable_global_guardrails, + data.metadata, # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # request models declare `metadata` as bare dict + user_api_key_dict, + entity="team", + existing_metadata=_existing_team_metadata if isinstance(_existing_team_metadata, dict) else None, # pyright: ignore[reportUnknownArgumentType] # existing_team_row.metadata is a bare dict + ) if data.soft_budget is not None: max_budget_to_check = data.max_budget if data.max_budget is not None else existing_team_row.max_budget diff --git a/tests/integration/authorization/_guardrail_opt_out.py b/tests/integration/authorization/_guardrail_opt_out.py new file mode 100644 index 00000000000..e813993edf4 --- /dev/null +++ b/tests/integration/authorization/_guardrail_opt_out.py @@ -0,0 +1,71 @@ +import json +import uuid +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import httpx +import yaml +from pydantic import JsonValue + +from integration._support.client import Gateway, Scenario, object_value +from integration._support.database import read_rows +from integration._support.wire import Reply, Request + +MANAGEMENT_ROUTES: Final = ["/key/*", "/team/new", "/team/update", "/v1/chat/completions"] + + +def denying_guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api" + return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic policy denial"}).encode()) + + +def guardrail_config(policy_url: str, path: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": "guardrail" + uuid.uuid4().hex, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": policy_url, + "api_key": "synthetic-guardrail-key", + }, + } + ] + path.write_text(yaml.safe_dump(config)) + return path + + +def stored_metadata(token: str) -> dict[str, object]: + rows: Final = read_rows( + 'SELECT metadata FROM "LiteLLM_VerificationToken" WHERE token = %s', (sha256(token.encode()).hexdigest(),) + ) + assert len(rows) == 1, rows + return rows[0]["metadata"] + + +def non_admin_caller(scenario: Scenario, member: str, team: str, model: str) -> str: + return scenario.key(user_id=member, team_id=team, models=[model], allowed_routes=MANAGEMENT_ROUTES) + + +def chat(candidate: Gateway, model: str, key: str, marker: str, *, stream: bool = False) -> httpx.Response: + return candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": marker}], "stream": stream}, + key=key, + ) + + +def upstream_observations(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]: + with httpx.Client(timeout=5, trust_env=False) as client: + drained: Final = object_value(client.get(f"{gateway.upstream_url}/__observations").json()) + requests: Final = drained["requests"] + assert isinstance(requests, list) + return tuple(object_value(entry) for entry in requests) + + +def upstream_hits(gateway: Gateway, marker: str) -> int: + return sum(1 for entry in upstream_observations(gateway) if marker in json.dumps(entry.get("body"))) diff --git a/tests/integration/authorization/test_key_guardrail_opt_out.py b/tests/integration/authorization/test_key_guardrail_opt_out.py new file mode 100644 index 00000000000..6591b0b3993 --- /dev/null +++ b/tests/integration/authorization/test_key_guardrail_opt_out.py @@ -0,0 +1,371 @@ +import uuid +from pathlib import Path +from typing import Final + +import httpx +import yaml + +from integration._support.client import Gateway, Scenario, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.wire import wire_server +from integration.authorization._guardrail_opt_out import ( + denying_guardrail, + guardrail_config, + non_admin_caller, + stored_metadata, +) + +_KEY_ROUTES: Final = ["/key/generate", "/key/update", "/key/regenerate", "/v1/chat/completions"] + + +def test_non_admin_cannot_opt_key_out_of_default_on_guardrail(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy: + config: Final = guardrail_config(policy.url, tmp_path / "default_on.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = scenario.key(user_id=member, models=[model], allowed_routes=_KEY_ROUTES) + own: Final = scenario.key(team_id=team, models=[model]) + + plain: Final = candidate.request("POST", "/key/generate", {"team_id": team, "models": [model]}, key=caller) + assert plain.status_code == 200, plain.text + scenario.cleanups.callback(scenario.delete_key, string_value(plain.json()["key"])) + + generated: Final = candidate.request( + "POST", + "/key/generate", + {"team_id": team, "models": [model], "disable_global_guardrails": True}, + key=caller, + ) + if generated.status_code == 200: + scenario.cleanups.callback(scenario.delete_key, string_value(generated.json()["key"])) + assert generated.status_code == 403, generated.text + assert "disable_global_guardrails" in generated.text + + smuggled: Final = candidate.request( + "POST", + "/key/generate", + {"team_id": team, "models": [model], "metadata": {"disable_global_guardrails": True}}, + key=caller, + ) + if smuggled.status_code == 200: + scenario.cleanups.callback(scenario.delete_key, string_value(smuggled.json()["key"])) + assert smuggled.status_code == 403, smuggled.text + + updated: Final = candidate.request( + "POST", "/key/update", {"key": own, "disable_global_guardrails": True}, key=caller + ) + assert updated.status_code == 403, updated.text + regenerated: Final = candidate.request( + "POST", "/key/regenerate", {"key": own, "disable_global_guardrails": True}, key=caller + ) + assert regenerated.status_code == 403, regenerated.text + assert "disable_global_guardrails" not in stored_metadata(own) + + blocked: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "synthetic denied marker"}]}, + key=own, + ) + assert blocked.status_code == 400 and "synthetic policy denial" in blocked.text, blocked.text + + exempt: Final = scenario.key(team_id=team, models=[model], disable_global_guardrails=True) + assert stored_metadata(exempt)["disable_global_guardrails"] is True + resaved: Final = candidate.request( + "POST", + "/key/update", + { + "key": exempt, + "key_alias": "renamed" + uuid.uuid4().hex, + "metadata": {"disable_global_guardrails": True}, + }, + key=caller, + ) + assert resaved.status_code == 200, resaved.text + assert stored_metadata(exempt)["disable_global_guardrails"] is True + served: Final = candidate.chat(model, key=exempt, text="synthetic denied marker") + assert object_value(served["usage"])["total_tokens"] == 40 + assert len(policy.drain()) == 1 + + +def _team_metadata(team_id: str) -> dict[str, object]: + rows: Final = read_rows('SELECT metadata FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team_id,)) + assert len(rows) == 1, rows + return rows[0]["metadata"] + + +def _drop_created_key(scenario: Scenario, response: httpx.Response) -> None: + if response.status_code == 200: + scenario.cleanups.callback(scenario.delete_key, string_value(response.json()["key"])) + + +def test_non_admin_flag_denied_on_every_key_write_route(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy: + config: Final = guardrail_config(policy.url, tmp_path / "denied-routes.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + own: Final = scenario.key(team_id=team, models=[model]) + + attempts: Final = ( + ("POST", "/key/generate", {"team_id": team, "models": [model], "disable_global_guardrails": True}), + ( + "POST", + "/key/generate", + {"team_id": team, "models": [model], "metadata": {"disable_global_guardrails": True}}, + ), + ( + "POST", + "/key/generate", + { + "team_id": team, + "models": [model], + "disable_global_guardrails": False, + "metadata": {"disable_global_guardrails": True}, + }, + ), + ("POST", "/key/update", {"key": own, "disable_global_guardrails": True}), + ("POST", "/key/update", {"key": own, "metadata": {"disable_global_guardrails": True}}), + ("POST", "/key/regenerate", {"key": own, "disable_global_guardrails": True}), + ("POST", f"/key/{own}/regenerate", {"disable_global_guardrails": True}), + ( + "POST", + "/key/service-account/generate", + {"team_id": team, "disable_global_guardrails": True}, + ), + ) + for method, path, body in attempts: + response: Final = candidate.request(method, path, body, key=caller) + _drop_created_key(scenario, response) + assert response.status_code == 403, f"{method} {path}: {response.text}" + assert "disable_global_guardrails" in response.text, response.text + assert "disable_global_guardrails" not in stored_metadata(own) + + service_alias: Final = "audit-sa-" + uuid.uuid4().hex + service_denied: Final = candidate.request( + "POST", + "/key/service-account/generate", + {"team_id": team, "key_alias": service_alias, "disable_global_guardrails": True}, + key=caller, + ) + _drop_created_key(scenario, service_denied) + assert ( + read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE key_alias = %s', (service_alias,)) == [] + ), service_denied.text + + +def test_non_admin_flag_denied_on_team_new(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy: + config: Final = guardrail_config(policy.url, tmp_path / "denied-team.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + + alias: Final = "audit-team-" + uuid.uuid4().hex + denied: Final = candidate.request( + "POST", + "/team/new", + {"team_alias": alias, "models": [model], "disable_global_guardrails": True}, + key=caller, + ) + created: Final = read_rows('SELECT team_id FROM "LiteLLM_TeamTable" WHERE team_alias = %s', (alias,)) + for row in created: + scenario.cleanups.callback(scenario.delete_team, str(row["team_id"])) + assert denied.status_code == 403, denied.text + assert "disable_global_guardrails" in denied.text, denied.text + + +def test_admin_flag_writes_succeed_on_all_routes(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy: + config: Final = guardrail_config(policy.url, tmp_path / "admin-routes.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + team: Final = scenario.team(models=[model]) + + generated: Final = candidate.post( + "/key/generate", {"team_id": team, "models": [model], "disable_global_guardrails": True} + ) + generated_key: Final = string_value(generated["key"]) + scenario.cleanups.callback(scenario.delete_key, generated_key) + assert stored_metadata(generated_key)["disable_global_guardrails"] is True + + plain: Final = scenario.key(team_id=team, models=[model]) + candidate.post("/key/update", {"key": plain, "disable_global_guardrails": True}) + assert stored_metadata(plain)["disable_global_guardrails"] is True + + regen_source: Final = string_value( + candidate.post("/key/generate", {"team_id": team, "models": [model]})["key"] + ) + regenerated: Final = candidate.post( + "/key/regenerate", {"key": regen_source, "disable_global_guardrails": True} + ) + regenerated_key: Final = string_value(regenerated["key"]) + scenario.cleanups.callback(scenario.delete_key, regenerated_key) + assert stored_metadata(regenerated_key)["disable_global_guardrails"] is True + + new_team: Final = candidate.post( + "/team/new", {"team_alias": "audit-admin-" + uuid.uuid4().hex, "disable_global_guardrails": True} + ) + new_team_id: Final = string_value(new_team["team_id"]) + scenario.cleanups.callback(scenario.delete_team, new_team_id) + assert _team_metadata(new_team_id)["disable_global_guardrails"] is True + + candidate.post("/team/update", {"team_id": team, "disable_global_guardrails": True}) + assert _team_metadata(team)["disable_global_guardrails"] is True + + +def test_non_admin_resave_omit_and_revoke_sequences(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy: + config: Final = guardrail_config(policy.url, tmp_path / "resave.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + exempt: Final = scenario.key(team_id=team, models=[model], disable_global_guardrails=True) + assert stored_metadata(exempt)["disable_global_guardrails"] is True + + resaved: Final = candidate.request( + "POST", + "/key/update", + { + "key": exempt, + "key_alias": "audit-resave-" + uuid.uuid4().hex, + "metadata": {"disable_global_guardrails": True}, + }, + key=caller, + ) + assert resaved.status_code == 200, resaved.text + assert stored_metadata(exempt)["disable_global_guardrails"] is True + + omitted: Final = candidate.request( + "POST", + "/key/update", + {"key": exempt, "key_alias": "audit-omit-" + uuid.uuid4().hex}, + key=caller, + ) + assert omitted.status_code == 200, omitted.text + + candidate.post("/key/update", {"key": exempt, "disable_global_guardrails": False}) + assert stored_metadata(exempt)["disable_global_guardrails"] is False + + rejected: Final = candidate.request( + "POST", "/key/update", {"key": exempt, "disable_global_guardrails": True}, key=caller + ) + assert rejected.status_code == 403, rejected.text + assert "disable_global_guardrails" in rejected.text, rejected.text + assert stored_metadata(exempt)["disable_global_guardrails"] is False + + +def test_generate_ignores_server_default_metadata_flag(gateway: Gateway, tmp_path: Path) -> None: + raw: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + raw.setdefault("litellm_settings", {})["default_key_generate_params"] = { + "metadata": {"disable_global_guardrails": True} + } + path: Final = tmp_path / "server-defaults.yaml" + path.write_text(yaml.safe_dump(raw)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + + generated: Final = candidate.post("/key/generate", {"team_id": team, "models": [model]}, key=caller) + generated_key: Final = string_value(generated["key"]) + scenario.cleanups.callback(scenario.delete_key, generated_key) + assert stored_metadata(generated_key)["disable_global_guardrails"] is True + + explicit: Final = candidate.request( + "POST", + "/key/generate", + {"team_id": team, "models": [model], "metadata": {"disable_global_guardrails": True}}, + key=caller, + ) + _drop_created_key(scenario, explicit) + assert explicit.status_code == 403, explicit.text + assert "disable_global_guardrails" in explicit.text, explicit.text + + +def test_sad_flag_inputs_on_key_generate(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {}) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + + denied_bodies: Final = ( + {"team_id": team, "models": [model], "disable_global_guardrails": "true"}, + {"team_id": team, "models": [model], "disable_global_guardrails": 1}, + {"team_id": team, "models": [model], "metadata": {"disable_global_guardrails": "true"}}, + {"team_id": team, "models": [model], "metadata": {"disable_global_guardrails": 1}}, + {"team_id": team, "models": [model], "metadata": {"disable_global_guardrails": "x" * 5120}}, + ) + for body in denied_bodies: + response: Final = candidate.request("POST", "/key/generate", body, key=caller) + _drop_created_key(scenario, response) + assert response.status_code == 403, response.text + assert "disable_global_guardrails" in response.text, response.text + + invalid_bodies: Final = ( + {"team_id": team, "models": [model], "disable_global_guardrails": []}, + {"team_id": team, "models": [model], "disable_global_guardrails": {}}, + ) + for body in invalid_bodies: + rejected: Final = candidate.request("POST", "/key/generate", body, key=caller) + assert rejected.status_code == 422, rejected.text + + unauthenticated: Final = candidate.request( + "POST", + "/key/generate", + {"team_id": team, "models": [model], "disable_global_guardrails": True}, + key="sk-not-a-real-key-" + uuid.uuid4().hex, + ) + assert unauthenticated.status_code == 401, unauthenticated.text + + repeat_alias: Final = "audit-repeat-" + uuid.uuid4().hex + for _ in range(2): + repeated: Final = candidate.request( + "POST", + "/key/generate", + {"team_id": team, "models": [model], "key_alias": repeat_alias, "disable_global_guardrails": True}, + key=caller, + ) + _drop_created_key(scenario, repeated) + assert repeated.status_code == 403, repeated.text + assert read_rows('SELECT token FROM "LiteLLM_VerificationToken" WHERE key_alias = %s', (repeat_alias,)) == [] + + +def test_falsy_metadata_flag_shapes_stay_stored_and_guarded(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy: + config: Final = guardrail_config(policy.url, tmp_path / "falsy.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + + for shape in ([], {}): + created: Final = candidate.request( + "POST", + "/key/generate", + {"team_id": team, "models": [model], "metadata": {"disable_global_guardrails": shape}}, + key=caller, + ) + _drop_created_key(scenario, created) + assert created.status_code == 200, created.text + token: Final = string_value(created.json()["key"]) + assert stored_metadata(token)["disable_global_guardrails"] == shape + blocked: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "synthetic denied marker"}]}, + key=token, + ) + assert blocked.status_code == 400 and "synthetic policy denial" in blocked.text, blocked.text diff --git a/tests/integration/authorization/test_key_guardrail_opt_out_chaos.py b/tests/integration/authorization/test_key_guardrail_opt_out_chaos.py new file mode 100644 index 00000000000..f205b3804f6 --- /dev/null +++ b/tests/integration/authorization/test_key_guardrail_opt_out_chaos.py @@ -0,0 +1,319 @@ +import json +import os +import signal +import socket +import threading +import time +import uuid +from concurrent.futures import ThreadPoolExecutor +from hashlib import sha256 +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +import psutil + +from integration._support.client import Gateway, eventually, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy, owned_proxy_process +from integration.authorization._guardrail_opt_out import ( + chat, + guardrail_config, + non_admin_caller, + stored_metadata, + upstream_hits, + upstream_observations, +) + + +class _GuardrailSink: + """Test-owned guardrail endpoint that can be stopped and restarted on the same port.""" + + def __init__(self, *, delay_seconds: float = 0.0, action: str = "BLOCKED") -> None: + self.received: SimpleQueue[bytes] = SimpleQueue() + self._delay: Final = delay_seconds + self._action: Final = action + self._server: ThreadingHTTPServer | None = None + self._thread: threading.Thread | None = None + self._port: Final = self._claim_port() + self.start() + + def _claim_port(self) -> int: + with socket.socket() as probe: + probe.bind(("127.0.0.1", 0)) + return probe.getsockname()[1] + + @property + def url(self) -> str: + return f"http://127.0.0.1:{self._port}" + + def start(self) -> None: + received = self.received + delay = self._delay + action = self._action + + class Handler(BaseHTTPRequestHandler): + def do_POST(self) -> None: + body: Final = self.rfile.read(int(self.headers.get("content-length", "0"))) + received.put(body) + if delay: + time.sleep(delay) + payload: Final = json.dumps({"action": action, "blocked_reason": "synthetic policy denial"}).encode() + self.send_response(200) + self.send_header("content-type", "application/json") + self.send_header("content-length", str(len(payload))) + self.end_headers() + self.wfile.write(payload) + + def log_message(self, format: str, *args: object) -> None: + pass + + class Server(ThreadingHTTPServer): + daemon_threads = True + allow_reuse_address = True + + self._server = Server(("127.0.0.1", self._port), Handler) + self._thread = threading.Thread(target=self._server.serve_forever, kwargs={"poll_interval": 0.05}) + self._thread.start() + + def stop(self) -> None: + assert self._server is not None and self._thread is not None + self._server.shutdown() + self._server.server_close() + self._thread.join(timeout=6) + assert not self._thread.is_alive() + self._server = None + + def drain(self) -> tuple[bytes, ...]: + return tuple(self.received.get_nowait() for _ in range(self.received.qsize())) + + def __enter__(self) -> "_GuardrailSink": + return self + + def __exit__(self, *exc_info: object) -> None: + if self._server is not None: + self.stop() + + +def test_concurrent_flag_writes_split_expected_outcomes(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {}) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + alias: Final = "audit-concurrent-" + uuid.uuid4().hex + + bodies: Final = [ + {"team_id": team, "models": [model], "key_alias": f"{alias}-{index}", "disable_global_guardrails": flag} + for index in range(20) + for flag in (True, False) + ] + with ThreadPoolExecutor(max_workers=20) as pool: + responses: Final = tuple( + pool.map(lambda body: candidate.request("POST", "/key/generate", body, key=caller), bodies) + ) + created_aliases: Final = [ + row["key_alias"] + for row in read_rows( + 'SELECT key_alias FROM "LiteLLM_VerificationToken" WHERE key_alias LIKE %s', (f"{alias}-%",) + ) + ] + for response in responses: + if response.status_code == 200: + scenario.cleanups.callback(scenario.delete_key, string_value(response.json()["key"])) + flagged: Final = tuple( + response for response, body in zip(responses, bodies) if body["disable_global_guardrails"] is True + ) + flagless: Final = tuple( + response for response, body in zip(responses, bodies) if body["disable_global_guardrails"] is False + ) + assert sorted(response.status_code for response in flagged) == [403] * 20, [ + response.text for response in flagged + ] + assert sorted(response.status_code for response in flagless) == [200] * 20, [ + response.text for response in flagless + ] + assert len(created_aliases) == 20, created_aliases + for entry in created_aliases: + stored: Final = read_rows('SELECT metadata FROM "LiteLLM_VerificationToken" WHERE key_alias = %s', (entry,)) + assert stored[0]["metadata"].get("disable_global_guardrails") is not True, entry + + +def test_revoked_exemption_denies_later_non_admin_resave(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {}) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + exempt: Final = scenario.key(team_id=team, models=[model], disable_global_guardrails=True) + assert stored_metadata(exempt)["disable_global_guardrails"] is True + + candidate.post("/key/update", {"key": exempt, "disable_global_guardrails": False}) + assert stored_metadata(exempt)["disable_global_guardrails"] is False + + resave: Final = candidate.request( + "POST", + "/key/update", + { + "key": exempt, + "key_alias": "audit-revoked-" + uuid.uuid4().hex, + "metadata": {"disable_global_guardrails": True}, + }, + key=caller, + ) + assert resave.status_code == 403, resave.text + assert "disable_global_guardrails" in resave.text, resave.text + assert stored_metadata(exempt)["disable_global_guardrails"] is False + + +def test_revoked_exemption_blocks_chats_on_both_workers(gateway: Gateway, tmp_path: Path) -> None: + with _GuardrailSink() as sink: + config: Final = guardrail_config(sink.url, tmp_path / "revoke-workers.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as first: + with owned_proxy(gateway, tmp_path, {}, config=config) as second: + with first.scenario() as scenario: + model: Final = scenario.model() + exempt: Final = scenario.key(models=[model], disable_global_guardrails=True) + for worker in (first, second): + served: Final = chat(worker, model, exempt, "audit-both-" + uuid.uuid4().hex) + assert served.status_code == 200, served.text + first.post("/key/update", {"key": exempt, "disable_global_guardrails": False}) + for worker in (first, second): + denied: Final = eventually( + lambda w=worker: chat(w, model, exempt, "audit-both-" + uuid.uuid4().hex), + lambda response: response.status_code == 400 and "synthetic policy denial" in response.text, + seconds=70, + ) + assert denied.status_code == 400, denied.text + + +def test_exempt_burst_survives_guardrail_sink_outage(gateway: Gateway, tmp_path: Path) -> None: + with _GuardrailSink() as sink: + config: Final = guardrail_config(sink.url, tmp_path / "sink-outage.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + exempt: Final = scenario.key(models=[model], disable_global_guardrails=True) + plain: Final = scenario.key(models=[model]) + + warm: Final = chat(candidate, model, plain, "warm-" + uuid.uuid4().hex) + assert warm.status_code == 400 and "synthetic policy denial" in warm.text, warm.text + assert sink.drain() != () + + def burst(keys: tuple[str, ...], tag: str) -> tuple[httpx.Response, ...]: + with ThreadPoolExecutor(max_workers=15) as pool: + return tuple( + pool.map( + lambda pair: chat(candidate, model, pair[1], f"{tag}-{pair[0]}-{uuid.uuid4().hex}"), + enumerate(keys * 10), + ) + ) + + outage_keys: Final = (exempt, plain) + with ThreadPoolExecutor(max_workers=2) as pool: + bursts: Final = pool.submit(burst, outage_keys, "outage") + eventually( + lambda: sink.received.qsize(), + lambda count: count >= 2, + seconds=30, + ) + sink.stop() + outage_responses: Final = bursts.result(timeout=90) + exempt_outage: Final = [response for index, response in enumerate(outage_responses) if index % 2 == 0] + non_exempt_outage: Final = [response for index, response in enumerate(outage_responses) if index % 2 == 1] + assert all(response.status_code == 200 for response in exempt_outage), [ + response.status_code for response in exempt_outage + ] + outage_statuses: Final = {response.status_code for response in non_exempt_outage} + assert outage_statuses <= {400, 500}, outage_statuses + assert all( + "synthetic policy denial" in response.text or response.status_code == 500 + for response in non_exempt_outage + ), [response.text for response in non_exempt_outage if response.status_code not in {400, 500}] + assert all(upstream_hits(gateway, f"outage-{index}-") == 0 for index in range(1, 20, 2)), ( + upstream_observations(gateway) + ) + + sink.start() + recovered: Final = chat(candidate, model, plain, "recovered-" + uuid.uuid4().hex) + assert recovered.status_code == 400 and "synthetic policy denial" in recovered.text, recovered.text + + +def test_flag_denial_survives_worker_kill(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, {}, workers=2) as owned: + candidate: Final = owned.gateway + with candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team(models=[model], members_with_roles=[{"role": "admin", "user_id": member}]) + caller: Final = non_admin_caller(scenario, member, team, model) + alias: Final = "audit-kill-" + uuid.uuid4().hex + + workers: Final = eventually( + lambda: psutil.Process(owned.process.pid).children(recursive=True), + lambda children: len(children) >= 2, + seconds=30, + ) + victim: Final = workers[0] + os.kill(victim.pid, signal.SIGKILL) + + probe: Final = eventually( + lambda: candidate.request( + "POST", + "/key/generate", + {"team_id": team, "models": [model], "key_alias": f"{alias}-probe"}, + key=caller, + ), + lambda response: response.status_code in (200, 403), + seconds=30, + ) + if probe.status_code == 200: + scenario.cleanups.callback(scenario.delete_key, string_value(probe.json()["key"])) + for index in range(10): + denied: Final = candidate.request( + "POST", + "/key/generate", + { + "team_id": team, + "models": [model], + "key_alias": f"{alias}-{index}", + "disable_global_guardrails": True, + }, + key=caller, + ) + if denied.status_code == 200: + scenario.cleanups.callback(scenario.delete_key, string_value(denied.json()["key"])) + assert denied.status_code == 403, denied.text + assert "disable_global_guardrails" in denied.text, denied.text + assert ( + read_rows( + 'SELECT token FROM "LiteLLM_VerificationToken" WHERE key_alias LIKE %s AND metadata::text LIKE %s', + (f"{alias}-%", '%"disable_global_guardrails": true%'), + ) + == [] + ) + + +def test_exempt_chats_do_not_wait_on_slow_guardrail_sink(gateway: Gateway, tmp_path: Path) -> None: + sink_delay: Final = 10.0 + with _GuardrailSink(delay_seconds=sink_delay) as sink: + config: Final = guardrail_config(sink.url, tmp_path / "slow-sink.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + exempt: Final = scenario.key(models=[model], disable_global_guardrails=True) + + started: Final = time.monotonic() + with ThreadPoolExecutor(max_workers=10) as pool: + responses: Final = tuple( + pool.map( + lambda index: chat(candidate, model, exempt, f"slow-sink-{index}-{uuid.uuid4().hex}"), + range(10), + ) + ) + elapsed: Final = time.monotonic() - started + assert all(response.status_code == 200 for response in responses), [ + (response.status_code, response.text) for response in responses + ] + assert elapsed < sink_delay, f"exempt chats waited on the guardrail sink: {elapsed}s" + assert sink.drain() == () diff --git a/tests/integration/authorization/test_key_guardrail_opt_out_runtime.py b/tests/integration/authorization/test_key_guardrail_opt_out_runtime.py new file mode 100644 index 00000000000..dc549d3be8d --- /dev/null +++ b/tests/integration/authorization/test_key_guardrail_opt_out_runtime.py @@ -0,0 +1,230 @@ +import asyncio +import json +import uuid +from pathlib import Path +from typing import Final + +import httpx +from anthropic import Anthropic +from openai import AsyncOpenAI, OpenAI + +from integration._support.client import Gateway, eventually, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server +from integration.authorization._guardrail_opt_out import ( + chat, + denying_guardrail, + guardrail_config, + stored_metadata, + upstream_hits, +) + + +def _wire_hits(wire: Wire, marker: str) -> int: + return sum(1 for request in wire.drain() if marker.encode() in request.body) + + +def _sink_hits(policy: Wire, marker: str) -> int: + return sum(1 for request in policy.drain() if marker.encode() in request.body) + + +def _anthropic_provider(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages", request.target + body: Final = json.loads(request.body) + if body.get("stream") is True: + identity: Final = "msg_" + uuid.uuid4().hex + frames: Final = ( + { + "type": "message_start", + "message": { + "id": identity, + "type": "message", + "role": "assistant", + "model": body["model"], + "content": [], + "stop_reason": None, + "usage": {"input_tokens": 10, "output_tokens": 1}, + }, + }, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "synthetic"}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 4}}, + {"type": "message_stop"}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {frame['type']}\ndata: {json.dumps(frame)}\n\n".encode() for frame in frames), + ) + return Reply( + body=json.dumps( + { + "id": "msg_" + uuid.uuid4().hex, + "type": "message", + "role": "assistant", + "model": body["model"], + "content": [{"type": "text", "text": "synthetic"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 4}, + } + ).encode() + ) + + +def _messages(candidate: Gateway, model: str, key: str, marker: str, *, stream: bool) -> httpx.Response: + return candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "messages": [{"role": "user", "content": marker}], + "max_tokens": 16, + "stream": stream, + }, + key=key, + ) + + +def _responses(candidate: Gateway, model: str, key: str, marker: str, *, stream: bool) -> httpx.Response: + return candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": marker, "stream": stream}, + key=key, + ) + + +def test_guardrail_denies_non_exempt_key_on_all_surfaces(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy, wire_server(_anthropic_provider) as anthropic_wire: + config: Final = guardrail_config(policy.url, tmp_path / "denied.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + openai_model: Final = scenario.model() + claude_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=anthropic_wire.url, + api_key="synthetic-anthropic-key", + ) + deepseek_model: Final = scenario.model(model="deepseek/gpt-4o-mini", api_base=gateway.upstream_url + "/v1") + key: Final = scenario.key(models=[openai_model, claude_model, deepseek_model]) + surfaces: Final = ( + ("chat", openai_model, chat), + ("messages", claude_model, _messages), + ("responses", deepseek_model, _responses), + ) + for surface, model, call in surfaces: + for stream in (False, True): + marker: Final = f"denied-{surface}-{stream}-{uuid.uuid4().hex}" + response: Final = call(candidate, model, key, marker, stream=stream) + response.read() + assert response.status_code == 400, f"{surface} stream={stream}: {response.text}" + assert "synthetic policy denial" in response.text, response.text + assert _sink_hits(policy, marker) == 1 + assert upstream_hits(gateway, marker) == 0 + assert _wire_hits(anthropic_wire, marker) == 0 + + +def test_guardrail_skipped_for_admin_exempt_key_on_all_surfaces_and_clients(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy, wire_server(_anthropic_provider) as anthropic_wire: + config: Final = guardrail_config(policy.url, tmp_path / "exempt.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + openai_model: Final = scenario.model() + claude_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=anthropic_wire.url, + api_key="synthetic-anthropic-key", + ) + deepseek_model: Final = scenario.model(model="deepseek/gpt-4o-mini", api_base=gateway.upstream_url + "/v1") + exempt: Final = scenario.key( + models=[openai_model, claude_model, deepseek_model], disable_global_guardrails=True + ) + assert stored_metadata(exempt)["disable_global_guardrails"] is True + surfaces: Final = ( + ("chat", openai_model, chat), + ("messages", claude_model, _messages), + ("responses", deepseek_model, _responses), + ) + for surface, model, call in surfaces: + for stream in (False, True): + marker: Final = f"exempt-{surface}-{stream}-{uuid.uuid4().hex}" + response: Final = call(candidate, model, exempt, marker, stream=stream) + response.read() + assert response.status_code == 200, f"{surface} stream={stream}: {response.text}" + assert "synthetic policy denial" not in response.text + provider_hits: Final = ( + _wire_hits(anthropic_wire, marker) if surface == "messages" else upstream_hits(gateway, marker) + ) + assert provider_hits == 1, f"{surface} stream={stream} marker={marker}" + assert _sink_hits(policy, marker) == 0 + + base_url: Final = str(candidate.client.base_url).rstrip("/") + "/v1" + sync_marker: Final = "exempt-sdk-sync-" + uuid.uuid4().hex + OpenAI(api_key=exempt, base_url=base_url, max_retries=0).chat.completions.create( + model=openai_model, messages=[{"role": "user", "content": sync_marker}] + ) + assert upstream_hits(gateway, sync_marker) == 1 + + async_marker: Final = "exempt-sdk-async-" + uuid.uuid4().hex + + async def _asyncchat() -> None: + async with AsyncOpenAI(api_key=exempt, base_url=base_url, max_retries=0) as client: + await client.chat.completions.create( + model=openai_model, messages=[{"role": "user", "content": async_marker}] + ) + + asyncio.run(_asyncchat()) + assert upstream_hits(gateway, async_marker) == 1 + + anthropic_marker: Final = "exempt-anthropic-" + uuid.uuid4().hex + Anthropic( + api_key=exempt, base_url=str(candidate.client.base_url).rstrip("/"), max_retries=0 + ).messages.create( + model=claude_model, max_tokens=16, messages=[{"role": "user", "content": anthropic_marker}] + ) + assert _wire_hits(anthropic_wire, anthropic_marker) == 1 + + +def test_team_flag_resaved_key_and_spend_log(gateway: Gateway, tmp_path: Path) -> None: + with wire_server(denying_guardrail) as policy: + config: Final = guardrail_config(policy.url, tmp_path / "team-exempt.yaml") + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model() + member: Final = scenario.user(user_role="internal_user") + + exempt_team: Final = scenario.team(models=[model], disable_global_guardrails=True) + team_key: Final = scenario.key(team_id=exempt_team, models=[model]) + team_marker: Final = "team-exempt-" + uuid.uuid4().hex + team_response: Final = chat(candidate, model, team_key, team_marker, stream=False) + assert team_response.status_code == 200, team_response.text + assert upstream_hits(gateway, team_marker) == 1 + assert _sink_hits(policy, team_marker) == 0 + + caller_team: Final = scenario.team( + models=[model], members_with_roles=[{"role": "admin", "user_id": member}] + ) + admin_exempt: Final = scenario.key(team_id=caller_team, models=[model], disable_global_guardrails=True) + resave_caller: Final = scenario.key( + user_id=member, team_id=caller_team, models=[model], allowed_routes=["/key/*", "/v1/chat/completions"] + ) + resaved: Final = candidate.request( + "POST", + "/key/update", + { + "key": admin_exempt, + "key_alias": "audit-runtime-resave-" + uuid.uuid4().hex, + "metadata": {"disable_global_guardrails": True}, + }, + key=resave_caller, + ) + assert resaved.status_code == 200, resaved.text + resave_marker: Final = "resaved-exempt-" + uuid.uuid4().hex + resave_response: Final = chat(candidate, model, admin_exempt, resave_marker, stream=False) + assert resave_response.status_code == 200, resave_response.text + response_id: Final = string_value(resave_response.json()["id"]) + assert upstream_hits(gateway, resave_marker) == 1 + assert _sink_hits(policy, resave_marker) == 0 + eventually( + lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (response_id,)), + lambda rows: len(rows) == 1, + seconds=70, + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index 2b614632346..69013408962 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -774,6 +774,135 @@ class TestCheckPassthroughRoutesCallerPermission: ) +class TestCheckDisableGlobalGuardrailsCallerPermission: + """Only proxy admins may set disable_global_guardrails (top-level or under + metadata); non-admins get a 403 naming the entity.""" + + def _non_admin(self): + return UserAPIKeyAuth( + user_id="u1", api_key="sk-x", user_role=LitellmUserRoles.INTERNAL_USER + ) + + def _admin(self): + return UserAPIKeyAuth( + user_id="u2", api_key="sk-y", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + def test_top_level_flag_rejected_with_default_entity(self): + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + with pytest.raises(HTTPException) as exc_info: + _check_disable_global_guardrails_caller_permission(True, None, self._non_admin()) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a key."} + + def test_metadata_flag_rejected_with_default_entity(self): + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + with pytest.raises(HTTPException) as exc_info: + _check_disable_global_guardrails_caller_permission( + None, {"disable_global_guardrails": True}, self._non_admin() + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a key."} + + def test_explicit_false_with_metadata_true_is_rejected(self): + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + with pytest.raises(HTTPException) as exc_info: + _check_disable_global_guardrails_caller_permission( + False, {"disable_global_guardrails": True}, self._non_admin() + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a key."} + + def test_rejection_names_the_team_entity(self): + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + with pytest.raises(HTTPException) as exc_info: + _check_disable_global_guardrails_caller_permission(True, None, self._non_admin(), entity="team") + + assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a team."} + + def test_false_and_absent_flag_do_not_raise(self): + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + non_admin = self._non_admin() + assert _check_disable_global_guardrails_caller_permission(False, None, non_admin) is None + assert _check_disable_global_guardrails_caller_permission(None, None, non_admin) is None + assert _check_disable_global_guardrails_caller_permission(None, {}, non_admin) is None + assert ( + _check_disable_global_guardrails_caller_permission(None, {"disable_global_guardrails": False}, non_admin) + is None + ) + + def test_unchanged_stored_flag_does_not_raise(self): + """Re-sending a flag that is already stored is not an opt-out.""" + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + non_admin = self._non_admin() + assert ( + _check_disable_global_guardrails_caller_permission( + True, + {"disable_global_guardrails": True}, + non_admin, + existing_metadata={"disable_global_guardrails": True}, + ) + is None + ) + + def test_stored_false_does_not_exempt(self): + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + with pytest.raises(HTTPException) as exc_info: + _check_disable_global_guardrails_caller_permission( + True, + None, + self._non_admin(), + existing_metadata={"disable_global_guardrails": False}, + ) + + assert exc_info.value.status_code == 403 + assert exc_info.value.detail == {"error": "Only proxy admins can set `disable_global_guardrails` on a key."} + + def test_proxy_admin_may_set_the_flag(self): + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + assert ( + _check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, self._admin()) + is None + ) + + class TestIsUserOrgAdminForTeam: """The caller must be looked up with its exact identity; a nulled or omitted lookup argument would silently mis-resolve org-admin status.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 69f66ca3939..3b86f1f6d20 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -18020,6 +18020,199 @@ async def test_regenerate_key_non_admin_permissions_rejected_before_enterprise_g assert "Enterprise" not in str(exc.value.message) +@pytest.mark.asyncio +async def test_generate_key_non_admin_disable_global_guardrails_rejected(monkeypatch): + """`_common_key_generation_helper` rejects a non-admin setting + `disable_global_guardrails` on the request body.""" + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.litellm.default_key_generate_params", + None, + raising=False, + ) + caller = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user-1", + max_budget=100.0, + ) + request = GenerateKeyRequest(disable_global_guardrails=True) + with pytest.raises(HTTPException) as exc_info: + await _common_key_generation_helper( + data=request, + user_api_key_dict=caller, + litellm_changed_by=None, + team_table=None, + ) + assert exc_info.value.status_code == 403 + assert "disable_global_guardrails" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_generate_key_non_admin_metadata_disable_global_guardrails_rejected(monkeypatch): + """`_common_key_generation_helper` rejects a non-admin smuggling + `disable_global_guardrails` under `metadata`.""" + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.litellm.default_key_generate_params", + None, + raising=False, + ) + caller = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user-1", + max_budget=100.0, + ) + request = GenerateKeyRequest(metadata={"disable_global_guardrails": True}) + with pytest.raises(HTTPException) as exc_info: + await _common_key_generation_helper( + data=request, + user_api_key_dict=caller, + litellm_changed_by=None, + team_table=None, + ) + assert exc_info.value.status_code == 403 + assert "disable_global_guardrails" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_generate_key_non_admin_server_default_guardrail_flag_not_treated_as_requested(monkeypatch): + """An admin-configured `default_key_generate_params.metadata` containing + `disable_global_guardrails: true` must not 403 a non-admin who sent no flag; + only caller-sent metadata counts as requesting the opt-out.""" + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.litellm.default_key_generate_params", + {"metadata": {"disable_global_guardrails": True}}, + raising=False, + ) + caller = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user-1", + max_budget=100.0, + ) + + raised: Exception | None = None + try: + await _common_key_generation_helper( + data=GenerateKeyRequest(team_id="team-1", models=["gpt-4o"]), + user_api_key_dict=caller, + litellm_changed_by=None, + team_table=None, + ) + except Exception as exc: + raised = exc + assert not (isinstance(raised, HTTPException) and "disable_global_guardrails" in str(raised.detail)), raised + + with pytest.raises(HTTPException) as exc_info: + await _common_key_generation_helper( + data=GenerateKeyRequest( + team_id="team-1", + models=["gpt-4o"], + metadata={"disable_global_guardrails": True}, + ), + user_api_key_dict=caller, + litellm_changed_by=None, + team_table=None, + ) + assert exc_info.value.status_code == 403 + assert "disable_global_guardrails" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_update_key_non_admin_disable_global_guardrails_rejected(monkeypatch): + """`_validate_update_key_data` rejects a non-admin when + `disable_global_guardrails` is true in the request body.""" + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + data = UpdateKeyRequest( + key="sk-alice-personal", + disable_global_guardrails=True, + ) + + with pytest.raises(HTTPException) as exc: + await _validate_update_key_data( + data=data, + existing_key_row=_make_personal_key_row_for_alice(), + user_api_key_dict=_make_alice_internal_user(), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=MagicMock(), + ) + assert exc.value.status_code == 403 + assert "disable_global_guardrails" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_update_key_non_admin_resending_stored_disable_global_guardrails_allowed(monkeypatch): + """`_validate_update_key_data` must not 403 when a non-admin edit form + re-sends `metadata.disable_global_guardrails` that is already stored on + the key (the Admin UI edit form round-trips the whole metadata JSON).""" + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = lambda data: data + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + existing_key_row = _make_personal_key_row_for_alice() + existing_key_row.metadata = {"disable_global_guardrails": True} + data = UpdateKeyRequest( + key="sk-alice-personal", + metadata={"disable_global_guardrails": True, "x": 1}, + ) + + raised: HTTPException | None = None + try: + await _validate_update_key_data( + data=data, + existing_key_row=existing_key_row, + user_api_key_dict=_make_alice_internal_user(), + llm_router=None, + premium_user=True, + prisma_client=mock_prisma_client, + user_api_key_cache=MagicMock(), + ) + except HTTPException as exc: + raised = exc + assert raised is None or "disable_global_guardrails" not in str(raised.detail) + + +@pytest.mark.asyncio +async def test_regenerate_key_non_admin_disable_global_guardrails_rejected(monkeypatch): + """`regenerate_key_fn` rejects a non-admin setting + `disable_global_guardrails` once the stored key row is loaded (the + already-stored exemption check needs the row's metadata).""" + from litellm.proxy._types import RegenerateKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import ( + regenerate_key_fn, + ) + + existing_key = _make_regenerate_existing_key() + mock_prisma_client = AsyncMock() + mock_repo = MagicMock() + mock_repo.table.find_unique = AsyncMock(return_value=existing_key) + + data = RegenerateKeyRequest( + key="sk-alice-personal", + disable_global_guardrails=True, + ) + + with ( + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.VerificationTokenRepository", + return_value=mock_repo, + ), + pytest.raises(ProxyException) as exc, + ): + await regenerate_key_fn( + key=None, + data=data, + user_api_key_dict=_make_alice_internal_user(), + litellm_changed_by=None, + ) + assert int(exc.value.code) == 403 + assert "disable_global_guardrails" in str(exc.value.message) + + def test_generate_key_helper_fn_accepts_per_tag_rate_limits(): """ Regression: new_user / SSO sign-in forward NewUserRequest fields to diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 7a9b6b66946..b066b3b80e6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -11592,6 +11592,83 @@ async def test_update_team_blocks_non_admin_passthrough_routes(mock_db_client): assert "allowed_passthrough_routes" in str(exc.value.message) +def test_check_disable_global_guardrails_caller_permission_team(): + from litellm.proxy._types import NewTeamRequest + from litellm.proxy.management_endpoints.common_utils import ( + _check_disable_global_guardrails_caller_permission, + ) + + admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + non_admin = _non_admin_auth() + + _check_disable_global_guardrails_caller_permission(True, {"disable_global_guardrails": True}, admin, entity="team") + _check_disable_global_guardrails_caller_permission(None, None, non_admin, entity="team") + _check_disable_global_guardrails_caller_permission(False, None, non_admin, entity="team") + + with pytest.raises(HTTPException) as exc: + _check_disable_global_guardrails_caller_permission(True, None, non_admin, entity="team") + assert exc.value.status_code == 403 + assert "disable_global_guardrails" in str(exc.value.detail) + assert "team" in str(exc.value.detail) + + with pytest.raises(HTTPException) as exc: + _check_disable_global_guardrails_caller_permission( + None, {"disable_global_guardrails": True}, non_admin, entity="team" + ) + assert exc.value.status_code == 403 + assert "disable_global_guardrails" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_new_team_blocks_non_admin_disable_global_guardrails(mock_db_client): + """A non-proxy-admin cannot opt a team out of global guardrails via /team/new.""" + mock_db_client.db.litellm_teamtable.count = AsyncMock(return_value=0) + from fastapi import Request + + from litellm.proxy._types import NewTeamRequest, ProxyException + from litellm.proxy.management_endpoints.team_endpoints import new_team + + with patch( + "litellm.proxy.management_endpoints.team_endpoints._check_user_team_limits", + AsyncMock(return_value=None), + ): + with pytest.raises(ProxyException) as exc: + await new_team( + data=NewTeamRequest(team_alias="t", disable_global_guardrails=True), + http_request=MagicMock(spec=Request), + user_api_key_dict=_non_admin_auth(), + ) + assert str(exc.value.code) == "403" + assert "disable_global_guardrails" in str(exc.value.message) + + +@pytest.mark.asyncio +async def test_update_team_blocks_non_admin_disable_global_guardrails(mock_db_client): + """Even a team manager (non-proxy-admin) cannot set + disable_global_guardrails via /team/update.""" + from fastapi import Request + + from litellm.proxy._types import ProxyException, UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import update_team + + existing = MagicMock() + existing.model_dump.return_value = {"team_id": "t1"} + mock_db_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing) + + with patch( + "litellm.proxy.management_endpoints.team_endpoints._resolve_team_access", + AsyncMock(return_value="org_admin"), + ): + with pytest.raises(ProxyException) as exc: + await update_team( + data=UpdateTeamRequest(team_id="t1", disable_global_guardrails=True), + http_request=MagicMock(spec=Request), + user_api_key_dict=_non_admin_auth(), + ) + assert str(exc.value.code) == "403" + assert "disable_global_guardrails" in str(exc.value.message) + + def test_set_budget_reset_at_clears_when_budget_duration_null(): """ When budget_duration is explicitly set to null, _set_budget_reset_at diff --git a/ui/litellm-dashboard/src/components/Teams.test.tsx b/ui/litellm-dashboard/src/components/Teams.test.tsx index b6c6bbbcdd9..5ff95d2af0c 100644 --- a/ui/litellm-dashboard/src/components/Teams.test.tsx +++ b/ui/litellm-dashboard/src/components/Teams.test.tsx @@ -1779,3 +1779,50 @@ describe("Teams - the create form keeps the organization and models picks while expect(modelsField()).toHaveValue(""); }); }); + +describe("Teams - disable_global_guardrails switch gating", () => { + const openCreateModal = async () => { + act(() => { + fireEvent.click(screen.getAllByRole("button", { name: /create team/i })[0]); + }); + await screen.findByLabelText(/team name/i); + }; + + beforeEach(() => { + vi.clearAllMocks(); + mockTeamInfoView.mockClear(); + vi.mocked(fetchAvailableModelsForTeamOrKey).mockResolvedValue(["gpt-4"]); + vi.mocked(fetchMCPAccessGroups).mockResolvedValue([]); + vi.mocked(getGuardrailsList).mockResolvedValue({ guardrails: [] }); + vi.mocked(getDefaultTeamSettings).mockResolvedValue({ values: {} }); + mockUseOrganizations.mockReturnValue({ data: null }); + }); + + it("hides the Disable Global Guardrails switch from a non-admin", async () => { + mockUseOrganizations.mockReturnValue({ + data: [ + { + organization_id: "org-1", + organization_alias: "Org 1", + models: [], + members: [{ user_id: "user-123", user_role: "org_admin" }], + }, + ], + }); + renderWithQueryClient(); + await openCreateModal(); + + fireEvent.click(screen.getByText("Additional Settings")); + + expect(screen.queryByRole("switch", { name: /Disable Global Guardrails/i })).not.toBeInTheDocument(); + }); + + it("shows the Disable Global Guardrails switch to a proxy admin", async () => { + renderWithQueryClient(); + await openCreateModal(); + + fireEvent.click(screen.getByText("Additional Settings")); + + expect(await screen.findByRole("switch", { name: /Disable Global Guardrails/i })).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Teams.tsx b/ui/litellm-dashboard/src/components/Teams.tsx index 7214d16f665..c2a23cef83a 100644 --- a/ui/litellm-dashboard/src/components/Teams.tsx +++ b/ui/litellm-dashboard/src/components/Teams.tsx @@ -983,29 +983,31 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser /> )} - - {({ id, value, onChange }) => ( - - )} - + {isProxyAdminRole(userRole || "") && ( + + {({ id, value, onChange }) => ( + + )} + + )} {canViewPolicies && ( { expect((await createdPayload()).disable_global_guardrails).toBe(true); }); + it("hides the disable_global_guardrails switch from a non-admin", async () => { + state.authorized = { ...state.authorized, userRole: "Internal User" }; + await openModal(); + await openSection(/Optional Settings/i); + + expect(screen.queryByRole("switch", { name: /Disable Global Guardrails/i })).not.toBeInTheDocument(); + }); + + it("shows the disable_global_guardrails switch to a proxy admin", async () => { + await openModal(); + await openSection(/Optional Settings/i); + + expect(await screen.findByRole("switch", { name: /Disable Global Guardrails/i })).toBeInTheDocument(); + }); + it("folds a metadata JSON string back through JSON.stringify", async () => { await openModal(); await nameTheKey(); diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index 25b986e4c9e..45245ff95b3 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -1293,40 +1293,42 @@ const CreateKey: React.FC = ({ team, teams, data, addKey, autoOp /> )} - - Disable Global Guardrails{" "} - - e.stopPropagation()} // Prevent accordion from collapsing when clicking link - > - - - - - } - name="disable_global_guardrails" - className="mt-4" - help={ - canEditGuardrails - ? "Bypass global guardrails for this key" - : "Premium feature - Upgrade to disable global guardrails by key" - } - > - {(control) => ( - - )} - + {userRole != null && isProxyAdminRole(userRole) && ( + + Disable Global Guardrails{" "} + + e.stopPropagation()} // Prevent accordion from collapsing when clicking link + > + + + + + } + name="disable_global_guardrails" + className="mt-4" + help={ + canEditGuardrails + ? "Bypass global guardrails for this key" + : "Premium feature - Upgrade to disable global guardrails by key" + } + > + {(control) => ( + + )} + + )} {canViewPolicies && ( { errorToast.mockRestore(); }); }); + +describe("TeamInfoView - disable_global_guardrails switch gating", () => { + beforeEach(() => { + seedDefaultMocks(); + vi.mocked(networking.teamInfoCall).mockResolvedValue(createMockTeamData()); + }); + + afterEach(() => { + vi.clearAllMocks(); + authState.userRole = "Admin"; + }); + + const props = { + teamId: "123", + onUpdate: vi.fn(), + onClose: vi.fn(), + accessToken: "test-token", + is_team_admin: true, + is_proxy_admin: true, + userModels: ["gpt-4"], + editTeam: false, + premiumUser: false, + }; + + const openEditForm = async () => { + const user = userEvent.setup({ delay: null }); + await waitFor(() => expect(screen.queryAllByText("Test Team").length).toBeGreaterThan(0)); + await user.click(screen.getByRole("tab", { name: "Settings" })); + await user.click(await screen.findByRole("button", { name: /edit settings/i })); + await screen.findByLabelText("Team Name"); + }; + + it("hides the Disable all global guardrails switch from a non-admin", async () => { + authState.userRole = "Internal User"; + renderWithProviders(); + await openEditForm(); + + expect(screen.queryByRole("switch", { name: /Disable all global guardrails/i })).not.toBeInTheDocument(); + }); + + it("shows the Disable all global guardrails switch to a proxy admin", async () => { + renderWithProviders(); + await openEditForm(); + + expect(await screen.findByRole("switch", { name: /Disable all global guardrails/i })).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 22cc99b32c8..d9e308e9d6f 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -1785,25 +1785,27 @@ const TeamInfoView: React.FC = ({ )} - - {({ id, value, onChange }) => ( - { - onChange(checked); - applyKillSwitchToGuardrails(checked); - }} - /> - )} - + {is_proxy_admin && ( + + {({ id, value, onChange }) => ( + { + onChange(checked); + applyKillSwitchToGuardrails(checked); + }} + /> + )} + + )} {canViewPolicies && ( { }, ); }); + + describe("disable_global_guardrails toggle gating", () => { + const renderAs = (userRole: string) => + renderWithProviders( + {}} + onSubmit={async () => {}} + accessToken="test-token" + userID="test-user" + userRole={userRole} + premiumUser={true} + />, + ); + + it("hides the switch from a non-admin", async () => { + renderAs("Internal User"); + await screen.findByRole("button", { name: /save changes/i }); + + expect(screen.queryByRole("switch", { name: /disable global guardrails/i })).not.toBeInTheDocument(); + }); + + it("shows the switch to a proxy admin", async () => { + renderAs("Admin"); + + expect(await screen.findByRole("switch", { name: /disable global guardrails/i })).toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index c668958be74..9cd97f4ef98 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -618,18 +618,20 @@ export function KeyEditView({ } - - {({ value, onChange, ref: _ref, ...field }) => ( - - )} - + {userRole != null && isProxyAdminRole(userRole) && ( + + {({ value, onChange, ref: _ref, ...field }) => ( + + )} + + )} {canViewPolicies && (