refactor(guardrails): drop the prometheus fail-open counter from the agent 365 guardrail

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-26 00:30:30 +00:00
parent 0c4d59c716
commit ff59e6d25b
6 changed files with 22 additions and 149 deletions

View file

@ -3252,16 +3252,6 @@ class PrometheusLogger(CustomLogger):
except Exception as e:
verbose_logger.debug("Error recording guardrail metrics: %s", e)
def record_guardrail_fail_open(self, guardrail_name: str, hook_type: str) -> None:
try:
self.litellm_guardrail_errors_total.labels(
guardrail_name=guardrail_name,
error_type="fail_open",
hook_type=hook_type,
).inc()
except Exception as e:
verbose_logger.debug("Error recording guardrail fail-open metric: %s", e)
########################################
# Managed Batch Metric Recording Methods
########################################

View file

@ -14,7 +14,7 @@ import threading
import time
import uuid
from collections import OrderedDict
from collections.abc import Callable, Mapping
from collections.abc import Mapping
from typing import TYPE_CHECKING, ClassVar, Final, Literal, NoReturn
import httpx
@ -28,7 +28,6 @@ from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.integrations.prometheus import PrometheusLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -162,7 +161,6 @@ class Agent365Guardrail(CustomGuardrail):
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] = PrometheusLogger.get_instance,
**kwargs, # noqa: ANN003 # kwargs-ok: forwarded verbatim to CustomGuardrail (event_hook, default_on)
) -> None:
super().__init__(
@ -187,7 +185,6 @@ class Agent365Guardrail(CustomGuardrail):
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self._prometheus_logger_lookup = prometheus_logger_lookup
self._obo_token_cache: OrderedDict[str, tuple[str, float]] = OrderedDict() # mutable-ok: lock-guarded LRU
self._obo_cache_lock = threading.Lock()
verbose_proxy_logger.info("Initialized Microsoft Agent 365 guardrail: %s", guardrail_name)
@ -588,7 +585,6 @@ class Agent365Guardrail(CustomGuardrail):
reason,
tool_name,
)
self._count_fail_open()
self._record_verdict(
data=data,
verdict="Unscanned",
@ -616,13 +612,6 @@ class Agent365Guardrail(CustomGuardrail):
}
raise HTTPException(status_code=503, detail=unavailable_detail)
def _count_fail_open(self) -> None:
prometheus: Final = self._prometheus_logger_lookup()
if prometheus is not None:
prometheus.record_guardrail_fail_open(
guardrail_name=self.guardrail_name or type(self).__name__, hook_type="pre_call"
)
def _record_verdict(
self,
data: dict[str, object], # mutable-ok: standard guardrail logging appends into the request metadata in place

View file

@ -76,7 +76,7 @@ class Agent365GuardrailConfigModel(GuardrailConfigModel):
description=(
"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. "
"it as Unscanned in the logs and OpenTelemetry. "
"'fail_closed' blocks it with HTTP 503. Policy blocks, 4xx rejections, throttling and a rejected "
"caller token always block."
),

View file

@ -31,11 +31,9 @@ from integration._support.mcp import (
)
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, 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_TIMEOUT_SECONDS: Final = 1.0
SLOW_REPLY_SECONDS: Final = 3.0
@ -113,7 +111,6 @@ def _config(
sibling_url: 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,
@ -147,6 +144,7 @@ def _config(
),
]
path: Final = tmp_path / "agent_365.yaml"
tmp_path.mkdir(parents=True, exist_ok=True)
path.write_text(yaml.safe_dump(config))
return path
@ -221,8 +219,6 @@ def _rig(
) -> Iterator[Rig]:
"""``authority_env`` names the environment variable that carries the Entra double instead of the config."""
alias: Final = "a365" + uuid.uuid4().hex[:8]
prom_dir: Final = tmp_path / "prom"
prom_dir.mkdir(parents=True)
with (
wire_server(_entra) as entra,
wire_server(_agent_365) as agent_365,
@ -232,7 +228,6 @@ def _rig(
gateway,
tmp_path,
{
"PROMETHEUS_MULTIPROC_DIR": str(prom_dir),
"PROXY_CONFIG_RELOAD_INTERVAL_SECONDS": "2",
**(environment or {}),
**({authority_env: entra.url} if authority_env else {}),
@ -264,26 +259,13 @@ def _guardrail_statuses(key: str) -> dict[str, str]:
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 _chat(rig: Rig, model: str, marker: str) -> httpx.Response:
return rig.candidate.request(
"POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": marker}]}, key=rig.key
)
def test_default_lets_the_call_through_unscanned_and_counts_it_when_agent_365_cannot_evaluate(
def test_default_lets_the_call_through_unscanned_when_agent_365_cannot_evaluate(
gateway: Gateway, tmp_path: Path
) -> None:
with _rig(gateway, tmp_path, fallback=None) as rig:
@ -300,10 +282,6 @@ def test_default_lets_the_call_through_unscanned_and_counts_it_when_agent_365_ca
"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(
@ -316,7 +294,6 @@ def test_explicit_fail_closed_blocks_with_503_and_never_reaches_upstream_when_ag
assert rig.upstream_tool_names() == ("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}
def test_explicit_fail_open_matches_the_default_and_still_blocks_policy_denials(
@ -329,7 +306,6 @@ def test_explicit_fail_open_matches_the_default_and_still_blocks_policy_denials(
assert rig.upstream_tool_names() == ("outage",)
statuses: Final = eventually(lambda: _guardrail_statuses(rig.key), lambda seen: len(seen) >= 2, seconds=70)
assert statuses == {"outage": "guardrail_failed_to_respond", "denied": "guardrail_intervened"}, statuses
assert _guardrail_error_counts(rig.candidate, rig.guardrail_name) == {"fail_open": 1.0, "HTTPException": 1.0}
def test_default_fails_open_on_malformed_or_stalled_agent_365_replies(gateway: Gateway, tmp_path: Path) -> None:
@ -340,7 +316,6 @@ def test_default_fails_open_on_malformed_or_stalled_agent_365_replies(gateway: G
assert rig.upstream_tool_names() == ("nonjson", "nobool", "slow")
statuses: Final = eventually(lambda: _guardrail_statuses(rig.key), lambda seen: len(seen) >= 3, seconds=70)
assert statuses == dict.fromkeys(("nonjson", "nobool", "slow"), "guardrail_failed_to_respond"), statuses
assert _guardrail_error_counts(rig.candidate, rig.guardrail_name) == {"fail_open": 3.0}
def test_default_fails_open_when_entra_is_down_stalled_malformed_or_refuses_the_gateway_credentials(
@ -362,7 +337,6 @@ def test_default_fails_open_when_entra_is_down_stalled_malformed_or_refuses_the_
seconds=70,
)
assert [row["gi"][0]["guardrail_status"] for row in rows] == ["guardrail_failed_to_respond"] * len(cases)
assert _guardrail_error_counts(rig.candidate, rig.guardrail_name) == {"fail_open": float(len(cases))}
def test_throttling_and_ordinary_4xx_from_agent_365_keep_blocking_under_the_default(
@ -376,7 +350,6 @@ def test_throttling_and_ordinary_4xx_from_agent_365_keep_blocking_under_the_defa
assert rig.upstream_tool_names() == ()
statuses: Final = eventually(lambda: _guardrail_statuses(rig.key), lambda seen: len(seen) >= 2, seconds=70)
assert statuses == {"throttled": "guardrail_failed_to_respond", "rejected": "guardrail_intervened"}, statuses
assert _guardrail_error_counts(rig.candidate, rig.guardrail_name) == {"HTTPException": 2.0}
def test_caller_authentication_failures_keep_blocking_under_the_default(gateway: Gateway, tmp_path: Path) -> None:
@ -397,7 +370,6 @@ def test_caller_authentication_failures_keep_blocking_under_the_default(gateway:
seconds=70,
)
assert [row["gi"][0]["guardrail_status"] for row in rows] == ["guardrail_intervened"] * 3
assert _guardrail_error_counts(rig.candidate, rig.guardrail_name) == {"HTTPException": 3.0}
def test_every_mcp_entry_point_fails_open_on_outage_and_blocks_denials(gateway: Gateway, tmp_path: Path) -> None:
@ -409,11 +381,6 @@ def test_every_mcp_entry_point_fails_open_on_outage_and_blocks_denials(gateway:
denied: Final = caller.call(f"{rig.alias}-denied", {"entry": entry}, server_id=rig.server_id)
assert denied.error is not None and "Blocked by Microsoft Defender" in denied.raw, f"{entry}: {denied.raw}"
assert rig.upstream_tool_names() == ("outage",) * len(ENTRY_POINTS)
counts: Final = eventually(
lambda: _guardrail_error_counts(rig.candidate, rig.guardrail_name),
lambda seen: sum(seen.values()) >= 2 * len(ENTRY_POINTS),
)
assert counts == {"fail_open": float(len(ENTRY_POINTS)), "HTTPException": float(len(ENTRY_POINTS))}, counts
def test_chat_completions_never_touch_agent_365_while_a_sibling_guardrail_keeps_its_fail_closed_default(
@ -425,7 +392,7 @@ def test_chat_completions_never_touch_agent_365_while_a_sibling_guardrail_keeps_
assert chat.status_code == 500 and "Generic Guardrail API failed" in chat.text, chat.text
assert rig.entra.drain() == () and rig.agent_365.drain() == ()
assert rig.caller.call(f"{rig.alias}-outage", {"a": 1}).text == '{"a": 1}'
assert _guardrail_error_counts(rig.candidate, rig.guardrail_name) == {"fail_open": 1.0}
assert rig.upstream_tool_names() == ("outage",)
def test_chat_completions_are_unaffected_by_the_mcp_guardrail(gateway: Gateway, tmp_path: Path) -> None:
@ -441,7 +408,6 @@ def test_chat_completions_are_unaffected_by_the_mcp_guardrail(gateway: Gateway,
seconds=70,
)
assert rows[0]["gi"] is None, rows
assert _guardrail_error_counts(rig.candidate, rig.guardrail_name) == {}
def test_agent365_authority_host_env_wins_over_azure_authority_host_when_config_has_none(
@ -480,10 +446,7 @@ def test_thirty_call_burst_against_a_flapping_agent_365_reaches_upstream_exactly
)
assert seen == sorted(marker for _, marker in markers)
assert rig.caller.call(f"{rig.alias}-denied", {"a": 1}).error is not None
counts: Final = eventually(
lambda: _guardrail_error_counts(rig.candidate, rig.guardrail_name), lambda seen: len(seen) >= 2
)
assert counts == {"fail_open": 15.0, "HTTPException": 1.0}, counts
assert rig.upstream_tool_names() == ()
def test_default_survives_a_worker_kill_and_keeps_blocking_denials_on_two_workers(

View file

@ -12,7 +12,6 @@ import litellm
from litellm.caching.caching import DualCache
from litellm.exceptions import Timeout as LitellmTimeout
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.prometheus import PrometheusLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.proxy._types import UserAPIKeyAuth
@ -120,14 +119,6 @@ class FakeHandler:
return item
class FakePrometheus:
def __init__(self) -> None:
self.fail_opens: list[tuple[str, str]] = []
def record_guardrail_fail_open(self, guardrail_name: str, hook_type: str) -> None:
self.fail_opens.append((guardrail_name, hook_type))
def _make_guardrail(
handler: FakeHandler,
*,
@ -135,7 +126,6 @@ def _make_guardrail(
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(
guardrail_name="agent-365-guard",
@ -147,20 +137,18 @@ def _make_guardrail(
authority_host=authority_host,
unreachable_fallback=unreachable_fallback,
async_handler=handler,
prometheus_logger_lookup=lambda: prometheus,
event_hook="pre_mcp_call",
default_on=True,
)
def _default_fallback_guardrail(handler: FakeHandler, prometheus: FakePrometheus | None = None) -> Agent365Guardrail:
def _default_fallback_guardrail(handler: FakeHandler) -> Agent365Guardrail:
return Agent365Guardrail(
guardrail_name="agent-365-guard",
tenant_id="tenant-abc",
client_id="client-xyz",
client_secret="secret-123",
async_handler=handler,
prometheus_logger_lookup=lambda: prometheus,
event_hook="pre_mcp_call",
default_on=True,
)
@ -614,43 +602,11 @@ class TestFailOpenDefault:
[_response(503, text="entra down")],
],
)
async def test_each_fail_open_counts_once_in_prometheus(self, responses):
prometheus: Final = FakePrometheus()
guardrail: Final = _default_fallback_guardrail(FakeHandler(responses), prometheus)
async def test_each_availability_failure_lets_the_call_through_as_failed_to_respond(self, responses):
guardrail: Final = _default_fallback_guardrail(FakeHandler(responses))
data: Final = _mcp_data()
assert await _run(guardrail, data) is data
assert prometheus.fail_opens == [("agent-365-guard", "pre_call")]
@pytest.mark.asyncio
async def test_allowed_and_blocked_calls_do_not_count_as_fail_open(self):
prometheus: Final = FakePrometheus()
handler: Final = FakeHandler([_token_response(), _allow_response(), _block_response()])
guardrail: Final = _default_fallback_guardrail(handler, prometheus)
await _run(guardrail, _mcp_data())
with pytest.raises(HTTPException):
await _run(guardrail, _mcp_data())
assert prometheus.fail_opens == []
@pytest.mark.asyncio
async def test_fail_closed_does_not_count_as_fail_open(self):
prometheus: Final = FakePrometheus()
handler: Final = FakeHandler([_token_response(), httpx.ReadTimeout("timed out")])
guardrail: Final = _make_guardrail(handler, unreachable_fallback="fail_closed", prometheus=prometheus)
with pytest.raises(HTTPException):
await _run(guardrail, _mcp_data())
assert prometheus.fail_opens == []
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
assert _guardrail_info(data)["guardrail_status"] == "guardrail_failed_to_respond"
class TestUnreachableFallback:

View file

@ -158,7 +158,9 @@ class TestPrometheusQueueTimeMetric:
if len(call[0]) > 0 and call[0][0] == 0.5: # Our test queue time value
queue_time_called = True
break
assert not queue_time_called, "Queue time metric should not be recorded when queue_time_seconds is missing"
assert (
not queue_time_called
), "Queue time metric should not be recorded when queue_time_seconds is missing"
def test_queue_time_metric_not_recorded_when_negative(self):
"""Test that queue time metric is not recorded when queue_time_seconds is negative"""
@ -222,7 +224,9 @@ class TestPrometheusQueueTimeMetric:
if len(call[0]) > 0 and call[0][0] == -0.1:
negative_value_called = True
break
assert not negative_value_called, "Queue time metric should not be recorded for negative values"
assert (
not negative_value_called
), "Queue time metric should not be recorded for negative values"
class TestPrometheusTotalLatencyMetric:
@ -395,7 +399,9 @@ class TestPrometheusGuardrailMetrics:
error_type="none",
hook_type=hook_type,
)
mock_latency_metric.labels.return_value.observe.assert_called_once_with(latency_seconds)
mock_latency_metric.labels.return_value.observe.assert_called_once_with(
latency_seconds
)
# Assert - requests metric should be incremented
mock_requests_metric.labels.assert_called_once_with(
@ -444,7 +450,9 @@ class TestPrometheusGuardrailMetrics:
error_type=error_type,
hook_type=hook_type,
)
mock_latency_metric.labels.return_value.observe.assert_called_once_with(latency_seconds)
mock_latency_metric.labels.return_value.observe.assert_called_once_with(
latency_seconds
)
# Assert - requests metric should be incremented
mock_requests_metric.labels.assert_called_once_with(
@ -546,36 +554,3 @@ class TestPrometheusGuardrailMetrics:
mock_latency_metric.labels.assert_called_once()
call_kwargs = mock_latency_metric.labels.call_args[1]
assert call_kwargs["guardrail_name"] == guardrail_name
def test_record_guardrail_fail_open_increments_errors_total(self):
prometheus_logger = PrometheusLogger()
prometheus_logger.record_guardrail_fail_open(guardrail_name="a365", hook_type="pre_call")
prometheus_logger.record_guardrail_fail_open(guardrail_name="a365", hook_type="pre_call")
assert (
REGISTRY.get_sample_value(
"litellm_guardrail_errors_total",
{"guardrail_name": "a365", "error_type": "fail_open", "hook_type": "pre_call"},
)
== 2.0
)
assert (
REGISTRY.get_sample_value(
"litellm_guardrail_requests_total",
{"guardrail_name": "a365", "status": "success", "hook_type": "pre_call"},
)
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"
)