mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
parent
0c4d59c716
commit
ff59e6d25b
6 changed files with 22 additions and 149 deletions
|
|
@ -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
|
||||
########################################
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue