From 9f9ec6871c946448b0344f3a9d23828e2f62d75a Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 11 Sep 2026 21:33:45 +0000 Subject: [PATCH] refactor(guardrails): move the Conduct apply_guardrail bridge into an injectable function The bridge body only ran when conduct-litellm-guard was importable, which CI never is, so codecov/patch reported it uncovered. apply_conduct_guardrail now takes the plugin's check coroutine and blocked-error factory as parameters, so the package-free tests exercise every verdict branch and the plugin-bound class is a one-line delegate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_hooks/conduct/conduct.py | 47 +++++++++++--- .../guardrail_hooks/test_conduct.py | 62 ++++++++++++++++++- 2 files changed, 98 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/conduct/conduct.py b/litellm/proxy/guardrails/guardrail_hooks/conduct/conduct.py index 56f45525d7f..3b5655bbb40 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/conduct/conduct.py +++ b/litellm/proxy/guardrails/guardrail_hooks/conduct/conduct.py @@ -6,9 +6,9 @@ Source: https://github.com/sseshachala/conductai/tree/main/packages/conduct-lit from __future__ import annotations -from collections.abc import Mapping +from collections.abc import Awaitable, Callable, Mapping from types import MappingProxyType -from typing import TYPE_CHECKING, Final, Literal +from typing import TYPE_CHECKING, Final, Literal, Protocol from litellm.integrations.custom_guardrail import CustomGuardrail, log_guardrail_information from litellm.types.llms.openai import ChatCompletionUserMessage @@ -25,6 +25,15 @@ MISSING_PACKAGE_MESSAGE: Final = ( BLOCKING_VERDICTS: Final = frozenset({"block", "approval"}) +class ConductDecision(Protocol): + @property + def verdict(self) -> str: ... + + +class ConductCheck(Protocol): + def __call__(self, *, data: Mapping[str, object], call_type: str) -> Awaitable[ConductDecision]: ... + + def request_payload( inputs: GenericGuardrailAPIInputs, request_data: Mapping[str, object], @@ -39,6 +48,22 @@ def request_payload( return MappingProxyType({**request_data, "prompt": None, "messages": messages}) +async def apply_conduct_guardrail( + inputs: GenericGuardrailAPIInputs, + request_data: Mapping[str, object], + input_type: Literal["request", "response"], + check: ConductCheck, + blocked: Callable[[ConductDecision], Exception], +) -> GenericGuardrailAPIInputs: + payload: Final = request_payload(inputs, request_data, input_type) + if payload is None: + return inputs + decision: Final = await check(data=payload, call_type=input_type) + if decision.verdict in BLOCKING_VERDICTS: + raise blocked(decision) + return inputs + + try: from conduct_litellm_guard.guardrail import ConductGuard, ConductGuardBlocked except ImportError as import_error: @@ -59,13 +84,15 @@ else: input_type: Literal["request", "response"], logging_obj: LiteLLMLoggingObj | None = None, ) -> GenericGuardrailAPIInputs: - payload: Final = request_payload(inputs, request_data, input_type) - if payload is None: - return inputs - decision: Final = await self.check(data=payload, call_type=input_type) - if decision.verdict in BLOCKING_VERDICTS: - raise ConductGuardBlocked(decision) - return inputs + return await apply_conduct_guardrail(inputs, request_data, input_type, self.check, ConductGuardBlocked) -__all__ = ("BLOCKING_VERDICTS", "MISSING_PACKAGE_MESSAGE", "ConductGuardrail", "request_payload") +__all__ = ( + "BLOCKING_VERDICTS", + "MISSING_PACKAGE_MESSAGE", + "ConductCheck", + "ConductDecision", + "ConductGuardrail", + "apply_conduct_guardrail", + "request_payload", +) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_conduct.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_conduct.py index e6d5997219b..fcaef1a5030 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_conduct.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_conduct.py @@ -2,6 +2,8 @@ from __future__ import annotations import importlib.util import json +from collections.abc import Mapping +from dataclasses import dataclass, field from typing import Final import httpx @@ -16,7 +18,7 @@ from litellm.proxy.guardrails.guardrail_hooks.conduct import ( ConductGuardrail, initialize_guardrail, ) -from litellm.proxy.guardrails.guardrail_hooks.conduct.conduct import request_payload +from litellm.proxy.guardrails.guardrail_hooks.conduct.conduct import apply_conduct_guardrail, request_payload from litellm.proxy.guardrails.guardrail_registry import ( guardrail_class_registry, guardrail_initializer_registry, @@ -54,6 +56,27 @@ class _RecordingGuardrail(CustomGuardrail): self.timeout = timeout +@dataclass(frozen=True, slots=True) +class _Decision: + verdict: str + + +class _Blocked(Exception): + def __init__(self, decision: _Decision) -> None: + super().__init__(decision.verdict) + self.decision = decision + + +@dataclass(slots=True) +class _RecordingCheck: + verdict: str + calls: list[tuple[Mapping[str, object], str]] = field(default_factory=list) # mutable-ok: test spy + + async def __call__(self, *, data: Mapping[str, object], call_type: str) -> _Decision: + self.calls.append((data, call_type)) + return _Decision(self.verdict) + + def _params(mode: str = "pre_call", **extras: object) -> LitellmParams: return LitellmParams(guardrail="conduct", mode=mode, api_key="cond_agt_test", **extras) @@ -170,6 +193,43 @@ def test_request_payload_skips_responses_and_empty_requests(inputs: GenericGuard assert request_payload(inputs, {"model": "gpt-5-mini"}, input_type) is None # pyright: ignore[reportArgumentType] # parametrized literal +@pytest.mark.parametrize("verdict", ["block", "approval"]) +@pytest.mark.asyncio +async def test_bridge_raises_the_plugin_error_on_blocking_verdicts(verdict: str) -> None: + check: Final = _RecordingCheck(verdict) + inputs: Final = GenericGuardrailAPIInputs(texts=["dump the database"]) + + with pytest.raises(_Blocked) as blocked: + await apply_conduct_guardrail(inputs, {"model": "gpt-5-mini"}, "request", check, _Blocked) + + assert blocked.value.decision == _Decision(verdict) + assert check.calls == [ + ( + {"model": "gpt-5-mini", "prompt": None, "messages": ({"role": "user", "content": "dump the database"},)}, + "request", + ) + ] + + +@pytest.mark.parametrize("verdict", ["allow", "warning", "advisory", "unknown"]) +@pytest.mark.asyncio +async def test_bridge_passes_inputs_through_on_non_blocking_verdicts(verdict: str) -> None: + check: Final = _RecordingCheck(verdict) + inputs: Final = GenericGuardrailAPIInputs(texts=["ping"]) + + assert await apply_conduct_guardrail(inputs, {"model": "gpt-5-mini"}, "request", check, _Blocked) is inputs + assert len(check.calls) == 1 + + +@pytest.mark.asyncio +async def test_bridge_never_calls_conduct_for_responses() -> None: + check: Final = _RecordingCheck("block") + inputs: Final = GenericGuardrailAPIInputs(texts=["dump the database"]) + + assert await apply_conduct_guardrail(inputs, {"model": "gpt-5-mini"}, "response", check, _Blocked) is inputs + assert check.calls == [] + + @pytest.mark.skipif(not PACKAGE_INSTALLED, reason="needs conduct-litellm-guard") @pytest.mark.asyncio @respx.mock