diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3842a5cb504..44b30386e07 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -11585,6 +11585,18 @@ "description": "Authorization bearer token for IBM Guardrails API. Reads from IBM_GUARDRAILS_AUTH_TOKEN env var if None.", "title": "Auth Token" }, + "authority_host": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "Microsoft Entra authority host that issues the On-Behalf-Of token, for sovereign clouds. Defaults to https://login.microsoftonline.com. Falls back to the AGENT365_AUTHORITY_HOST, then AZURE_AUTHORITY_HOST environment variables.", + "title": "Authority Host" + }, "aws_access_key_id": { "anyOf": [ { diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py index ef4b46e3dfd..04bb00be7f1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/__init__.py @@ -2,6 +2,7 @@ from typing import TYPE_CHECKING, Final from litellm.types.guardrails import SupportedGuardrailIntegrations from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import ( + AGENT_365_DEFAULT_AUTHORITY_HOST, AGENT_365_PROD_API_BASE, AGENT_365_PROD_RESOURCE_APP_ID, ) @@ -23,6 +24,12 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" ) api_base: Final = litellm_params.api_base or get_secret_str("AGENT365_API_BASE") resource_app_id: Final = litellm_params.resource_app_id or get_secret_str("AGENT365_RESOURCE_APP_ID") + authority_host: Final = ( + litellm_params.authority_host + or get_secret_str("AGENT365_AUTHORITY_HOST") + or get_secret_str("AZURE_AUTHORITY_HOST") + or AGENT_365_DEFAULT_AUTHORITY_HOST + ) if not tenant_id: raise ValueError("Microsoft Agent 365: tenant_id is required") @@ -45,6 +52,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" api_base=api_base or AGENT_365_PROD_API_BASE, resource_app_id=resource_app_id or AGENT_365_PROD_RESOURCE_APP_ID, agent_id=litellm_params.agent_id, + authority_host=authority_host, request_timeout=litellm_params.timeout if litellm_params.timeout is not None else 10.0, unreachable_fallback=litellm_params.unreachable_fallback or "fail_open", event_hook=litellm_params.mode, diff --git a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py index 3d8007fc480..a25c4ba5cbd 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py +++ b/litellm/proxy/guardrails/guardrail_hooks/agent_365/agent_365.py @@ -22,7 +22,6 @@ from fastapi import HTTPException from pydantic import TypeAdapter, ValidationError from typing_extensions import ReadOnly, TypedDict -import litellm from litellm._logging import verbose_proxy_logger from litellm.exceptions import Timeout as LitellmTimeout from litellm.integrations.custom_guardrail import ( @@ -38,6 +37,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import ( + AGENT_365_DEFAULT_AUTHORITY_HOST, AGENT_365_PROD_API_BASE, AGENT_365_PROD_RESOURCE_APP_ID, AGENT_365_SCOPE_NAME, @@ -50,7 +50,7 @@ if TYPE_CHECKING: from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel from litellm.types.utils import GuardrailStatus -TOKEN_ENDPOINT_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token" +TOKEN_PATH_TEMPLATE: Final = "/{tenant_id}/oauth2/v2.0/token" EVALUATE_PATH: Final = "/agents/tool-evaluation/evaluate" MCP_SESSION_ID_HEADER: Final = "mcp-session-id" DEFENDER_STATUS_EVALUATED: Final = "Evaluated" @@ -83,10 +83,6 @@ def _parse_aadsts_codes(raw: object) -> tuple[int, ...]: return () -def registered_prometheus_logger() -> PrometheusLogger | None: - return next((cb for cb in litellm.callbacks if isinstance(cb, PrometheusLogger)), None) - - def entra_assertion(value: object) -> str | None: """``value`` when it is a compact JWS, the only bearer shape the OBO exchange accepts as its assertion. A LiteLLM virtual key, session bearer, or opaque upstream token in ``Authorization`` yields ``None``.""" @@ -162,10 +158,11 @@ class Agent365Guardrail(CustomGuardrail): api_base: str = AGENT_365_PROD_API_BASE, resource_app_id: str = AGENT_365_PROD_RESOURCE_APP_ID, agent_id: str | None = None, + authority_host: str = AGENT_365_DEFAULT_AUTHORITY_HOST, request_timeout: float = 10.0, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_open", async_handler: AsyncHTTPHandler | None = None, - prometheus_logger_lookup: Callable[[], PrometheusLogger | None] = registered_prometheus_logger, + prometheus_logger_lookup: Callable[[], PrometheusLogger | None] = PrometheusLogger.get_instance, **kwargs, # noqa: ANN003 # kwargs-ok: forwarded verbatim to CustomGuardrail (event_hook, default_on) ) -> None: super().__init__( @@ -181,6 +178,7 @@ class Agent365Guardrail(CustomGuardrail): self.api_base = api_base.rstrip("/") self.resource_app_id = resource_app_id self.agent_id = agent_id + self.authority_host = authority_host.rstrip("/") self.request_timeout = request_timeout self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( "fail_closed" if unreachable_fallback == "fail_closed" else "fail_open" @@ -456,7 +454,7 @@ class Agent365Guardrail(CustomGuardrail): return cached[0] response: Final = await self._post_allowing_error_status( - url=TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id), + url=f"{self.authority_host}{TOKEN_PATH_TEMPLATE.format(tenant_id=self.tenant_id)}", data={ # mutable-ok: OAuth form body; AsyncHTTPHandler.post requires dict "grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer", "client_id": self.client_id, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/agent_365.py b/litellm/types/proxy/guardrails/guardrail_hooks/agent_365.py index 83754df9c59..896df478577 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/agent_365.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/agent_365.py @@ -7,6 +7,7 @@ from .base import GuardrailConfigModel AGENT_365_PROD_API_BASE: Final = "https://agent365.svc.cloud.microsoft" AGENT_365_PROD_RESOURCE_APP_ID: Final = "ea9ffc3e-8a23-4a7d-836d-234d7c7565c1" AGENT_365_SCOPE_NAME: Final = "ThreatProtection.Evaluate.All" +AGENT_365_DEFAULT_AUTHORITY_HOST: Final = "https://login.microsoftonline.com" class Agent365GuardrailConfigModel(GuardrailConfigModel): @@ -61,13 +62,23 @@ class Agent365GuardrailConfigModel(GuardrailConfigModel): ), ) + authority_host: str | None = Field( + default=None, + description=( + "Microsoft Entra authority host that issues the On-Behalf-Of token, for sovereign clouds. " + f"Defaults to {AGENT_365_DEFAULT_AUTHORITY_HOST}. " + "Falls back to the AGENT365_AUTHORITY_HOST, then AZURE_AUTHORITY_HOST environment variables." + ), + ) + unreachable_fallback: Literal["fail_closed", "fail_open"] = Field( default="fail_open", description=( - "Behavior when Agent 365 or Entra is unreachable, times out, returns 5xx, or skips the evaluation. " - "'fail_open' (default) allows the tool call and records it as Unscanned in the logs, OpenTelemetry " - "and the litellm_guardrail_errors_total Prometheus counter. 'fail_closed' blocks it with HTTP 503. " - "Blocks, 4xx rejections, throttling and Entra token failures always block." + "Behavior when Agent 365 or Entra is unreachable, times out, returns 5xx, skips the evaluation, or " + "rejects the gateway's own client credentials. 'fail_open' (default) allows the tool call and records " + "it as Unscanned in the logs, OpenTelemetry and the litellm_guardrail_errors_total Prometheus counter. " + "'fail_closed' blocks it with HTTP 503. Policy blocks, 4xx rejections, throttling and a rejected " + "caller token always block." ), ) diff --git a/tests/integration/mcp/test_mcp_agent_365_guardrail.py b/tests/integration/mcp/test_mcp_agent_365_guardrail.py new file mode 100644 index 00000000000..53b9efb4818 --- /dev/null +++ b/tests/integration/mcp/test_mcp_agent_365_guardrail.py @@ -0,0 +1,160 @@ +"""Agent 365 guardrail when its own dependencies fail: Entra and Agent 365 are owned local doubles.""" + +import json +import uuid +from collections.abc import Iterator +from contextlib import contextmanager +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.mcp import McpCaller, McpPeer, echo_tool, register_mcp, scripted_peer, tool_calls +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from prometheus_client.parser import text_string_to_metric_families + +TENANT: Final = "00000000-0000-4000-8000-0000000a3650" +EVALUATE_PATH: Final = "/agents/tool-evaluation/evaluate" +GUARDRAIL_ERRORS: Final = "litellm_guardrail_errors_total" +CALLER_TOKEN: Final = "eyJhbGciOiJub25lIn0.eyJzdWIiOiJpbnRlZ3JhdGlvbiJ9.synthetic-signature" +GUARDRAIL_STATUSES: Final = ( + "SELECT metadata->'mcp_tool_call_metadata'->>'name' AS tool, gi->>'guardrail_status' AS status " + "FROM \"LiteLLM_SpendLogs\", jsonb_array_elements(metadata->'guardrail_information') gi " + "WHERE api_key = %s AND gi->>'guardrail_provider' = 'agent_365'" +) + + +def _entra(request: Request) -> Reply: + assert request.target == f"/{TENANT}/oauth2/v2.0/token", request.target + return Reply(body=json.dumps({"access_token": "obo-" + uuid.uuid4().hex, "expires_in": 3599}).encode()) + + +def _agent_365(request: Request) -> Reply: + assert request.target == EVALUATE_PATH, request.target + tool: Final = str(json.loads(request.body)["tool"]["name"]).rsplit("-", 1)[-1] + match tool: + case "outage": + return Reply(status=503, body=json.dumps({"error": "synthetic Agent 365 outage"}).encode()) + case "skipped": + return Reply(body=json.dumps({"allowed": True, "defender": {"status": "Skipped"}}).encode()) + case "denied": + verdict: Final = {"status": "Evaluated", "verdict": "Block", "message": "synthetic block"} + return Reply(body=json.dumps({"allowed": False, "defender": verdict, "correlationId": "denied-1"}).encode()) + case _: + return Reply(body=json.dumps({"allowed": True, "defender": {"status": "Evaluated"}}).encode()) + + +def _config(tmp_path: Path, name: str, entra_url: str, agent_365_url: str, fallback: str | None) -> Path: + config: dict = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"]["callbacks"] = ["prometheus"] + config["guardrails"] = [ + { + "guardrail_name": name, + "litellm_params": { + "guardrail": "agent_365", + "mode": "pre_mcp_call", + "default_on": True, + "tenant_id": TENANT, + "client_id": "synthetic-client-id", + "client_secret": "synthetic-client-secret", + "api_base": agent_365_url, + "authority_host": entra_url, + **({"unreachable_fallback": fallback} if fallback else {}), + }, + } + ] + path: Final = tmp_path / "agent_365.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@dataclass(frozen=True, slots=True) +class Rig: + candidate: Gateway + caller: McpCaller + key: str + alias: str + guardrail_name: str + peer: McpPeer + + +@contextmanager +def _rig(gateway: Gateway, tmp_path: Path, fallback: str | None) -> Iterator[Rig]: + alias: Final = "a365" + uuid.uuid4().hex[:8] + prom_dir: Final = tmp_path / "prom" + prom_dir.mkdir() + with ( + wire_server(_entra) as entra, + wire_server(_agent_365) as agent_365, + scripted_peer(echo_tool("outage"), echo_tool("skipped"), echo_tool("denied"), echo_tool("add")) as peer, + owned_proxy_process( + gateway, + tmp_path, + {"PROMETHEUS_MULTIPROC_DIR": str(prom_dir)}, + config=_config(tmp_path, alias, entra.url, agent_365.url, fallback), + ) as owned, + owned.gateway.scenario() as scenario, + ): + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + caller: Final = McpCaller(owned.gateway, key, "mcp", alias, headers={"Authorization": f"Bearer {CALLER_TOKEN}"}) + peer.drain() + yield Rig(owned.gateway, caller, key, alias, alias, peer) + + +def _guardrail_statuses(key: str) -> dict[str, str]: + rows: Final = read_rows(GUARDRAIL_STATUSES, (sha256(key.encode()).hexdigest(),)) + return {str(row["tool"]): str(row["status"]) for row in rows} + + +def _guardrail_error_counts(candidate: Gateway, guardrail_name: str) -> dict[str, float]: + response: Final = candidate.client.get( + "/metrics", headers={"Authorization": f"Bearer {candidate.key}"}, follow_redirects=True + ) + assert response.status_code == 200, f"GET /metrics: {response.status_code} {response.text[:300]}" + return { + sample.labels["error_type"]: float(sample.value) + for family in text_string_to_metric_families(response.text) + for sample in family.samples + if sample.name == GUARDRAIL_ERRORS and sample.labels.get("guardrail_name") == guardrail_name + } + + +def test_default_lets_the_call_through_unscanned_and_counts_it_when_agent_365_cannot_evaluate( + gateway: Gateway, tmp_path: Path +) -> None: + with _rig(gateway, tmp_path, fallback=None) as rig: + outage: Final = rig.caller.call(f"{rig.alias}-outage", {"a": 1}) + assert outage.text == '{"a": 1}', f"Agent 365 down must fail open by default: {outage.raw}" + skipped: Final = rig.caller.call(f"{rig.alias}-skipped", {"a": 2}) + assert skipped.text == '{"a": 2}', f"Defender skipping the call must fail open by default: {skipped.raw}" + denied: Final = rig.caller.call(f"{rig.alias}-denied", {"a": 3}) + assert denied.error is not None and "Blocked by Microsoft Defender" in denied.raw, denied.raw + assert tuple(call["body"]["params"]["name"] for call in tool_calls(rig.peer.drain())) == ("outage", "skipped") + statuses: Final = eventually(lambda: _guardrail_statuses(rig.key), lambda seen: len(seen) >= 3, seconds=70) + assert statuses == { + "outage": "guardrail_failed_to_respond", + "skipped": "guardrail_failed_to_respond", + "denied": "guardrail_intervened", + }, statuses + counts: Final = eventually( + lambda: _guardrail_error_counts(rig.candidate, rig.guardrail_name), lambda seen: len(seen) >= 2 + ) + assert counts == {"fail_open": 2.0, "HTTPException": 1.0}, counts + + +def test_explicit_fail_closed_blocks_with_503_and_never_reaches_upstream_when_agent_365_is_down( + gateway: Gateway, tmp_path: Path +) -> None: + with _rig(gateway, tmp_path, fallback="fail_closed") as rig: + outage: Final = rig.caller.call(f"{rig.alias}-outage", {"a": 1}) + assert outage.error is not None and "could not authorize the tool call" in outage.raw, outage.raw + assert rig.caller.call(f"{rig.alias}-add", {"a": 2}).text == '{"a": 2}' + assert tuple(call["body"]["params"]["name"] for call in tool_calls(rig.peer.drain())) == ("add",) + statuses: Final = eventually(lambda: _guardrail_statuses(rig.key), lambda seen: len(seen) >= 2, seconds=70) + assert statuses == {"outage": "guardrail_failed_to_respond", "add": "success"}, statuses + assert _guardrail_error_counts(rig.candidate, rig.guardrail_name) == {"HTTPException": 1.0} diff --git a/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py b/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py index 09ac4f3addf..f234d95bd7d 100644 --- a/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py +++ b/tests/test_litellm/integrations/test_prometheus_queue_guardrail_metrics.py @@ -567,3 +567,15 @@ class TestPrometheusGuardrailMetrics: ) is None ) + + def test_record_guardrail_fail_open_swallows_metric_errors(self): + prometheus_logger = PrometheusLogger() + broken_counter = MagicMock() + broken_counter.labels.side_effect = ValueError("registry exploded") + prometheus_logger.litellm_guardrail_errors_total = broken_counter + + prometheus_logger.record_guardrail_fail_open(guardrail_name="a365", hook_type="pre_call") + + broken_counter.labels.assert_called_once_with( + guardrail_name="a365", error_type="fail_open", hook_type="pre_call" + ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py index 7ccb30b960c..78871adbeda 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_agent_365.py @@ -1,3 +1,4 @@ +import logging import time import uuid from types import SimpleNamespace @@ -21,7 +22,6 @@ from litellm.proxy.guardrails.guardrail_hooks.agent_365 import ( guardrail_initializer_registry, initialize_guardrail, ) -from litellm.proxy.guardrails.guardrail_hooks.agent_365.agent_365 import registered_prometheus_logger from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import ( GuardrailEventHooks, @@ -29,6 +29,7 @@ from litellm.types.guardrails import ( SupportedGuardrailIntegrations, ) from litellm.types.proxy.guardrails.guardrail_hooks.agent_365 import ( + AGENT_365_DEFAULT_AUTHORITY_HOST, AGENT_365_PROD_API_BASE, AGENT_365_PROD_RESOURCE_APP_ID, Agent365GuardrailConfigModel, @@ -133,6 +134,7 @@ def _make_guardrail( unreachable_fallback: str = "fail_closed", agent_id: str | None = None, api_base: str = AGENT_365_PROD_API_BASE, + authority_host: str = AGENT_365_DEFAULT_AUTHORITY_HOST, prometheus: FakePrometheus | None = None, ) -> Agent365Guardrail: return Agent365Guardrail( @@ -142,6 +144,7 @@ def _make_guardrail( client_secret="secret-123", api_base=api_base, agent_id=agent_id, + authority_host=authority_host, unreachable_fallback=unreachable_fallback, async_handler=handler, prometheus_logger_lookup=lambda: prometheus, @@ -269,6 +272,41 @@ class TestInitializeGuardrail: assert LitellmParams(guardrail="generic_guardrail_api", mode="pre_call").unreachable_fallback is None assert Agent365GuardrailConfigModel.model_fields["unreachable_fallback"].default == "fail_open" + def test_authority_host_defaults_to_public_entra(self, monkeypatch): + monkeypatch.delenv("AGENT365_AUTHORITY_HOST", raising=False) + monkeypatch.delenv("AZURE_AUTHORITY_HOST", raising=False) + params: Final = LitellmParams( + guardrail="agent_365", mode="pre_mcp_call", tenant_id="t", client_id="c", client_secret="s" + ) + guardrail: Final = initialize_guardrail(params, {"guardrail_name": "a365-authority"}) + assert guardrail.authority_host == "https://login.microsoftonline.com" + + def test_authority_host_precedence_param_then_agent365_env_then_azure_env(self, monkeypatch): + monkeypatch.setenv("AZURE_AUTHORITY_HOST", "https://login.microsoftonline.us/") + azure_only: Final = initialize_guardrail( + LitellmParams(guardrail="agent_365", mode="pre_mcp_call", tenant_id="t", client_id="c", client_secret="s"), + {"guardrail_name": "a365-azure"}, + ) + assert azure_only.authority_host == "https://login.microsoftonline.us" + monkeypatch.setenv("AGENT365_AUTHORITY_HOST", "https://login.partner.microsoftonline.cn") + agent_env: Final = initialize_guardrail( + LitellmParams(guardrail="agent_365", mode="pre_mcp_call", tenant_id="t", client_id="c", client_secret="s"), + {"guardrail_name": "a365-agentenv"}, + ) + assert agent_env.authority_host == "https://login.partner.microsoftonline.cn" + explicit: Final = initialize_guardrail( + LitellmParams( + guardrail="agent_365", + mode="pre_mcp_call", + tenant_id="t", + client_id="c", + client_secret="s", + authority_host="http://127.0.0.1:9", + ), + {"guardrail_name": "a365-explicit"}, + ) + assert explicit.authority_host == "http://127.0.0.1:9" + def test_explicit_params_win(self, monkeypatch): monkeypatch.setenv("AGENT365_TENANT_ID", "env-tenant") params: Final = LitellmParams( @@ -335,6 +373,14 @@ class TestAllowFlow: assert token_call.data["client_secret"] == "secret-123" assert token_call.data["scope"] == f"{AGENT_365_PROD_RESOURCE_APP_ID}/ThreatProtection.Evaluate.All" + @pytest.mark.asyncio + async def test_obo_exchange_goes_to_the_configured_authority_host(self): + handler: Final = FakeHandler([_token_response(), _allow_response()]) + guardrail: Final = _make_guardrail(handler, authority_host="https://login.microsoftonline.us/") + await _run(guardrail, _mcp_data()) + assert handler.calls[0].url == "https://login.microsoftonline.us/tenant-abc/oauth2/v2.0/token" + assert handler.calls[1].url == EVALUATE_URL + @pytest.mark.asyncio async def test_evaluate_payload(self): handler: Final = FakeHandler([_token_response(), _allow_response()]) @@ -522,15 +568,18 @@ class TestDefenderNotEvaluated: class TestFailOpenDefault: @pytest.mark.asyncio - async def test_constructor_default_lets_timed_out_evaluation_through_unscanned(self): + async def test_constructor_default_lets_timed_out_evaluation_through_unscanned(self, caplog): handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")]) guardrail: Final = _default_fallback_guardrail(handler) assert guardrail.unreachable_fallback == "fail_open" data: Final = _mcp_data() - assert await _run(guardrail, data) is data + with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): + assert await _run(guardrail, data) is data info: Final = _guardrail_info(data) assert info["guardrail_status"] == "guardrail_failed_to_respond" assert info["guardrail_response"]["verdict"] == "Unscanned" + fail_open_logs: Final = [r for r in caplog.records if "unreachable_fallback='fail_open'" in r.getMessage()] + assert [r.levelno for r in fail_open_logs] == [logging.ERROR], caplog.text @pytest.mark.asyncio async def test_constructor_default_still_blocks_a_policy_block(self): @@ -576,12 +625,17 @@ class TestFailOpenDefault: await _run(guardrail, _mcp_data()) assert prometheus.fail_opens == [] - def test_registered_prometheus_logger_reads_litellm_callbacks(self, monkeypatch): - monkeypatch.setattr(litellm, "callbacks", []) - assert registered_prometheus_logger() is None - logger: Final = PrometheusLogger.__new__(PrometheusLogger) - monkeypatch.setattr(litellm, "callbacks", ["langfuse", logger]) - assert registered_prometheus_logger() is logger + def test_default_lookup_is_the_registered_prometheus_logger(self): + guardrail: Final = Agent365Guardrail( + guardrail_name="agent-365-guard", + tenant_id="tenant-abc", + client_id="client-xyz", + client_secret="secret-123", + async_handler=FakeHandler([]), + event_hook="pre_mcp_call", + default_on=True, + ) + assert guardrail._prometheus_logger_lookup is PrometheusLogger.get_instance class TestUnreachableFallback: diff --git a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py index fc2fb949143..d2b4bc2e52e 100644 --- a/tests/test_litellm/proxy/guardrails/test_init_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_init_guardrails.py @@ -1,9 +1,7 @@ import json -from unittest.mock import MagicMock, patch import pytest - from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2 from litellm.types.guardrails import SupportedGuardrailIntegrations @@ -167,6 +165,29 @@ def test_initialize_guardrail_sets_run_in_parallel(config_value, expected): assert custom_guardrail.run_in_parallel is expected +@pytest.mark.parametrize( + "guardrail, provider_params, expected", + [ + ("agent_365", {"tenant_id": "t", "client_id": "c", "client_secret": "s", "mode": "pre_mcp_call"}, "fail_open"), + ("typesafe", {"api_key": "k"}, "fail_open"), + ("generic_guardrail_api", {"api_base": "http://127.0.0.1:1/guard"}, "fail_closed"), + ("akto", {"akto_base_url": "http://127.0.0.1:1", "akto_api_key": "k", "akto_account_id": "1"}, "fail_closed"), + ("alice", {"api_key": "k", "api_base": "http://127.0.0.1:1"}, "fail_closed"), + ("deepkeep", {"api_base": "http://127.0.0.1:1", "api_key": "k", "deepkeep_firewall_id": "f"}, "fail_closed"), + ("repelloai", {"api_key": "k", "api_base": "http://127.0.0.1:1", "asset_id": "a"}, "fail_closed"), + ], +) +def test_unset_unreachable_fallback_applies_each_guardrails_own_default(guardrail, provider_params, expected): + litellm_params = {"guardrail": guardrail, "mode": "pre_call", **provider_params} + guardrail_handler = InMemoryGuardrailHandler() + result = guardrail_handler.initialize_guardrail( + guardrail={"guardrail_name": f"default-fallback-{guardrail}", "litellm_params": litellm_params}, + ) + + custom_guardrail = guardrail_handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]] + assert custom_guardrail.unreachable_fallback == expected, f"{guardrail} with unreachable_fallback unset" + + def test_initialize_presidio_forwards_analyze_chunk_size_bytes(): """Regression (LIT-4785): `presidio_analyze_chunk_size_bytes` set in config.yaml must reach the guardrail instance. The field lives on @@ -195,8 +216,7 @@ def test_initialize_presidio_forwards_analyze_chunk_size_bytes(): initialized = [ callback for callback in litellm.callbacks - if isinstance(callback, _OPTIONAL_PresidioPIIMasking) - and callback.guardrail_name == "test_presidio_chunk_size" + if isinstance(callback, _OPTIONAL_PresidioPIIMasking) and callback.guardrail_name == "test_presidio_chunk_size" ] assert initialized, "presidio guardrail was not registered as a callback" assert initialized[-1].presidio_analyze_chunk_size_bytes == 250_000 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/TeamGuardrailsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/TeamGuardrailsTab.tsx index 614fdaefade..5841d0162a5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/TeamGuardrailsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/TeamGuardrailsTab.tsx @@ -38,6 +38,7 @@ import { Textarea } from "@/components/ui/textarea"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { isValidUrl } from "@/lib/forms/urlValidation"; import { useZodForm } from "@/lib/forms/useZodForm"; +import { buildEquivalentConfigYaml, type TeamGuardrail, type TeamGuardrailStatus } from "./teamGuardrailConfigYaml"; import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog"; import { Button } from "@/components/ui/button"; @@ -97,32 +98,7 @@ const labelWithHint = (label: string, hint: string): React.ReactNode => ( ); -type GuardrailStatus = "active" | "pending" | "rejected"; - -type TeamGuardrail = { - id: string; - team: string; - name: string; - endpoint: string; - status: GuardrailStatus; - model: string; - forwardKey: boolean; - description: string; - method: "POST" | "GET"; - customHeaders: { - key: string; - value: string; - }[]; - extraHeaders: string[]; - submittedAt: string; - submittedBy: string; - mode?: string; - unreachable_fallback?: string; - additionalProviderParams?: Record; - guardrailType?: string; -}; - -function mapStatus(apiStatus: string): GuardrailStatus { +function mapStatus(apiStatus: string): TeamGuardrailStatus { if (apiStatus === "pending_review") return "pending"; if (apiStatus === "active" || apiStatus === "rejected") return apiStatus; return "active"; @@ -180,7 +156,7 @@ function submissionToTeamGuardrail(item: GuardrailSubmissionItem): TeamGuardrail }; } -const STATUS_CONFIG: Record = { +const STATUS_CONFIG: Record = { active: { label: "Active", bg: "bg-success/10", @@ -210,52 +186,6 @@ const TEAM_COLORS: Record = { Finance: "bg-success/15 text-success", }; -const FAIL_OPEN_BY_DEFAULT_GUARDRAILS: ReadonlySet = new Set(["agent_365", "typesafe"]); - -function defaultUnreachableFallback(guardrailType: string | undefined): "fail_open" | "fail_closed" { - return guardrailType !== undefined && FAIL_OPEN_BY_DEFAULT_GUARDRAILS.has(guardrailType) - ? "fail_open" - : "fail_closed"; -} - -function buildEquivalentConfigYaml(g: TeamGuardrail): string { - const lines: string[] = [ - "litellm_settings:", - " guardrails:", - ` - guardrail_name: "${g.name.replace(/\\/g, "\\\\").replace(/"/g, '\\"')}"`, - " litellm_params:", - ` guardrail: ${g.guardrailType ?? "generic_guardrail_api"}`, - ` mode: ${g.mode ?? "pre_call"} # or post_call, during_call`, - ` api_base: ${g.endpoint || "https://your-guardrail-api.com"}`, - " api_key: os.environ/YOUR_GUARDRAIL_API_KEY # optional", - ` unreachable_fallback: ${g.unreachable_fallback ?? defaultUnreachableFallback(g.guardrailType)} # fail_closed blocks, fail_open proceeds when the guardrail endpoint is unreachable. Shown value is this guardrail's default.`, - ` forward_api_key: ${g.forwardKey}`, - ]; - if (g.model && g.model !== "—") { - lines.push(` model: "${g.model}" # LLM model name sent to the guardrail for context`); - } - if (g.customHeaders.length > 0) { - lines.push(" headers: # static headers (sent with every request)"); - for (const h of g.customHeaders) { - lines.push(` ${h.key}: "${String(h.value).replace(/\\/g, "\\\\").replace(/"/g, '\\"')}"`); - } - } - if (g.extraHeaders.length > 0) { - lines.push(" extra_headers: # forward these client request headers to the guardrail"); - for (const name of g.extraHeaders) { - lines.push(` - ${name}`); - } - } - if (g.additionalProviderParams && Object.keys(g.additionalProviderParams).length > 0) { - lines.push(" additional_provider_specific_params:"); - for (const [k, v] of Object.entries(g.additionalProviderParams)) { - const val = typeof v === "string" ? `"${v}"` : String(v); - lines.push(` ${k}: ${val}`); - } - } - return lines.join("\n"); -} - function StatCard({ label, value, color }: { label: string; value: number; color: string }) { return (
@@ -833,7 +763,7 @@ export function TeamGuardrailsTab({ accessToken }: TeamGuardrailsTabProps) { }); const [search, setSearch] = useState(""); const [searchDebounced] = useDebouncedValue(search, { wait: DEBOUNCE_WAIT_MS }); - const [statusFilter, setStatusFilter] = useState<"all" | GuardrailStatus>("all"); + const [statusFilter, setStatusFilter] = useState<"all" | TeamGuardrailStatus>("all"); const [selectedId, setSelectedId] = useState(null); const [expandedHeaders, setExpandedHeaders] = useState>(new Set()); const [confirmAction, setConfirmAction] = useState<{ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/teamGuardrailConfigYaml.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/teamGuardrailConfigYaml.test.ts new file mode 100644 index 00000000000..e521dd49e3f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/teamGuardrailConfigYaml.test.ts @@ -0,0 +1,85 @@ +import { describe, expect, it } from "vitest"; + +import { buildEquivalentConfigYaml, defaultUnreachableFallback, type TeamGuardrail } from "./teamGuardrailConfigYaml"; + +function guardrail(overrides: Partial): TeamGuardrail { + return { + id: "g-1", + team: "Security", + name: "agent365-mcp", + endpoint: "", + status: "active", + model: "—", + forwardKey: false, + description: "", + method: "POST", + customHeaders: [], + extraHeaders: [], + submittedAt: "2026-01-01", + submittedBy: "ops@example.com", + ...overrides, + }; +} + +function fallbackLine(yaml: string): string { + const line = yaml.split("\n").find((l) => l.includes("unreachable_fallback:")); + if (line === undefined) throw new Error(`no unreachable_fallback line in:\n${yaml}`); + return line.trim().split(" #")[0]; +} + +describe("defaultUnreachableFallback", () => { + it("fails open for agent_365 and typesafe, closed for everything else including unknown", () => { + expect(defaultUnreachableFallback("agent_365")).toBe("fail_open"); + expect(defaultUnreachableFallback("typesafe")).toBe("fail_open"); + expect(defaultUnreachableFallback("generic_guardrail_api")).toBe("fail_closed"); + expect(defaultUnreachableFallback("akto")).toBe("fail_closed"); + expect(defaultUnreachableFallback(undefined)).toBe("fail_closed"); + }); +}); + +describe("buildEquivalentConfigYaml unreachable_fallback line", () => { + it("shows fail_open for an agent_365 guardrail with no explicit fallback", () => { + const yaml = buildEquivalentConfigYaml(guardrail({ guardrailType: "agent_365" })); + expect(fallbackLine(yaml)).toBe("unreachable_fallback: fail_open"); + expect(yaml).toContain(" guardrail: agent_365"); + }); + + it("shows fail_open for a typesafe guardrail with no explicit fallback", () => { + expect(fallbackLine(buildEquivalentConfigYaml(guardrail({ guardrailType: "typesafe" })))).toBe( + "unreachable_fallback: fail_open", + ); + }); + + it("shows fail_closed for a generic guardrail and when the type is unknown", () => { + expect(fallbackLine(buildEquivalentConfigYaml(guardrail({ guardrailType: "generic_guardrail_api" })))).toBe( + "unreachable_fallback: fail_closed", + ); + const untyped = buildEquivalentConfigYaml(guardrail({})); + expect(fallbackLine(untyped)).toBe("unreachable_fallback: fail_closed"); + expect(untyped).toContain(" guardrail: generic_guardrail_api"); + }); + + it("keeps an explicit override over the per-guardrail default", () => { + expect( + fallbackLine( + buildEquivalentConfigYaml(guardrail({ guardrailType: "agent_365", unreachable_fallback: "fail_closed" })), + ), + ).toBe("unreachable_fallback: fail_closed"); + expect( + fallbackLine(buildEquivalentConfigYaml(guardrail({ guardrailType: "akto", unreachable_fallback: "fail_open" }))), + ).toBe("unreachable_fallback: fail_open"); + }); + + it("explains the shown value is the guardrail's default only when nothing was set explicitly", () => { + const rawLine = (g: TeamGuardrail) => + buildEquivalentConfigYaml(g) + .split("\n") + .find((l) => l.includes("unreachable_fallback:")); + expect(rawLine(guardrail({ guardrailType: "agent_365" }))).toBe( + " unreachable_fallback: fail_open # fail_closed blocks, fail_open proceeds when the guardrail endpoint is unreachable. Shown value is this guardrail's default.", + ); + expect(rawLine(guardrail({ guardrailType: "agent_365", unreachable_fallback: "fail_closed" }))).toBe( + " unreachable_fallback: fail_closed # fail_closed blocks, fail_open proceeds when the guardrail endpoint is unreachable", + ); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/teamGuardrailConfigYaml.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/teamGuardrailConfigYaml.ts new file mode 100644 index 00000000000..824bf37b8f2 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/teamGuardrailConfigYaml.ts @@ -0,0 +1,78 @@ +export type TeamGuardrailStatus = "active" | "pending" | "rejected"; + +export type TeamGuardrail = { + id: string; + team: string; + name: string; + endpoint: string; + status: TeamGuardrailStatus; + model: string; + forwardKey: boolean; + description: string; + method: "POST" | "GET"; + customHeaders: { + key: string; + value: string; + }[]; + extraHeaders: string[]; + submittedAt: string; + submittedBy: string; + mode?: string; + unreachable_fallback?: string; + additionalProviderParams?: Record; + guardrailType?: string; +}; + +const FAIL_OPEN_BY_DEFAULT_GUARDRAILS: ReadonlySet = new Set(["agent_365", "typesafe"]); + +export function defaultUnreachableFallback(guardrailType: string | undefined): "fail_open" | "fail_closed" { + return guardrailType !== undefined && FAIL_OPEN_BY_DEFAULT_GUARDRAILS.has(guardrailType) + ? "fail_open" + : "fail_closed"; +} + +function unreachableFallbackLine(g: TeamGuardrail): string { + const hint = "fail_closed blocks, fail_open proceeds when the guardrail endpoint is unreachable"; + if (g.unreachable_fallback !== undefined) { + return ` unreachable_fallback: ${g.unreachable_fallback} # ${hint}`; + } + return ` unreachable_fallback: ${defaultUnreachableFallback(g.guardrailType)} # ${hint}. Shown value is this guardrail's default.`; +} + +export function buildEquivalentConfigYaml(g: TeamGuardrail): string { + const lines: string[] = [ + "litellm_settings:", + " guardrails:", + ` - guardrail_name: "${g.name.replace(/\\/g, "\\\\").replace(/"/g, '\\"')}"`, + " litellm_params:", + ` guardrail: ${g.guardrailType ?? "generic_guardrail_api"}`, + ` mode: ${g.mode ?? "pre_call"} # or post_call, during_call`, + ` api_base: ${g.endpoint || "https://your-guardrail-api.com"}`, + " api_key: os.environ/YOUR_GUARDRAIL_API_KEY # optional", + unreachableFallbackLine(g), + ` forward_api_key: ${g.forwardKey}`, + ]; + if (g.model && g.model !== "—") { + lines.push(` model: "${g.model}" # LLM model name sent to the guardrail for context`); + } + if (g.customHeaders.length > 0) { + lines.push(" headers: # static headers (sent with every request)"); + for (const h of g.customHeaders) { + lines.push(` ${h.key}: "${String(h.value).replace(/\\/g, "\\\\").replace(/"/g, '\\"')}"`); + } + } + if (g.extraHeaders.length > 0) { + lines.push(" extra_headers: # forward these client request headers to the guardrail"); + for (const name of g.extraHeaders) { + lines.push(` - ${name}`); + } + } + if (g.additionalProviderParams && Object.keys(g.additionalProviderParams).length > 0) { + lines.push(" additional_provider_specific_params:"); + for (const [k, v] of Object.entries(g.additionalProviderParams)) { + const val = typeof v === "string" ? `"${v}"` : String(v); + lines.push(` ${k}: ${val}`); + } + } + return lines.join("\n"); +} diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index cff334981ca..758eccfd075 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -34084,6 +34084,11 @@ export interface components { * @description Authorization bearer token for IBM Guardrails API. Reads from IBM_GUARDRAILS_AUTH_TOKEN env var if None. */ auth_token?: string | null; + /** + * Authority Host + * @description Microsoft Entra authority host that issues the On-Behalf-Of token, for sovereign clouds. Defaults to https://login.microsoftonline.com. Falls back to the AGENT365_AUTHORITY_HOST, then AZURE_AUTHORITY_HOST environment variables. + */ + authority_host?: string | null; /** * Aws Access Key Id * @description AWS access key ID for authentication