mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(guardrails): treat an unknown straiker api_version as unset instead of skipping the guardrail (#43956)
* fix(guardrails): treat an unknown straiker api_version as unset instead of skipping the guardrail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): patch the shared proxy logger directly in the straiker api_version test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): keep the straiker stray-version block marker separate from the shared block marker Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): drop redundant comments on the straiker api_version integration tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng <yucheng@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2c3866ebb4
commit
a3a7650569
3 changed files with 105 additions and 3 deletions
|
|
@ -1,8 +1,9 @@
|
|||
from typing import TYPE_CHECKING, Final, Literal
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, field_validator
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
from .straiker import StraikerGuardrail
|
||||
|
|
@ -17,6 +18,18 @@ class _V3Routing(BaseModel):
|
|||
client: str | None = None
|
||||
format_hint: Literal["anthropic.messages", "openai.chat"] | None = None
|
||||
|
||||
@field_validator("api_version", mode="before")
|
||||
@classmethod
|
||||
def _unknown_api_version_is_unset(cls, value: object) -> object:
|
||||
if value is None or value in ("v1", "v3"):
|
||||
return value
|
||||
verbose_proxy_logger.warning(
|
||||
"Straiker guardrail: ignoring api_version %r, expected 'v1', 'v3' or unset; "
|
||||
"the route follows the api_key prefix",
|
||||
value,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
_OPTIONAL_INIT_FIELDS: Final = (
|
||||
"timeout",
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ V1_KEY: Final = "synthetic-v1-collection-key"
|
|||
V3_PATH: Final = "/api/v3/detect"
|
||||
V1_PATH: Final = "/api/v1/detect/webhook"
|
||||
BLOCK_MARK: Final = "SYNTHETIC-INJECTION"
|
||||
STRAY_V3_BLOCK_MARK: Final = "SYNTHETIC-STRAY-VERSION-BLOCK"
|
||||
KILL_MARK: Final = "SYNTHETIC-KILLSWITCH"
|
||||
DENY_MARK: Final = "SYNTHETIC-DENY"
|
||||
SINK_500_MARK: Final = "SYNTHETIC-SINK-500"
|
||||
|
|
@ -144,7 +145,11 @@ def _verdict(seen: Seen, text: str) -> tuple[int, bytes]:
|
|||
return 200, json.dumps({"action": "NONE"}).encode()
|
||||
assert seen.target == V3_PATH, seen.target
|
||||
turn: Final = "turn-" + hashlib.sha256(text.encode()).hexdigest()[:12]
|
||||
if BLOCK_MARK in text or (LOG_BLOCK_MARK in text and agent == LOG_AGENT):
|
||||
if (
|
||||
BLOCK_MARK in text
|
||||
or (STRAY_V3_BLOCK_MARK in text and agent is None)
|
||||
or (LOG_BLOCK_MARK in text and agent == LOG_AGENT)
|
||||
):
|
||||
return 200, json.dumps(
|
||||
{
|
||||
"hookSpecificOutput": {"permissionDecision": "block"},
|
||||
|
|
@ -363,7 +368,9 @@ def _rig_config(sink_url: str, root: Path) -> Path:
|
|||
format_hint="anthropic.messages",
|
||||
),
|
||||
_guardrail("straiker-v3-as-v1", V3_KEY, sink_url, "pre_call", False, api_version="v1"),
|
||||
_guardrail("straiker-v3-stray-version", V3_KEY, sink_url, "pre_call", False, api_version="2024-09-01"),
|
||||
_guardrail("straiker-v1", V1_KEY, sink_url, "pre_call", False),
|
||||
_guardrail("straiker-v1-empty-version", V1_KEY, sink_url, "pre_call", False, api_version=""),
|
||||
_guardrail("straiker-v1-post", V1_KEY, sink_url, "post_call", False),
|
||||
]
|
||||
path: Final = root / "straiker.yaml"
|
||||
|
|
@ -786,6 +793,36 @@ def test_explicit_api_version_v1_overrides_key_prefix(rig: Rig) -> None:
|
|||
assert calls[0].headers["x-straiker-webhook-format"] == "litellm"
|
||||
|
||||
|
||||
def test_stray_api_version_with_v3_key_still_enforces_on_v3(rig: Rig) -> None:
|
||||
allowed_marker: Final = rig.marker()
|
||||
allowed: Final = _chat(rig, "stray version " + allowed_marker, guardrails=["straiker-v3-stray-version"])
|
||||
assert allowed.status_code == 200, allowed.text
|
||||
assert len(_v3_request_calls(rig, allowed_marker, agent=None)) == 1
|
||||
assert len(rig.provider_calls(allowed_marker, rig.provider_drain())) == 1
|
||||
|
||||
blocked_marker: Final = rig.marker()
|
||||
blocked: Final = _chat(rig, f"{STRAY_V3_BLOCK_MARK} {blocked_marker}", guardrails=["straiker-v3-stray-version"])
|
||||
assert blocked.status_code == 400, blocked.text
|
||||
assert blocked.json()["error"]["message"] == BLOCK_MESSAGE, blocked.text
|
||||
assert len(_v3_request_calls(rig, blocked_marker, agent=None)) == 1
|
||||
assert rig.provider_calls(blocked_marker, rig.provider_drain()) == ()
|
||||
|
||||
|
||||
def test_empty_api_version_with_v1_key_still_enforces_on_v1(rig: Rig) -> None:
|
||||
allowed_marker: Final = rig.marker()
|
||||
allowed: Final = _chat(rig, "empty version " + allowed_marker, guardrails=["straiker-v1-empty-version"])
|
||||
assert allowed.status_code == 200, allowed.text
|
||||
assert len(_v1_calls(rig, allowed_marker, V1_KEY)) == 1
|
||||
assert len(rig.provider_calls(allowed_marker, rig.provider_drain())) == 1
|
||||
|
||||
blocked_marker: Final = rig.marker()
|
||||
blocked: Final = _chat(rig, f"{V1_BLOCK_MARK} {blocked_marker}", guardrails=["straiker-v1-empty-version"])
|
||||
assert blocked.status_code == 400, blocked.text
|
||||
assert blocked.json()["error"]["message"] == BLOCK_MESSAGE, blocked.text
|
||||
assert len(_v1_calls(rig, blocked_marker, V1_KEY)) == 1
|
||||
assert rig.provider_calls(blocked_marker, rig.provider_drain()) == ()
|
||||
|
||||
|
||||
# E: configured client and format_hint ride as headers; request header for agent fills in when YAML has none
|
||||
def test_v3_client_and_format_hint_headers_and_request_agent_header(rig: Rig) -> None:
|
||||
marker: Final = rig.marker()
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -1216,6 +1216,58 @@ def test_v3_initializer_reads_api_version_from_config():
|
|||
assert g._webhook_url().endswith("/api/v3/detect")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("api_version", ["2024-09-01", "", "v2"])
|
||||
@pytest.mark.parametrize(("api_key", "expected"), [("c4ac433a-uuid", "v1"), (V3_KEY, "v3")])
|
||||
def test_unknown_api_version_follows_key_prefix(api_version, api_key, expected, monkeypatch):
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", litellm.callbacks.copy())
|
||||
|
||||
with patch.object(verbose_proxy_logger, "warning") as warning:
|
||||
g = initialize_guardrail(
|
||||
LitellmParams(guardrail="straiker", mode="pre_call", api_key=api_key, api_version=api_version),
|
||||
Guardrail(guardrail_name="straiker", litellm_params={"guardrail": "straiker", "mode": "pre_call"}),
|
||||
)
|
||||
|
||||
assert g.api_version == expected
|
||||
expected_path = "/api/v3/detect" if expected == "v3" else "/api/v1/detect/webhook"
|
||||
assert g._webhook_url().endswith(expected_path)
|
||||
warning.assert_called_once()
|
||||
assert warning.call_args.args[-1] == api_version
|
||||
|
||||
|
||||
def test_init_guardrails_v2_registers_straiker_with_unknown_api_version(monkeypatch):
|
||||
import litellm
|
||||
from litellm.proxy.guardrails import guardrail_registry
|
||||
from litellm.proxy.guardrails.guardrail_registry import InMemoryGuardrailHandler
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
|
||||
handler = InMemoryGuardrailHandler()
|
||||
monkeypatch.setattr(guardrail_registry, "IN_MEMORY_GUARDRAIL_HANDLER", handler)
|
||||
monkeypatch.setattr(litellm, "callbacks", litellm.callbacks.copy())
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
{
|
||||
"guardrail_name": "straiker-unknown-version",
|
||||
"litellm_params": {
|
||||
"guardrail": "straiker",
|
||||
"mode": "pre_call",
|
||||
"api_key": V3_KEY,
|
||||
"api_version": "2024-09-01",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
callbacks = tuple(handler.guardrail_id_to_custom_guardrail.values())
|
||||
assert len(callbacks) == 1
|
||||
assert isinstance(callbacks[0], StraikerGuardrail)
|
||||
assert callbacks[0].api_version == "v3"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_v3_request_phase_relays_the_provider_body_and_nothing_else():
|
||||
g = _make_guardrail(api_key=V3_KEY, source="Yum Gateway")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue