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>
This commit is contained in:
yucheng 2026-09-11 21:33:45 +00:00
parent 0f729bd767
commit 9f9ec6871c
2 changed files with 98 additions and 11 deletions

View file

@ -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",
)

View file

@ -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