mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
parent
0f729bd767
commit
9f9ec6871c
2 changed files with 98 additions and 11 deletions
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue