diff --git a/litellm/proxy/guardrails/guardrail_hooks/straiker/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/straiker/__init__.py index 7e3f23fec86..cb037fb7513 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/straiker/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/straiker/__init__.py @@ -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", diff --git a/tests/integration/observability/test_straiker_v3_platform.py b/tests/integration/observability/test_straiker_v3_platform.py index e44abf4e066..c4d34b1a9a7 100644 --- a/tests/integration/observability/test_straiker_v3_platform.py +++ b/tests/integration/observability/test_straiker_v3_platform.py @@ -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() diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py index 05260cfe5e3..e52f9c96971 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py @@ -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")