fix(guardrails): treat an unknown straiker api_version as unset instead of skipping the guardrail (#43956)

The integration test hunks are dropped because tests/integration/observability/test_straiker_v3_platform.py is not on this line

(cherry picked from commit a3a7650569)
This commit is contained in:
devin-ai-integration[bot] 2026-09-30 18:20:57 -07:00 • committed by yucheng
parent 7d55a4ebf8
commit d34f5b3a4e
2 changed files with 67 additions and 2 deletions

View file

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

View file

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