diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py index acfbf176d0a..b1dc368b8ee 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py @@ -8,6 +8,7 @@ import datetime import hashlib import os import re +import time from collections import OrderedDict from typing import TYPE_CHECKING, Any, Dict, List, Literal, NamedTuple, Optional, Type @@ -32,6 +33,8 @@ if TYPE_CHECKING: BLOCKED_BY_OVALIX_FALLBACK_MESSAGE = "This message was blocked by Ovalix" BLOCKED_ACTION_TYPE = "block" +_ROUTING_CACHE_TTL_SECONDS = 3600 +_ROUTING_CACHE_MAX_SIZE = 1000 class ResolvedRouting(NamedTuple): @@ -302,6 +305,109 @@ class OvalixGuardrail(CustomGuardrail): return modified["content"] return None + def _get_key_alias(self, request_data: dict) -> str | None: + metadata = {**(request_data.get("metadata") or {}), **(request_data.get("litellm_metadata") or {})} + return metadata.get("user_api_key_alias") or metadata.get("user_api_key_key_alias") + + async def _get_app_name_regex(self) -> re.Pattern[str]: + if self._app_name_regex is not None: + return self._app_name_regex + url = f"{self._tracker_api_base}/tracking/custom_application/litellm_app_name_regex" + try: + response = await self._async_handler.get(url, headers=dict(self._tracker_headers)) + response.raise_for_status() + compiled = re.compile(response.json()["regex"]) + except Exception as e: + verbose_proxy_logger.exception("Ovalix app-name regex fetch failed: %s", e) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Ovalix guardrail error: app-name regex fetch failed: {e!s}", + should_wrap_with_default_message=False, + ) from e + self._app_name_regex = compiled + return compiled + + def _extract_application_name(self, alias: str, regex: re.Pattern[str]) -> str | None: + match = regex.search(alias) + if not match: + return None + name = (match.group(1) if match.groups() else match.group(0)).strip() + return name or None + + def _routing_cache_get(self, name: str) -> ResolvedRouting | None: + entry = self._routing_cache.get(name) + if entry is None: + return None + stored_at, routing = entry + if time.monotonic() - stored_at >= _ROUTING_CACHE_TTL_SECONDS: + del self._routing_cache[name] + return None + self._routing_cache.move_to_end(name) + return routing + + def _routing_cache_put(self, name: str, routing: ResolvedRouting) -> None: + self._routing_cache[name] = (time.monotonic(), routing) + self._routing_cache.move_to_end(name) + while len(self._routing_cache) > _ROUTING_CACHE_MAX_SIZE: + self._routing_cache.popitem(last=False) + + async def _resolve_routing(self, request_data: dict) -> ResolvedRouting: + if self._application_id: + return ResolvedRouting( + self._application_id, + self._pre_checkpoint_id, + self._post_checkpoint_id, + self._file_checkpoint_id, + None, + ) + alias = self._get_key_alias(request_data) + if not alias: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="Ovalix guardrail error: no application_id configured and no user_api_key_alias to resolve by", + should_wrap_with_default_message=False, + ) + regex = await self._get_app_name_regex() + name = self._extract_application_name(alias, regex) + if not name: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message="Ovalix guardrail error: could not extract an application name from the api key alias", + should_wrap_with_default_message=False, + ) + if self._enable_routing_cache: + cached = self._routing_cache_get(name) + if cached is not None: + return cached + routing = await self._resolve_via_tracker(name) + if self._enable_routing_cache: + self._routing_cache_put(name, routing) + return routing + + async def _resolve_via_tracker(self, application_name: str) -> ResolvedRouting: + url = f"{self._tracker_api_base}/tracking/custom_application/resolve_litellm_application" + try: + response = await self._async_handler.post( + url, headers=dict(self._tracker_headers), json={"application_name": application_name} + ) + response.raise_for_status() + body = response.json() + routing = ResolvedRouting( + application_id=str(body["application_id"]), + checkpoint_id_pre=body.get("checkpoint_id_pre"), + checkpoint_id_post=body.get("checkpoint_id_post"), + checkpoint_id_pre_file=body.get("checkpoint_id_pre_file"), + checkpoint_id_post_file=body.get("checkpoint_id_post_file"), + ) + except Exception as e: + verbose_proxy_logger.exception("Ovalix routing resolution failed: %s", e) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Ovalix guardrail error: routing resolution failed: {e!s}", + should_wrap_with_default_message=False, + ) from e + return routing + @staticmethod def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py index 0bf456830f0..d6a8f3bf2e1 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py @@ -15,9 +15,18 @@ from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix import ( OvalixGuardrail, OvalixGuardrailBlockedException, OvalixGuardrailMissingSecrets, + ResolvedRouting, ) from litellm.types.utils import GenericGuardrailAPIInputs + +@pytest.fixture(autouse=True) +def _clear_ovalix_env(monkeypatch): + for key in list(os.environ.keys()): + if key.startswith("OVALIX_"): + monkeypatch.delenv(key, raising=False) + + # Example Tracker responses (as returned by the checkpoint API) TRACKER_RESPONSE_ALLOW = { "action_type": "allow", @@ -217,9 +226,7 @@ class TestOvalixGuardrail: mock_response.json.return_value = TRACKER_RESPONSE_ALLOW mock_response.raise_for_status = MagicMock() - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.return_value = mock_response result = await guardrail._call_checkpoint( content="hello", @@ -231,9 +238,7 @@ class TestOvalixGuardrail: assert result == TRACKER_RESPONSE_ALLOW mock_post.assert_called_once() call_args = mock_post.call_args - assert call_args.args[0] == ( - "https://tracker.test/tracking/custom_application/checkpoint" - ) + assert call_args.args[0] == ("https://tracker.test/tracking/custom_application/checkpoint") body = call_args.kwargs["json"] assert body["application_id"] == "app-1" assert body["checkpoint_id"] == "pre-1" @@ -263,9 +268,7 @@ class TestOvalixGuardrail: mock_response.json.return_value = TRACKER_RESPONSE_ALLOW mock_response.raise_for_status = MagicMock() - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.return_value = mock_response result = await guardrail.apply_guardrail( inputs=inputs, @@ -289,9 +292,7 @@ class TestOvalixGuardrail: try: guardrail = OvalixGuardrail(**_guardrail_kwargs()) inputs = GenericGuardrailAPIInputs( - structured_messages=[ - {"role": "user", "content": "Hello, my name is David."} - ], + structured_messages=[{"role": "user", "content": "Hello, my name is David."}], texts=["Hello, my name is David."], ) request_data = {} @@ -300,9 +301,7 @@ class TestOvalixGuardrail: mock_response.json.return_value = TRACKER_RESPONSE_ANONYMIZE mock_response.raise_for_status = MagicMock() - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.return_value = mock_response result = await guardrail.apply_guardrail( inputs=inputs, @@ -335,9 +334,7 @@ class TestOvalixGuardrail: mock_response.json.return_value = TRACKER_RESPONSE_BLOCK mock_response.raise_for_status = MagicMock() - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.return_value = mock_response with pytest.raises(OvalixGuardrailBlockedException) as exc_info: await guardrail.apply_guardrail( @@ -382,9 +379,7 @@ class TestOvalixGuardrail: resp.raise_for_status = MagicMock() return resp - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.side_effect = side_effect result = await guardrail.apply_guardrail( inputs=inputs, @@ -411,9 +406,7 @@ class TestOvalixGuardrail: try: guardrail = OvalixGuardrail(**_guardrail_kwargs()) inputs = GenericGuardrailAPIInputs( - structured_messages=[ - {"role": "assistant", "content": "Safe assistant reply"} - ], + structured_messages=[{"role": "assistant", "content": "Safe assistant reply"}], texts=["Safe assistant reply"], ) request_data = {} @@ -422,9 +415,7 @@ class TestOvalixGuardrail: mock_response.json.return_value = TRACKER_RESPONSE_ALLOW mock_response.raise_for_status = MagicMock() - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.return_value = mock_response result = await guardrail.apply_guardrail( inputs=inputs, @@ -454,9 +445,7 @@ class TestOvalixGuardrail: mock_response.json.return_value = TRACKER_RESPONSE_BLOCK mock_response.raise_for_status = MagicMock() - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.return_value = mock_response with pytest.raises(OvalixGuardrailBlockedException) as exc_info: await guardrail.apply_guardrail( @@ -471,9 +460,7 @@ class TestOvalixGuardrail: assert mock_post.call_count == 1 @pytest.mark.asyncio - async def test_apply_guardrail_request_missing_modified_data_uses_original_content( - self, guardrail_with_env - ): + async def test_apply_guardrail_request_missing_modified_data_uses_original_content(self, guardrail_with_env): """When Tracker response has no modified_data.content, original content is used.""" guardrail = guardrail_with_env inputs = GenericGuardrailAPIInputs( @@ -492,9 +479,7 @@ class TestOvalixGuardrail: } mock_response.raise_for_status = MagicMock() - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.return_value = mock_response result = await guardrail.apply_guardrail( inputs=inputs, @@ -507,9 +492,7 @@ class TestOvalixGuardrail: assert mock_post.call_count == 1 @pytest.mark.asyncio - async def test_apply_guardrail_tracker_http_error_raises_guardrail_exception( - self, guardrail_with_env - ): + async def test_apply_guardrail_tracker_http_error_raises_guardrail_exception(self, guardrail_with_env): """When Tracker returns HTTP error (e.g. 400), GuardrailRaisedException is raised.""" guardrail = guardrail_with_env inputs = GenericGuardrailAPIInputs( @@ -525,9 +508,7 @@ class TestOvalixGuardrail: response=MagicMock(status_code=400), ) - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: mock_post.return_value = mock_response with pytest.raises(GuardrailRaisedException): await guardrail.apply_guardrail( @@ -580,9 +561,7 @@ class TestOvalixGuardrail: inputs = GenericGuardrailAPIInputs(structured_messages=[], texts=[]) request_data = {} - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: result = await guardrail.apply_guardrail( inputs=inputs, request_data=request_data, @@ -603,22 +582,9 @@ class TestOvalixGuardrail: os.environ[k] = v try: guardrail = OvalixGuardrail(**_guardrail_kwargs()) - assert ( - guardrail._get_actor( - {"metadata": {"user_api_key_user_email": "a@b.com"}} - ) - == "a@b.com" - ) - assert ( - guardrail._get_actor({"metadata": {"user_api_key_user_id": "uid-1"}}) - == "uid-1" - ) - assert ( - guardrail._get_actor( - {"litellm_metadata": {"user_api_key_user_id": "uid-2"}} - ) - == "uid-2" - ) + assert guardrail._get_actor({"metadata": {"user_api_key_user_email": "a@b.com"}}) == "a@b.com" + assert guardrail._get_actor({"metadata": {"user_api_key_user_id": "uid-1"}}) == "uid-1" + assert guardrail._get_actor({"litellm_metadata": {"user_api_key_user_id": "uid-2"}}) == "uid-2" assert guardrail._get_actor({}) == "unknown" finally: for k in _ovalix_env(): @@ -656,9 +622,7 @@ class TestOvalixGuardrail: assert session_id_1 == session_id_2 assert "app-1" in session_id_1 - def test_block_current_message_raises_ovalix_blocked_exception( - self, guardrail_with_env - ): + def test_block_current_message_raises_ovalix_blocked_exception(self, guardrail_with_env): """_block_current_message raises OvalixGuardrailBlockedException with status_code 400.""" guardrail = guardrail_with_env with pytest.raises(OvalixGuardrailBlockedException) as exc_info: @@ -670,16 +634,11 @@ class TestOvalixGuardrail: """_get_trackers_corrected_message returns modified_data.content or None.""" guardrail = guardrail_with_env assert ( - guardrail._get_trackers_corrected_message( - {"modified_data": {"content": "corrected text"}} - ) + guardrail._get_trackers_corrected_message({"modified_data": {"content": "corrected text"}}) == "corrected text" ) assert guardrail._get_trackers_corrected_message({"modified_data": {}}) is None - assert ( - guardrail._get_trackers_corrected_message({"modified_data": "not-a-dict"}) - is None - ) + assert guardrail._get_trackers_corrected_message({"modified_data": "not-a-dict"}) is None @pytest.mark.asyncio async def test_apply_guardrail_response_no_texts_returns_unchanged(self): @@ -691,9 +650,7 @@ class TestOvalixGuardrail: inputs = GenericGuardrailAPIInputs() request_data = {} - with patch.object( - guardrail._async_handler, "post", new_callable=AsyncMock - ) as mock_post: + with patch.object(guardrail._async_handler, "post", new_callable=AsyncMock) as mock_post: result = await guardrail.apply_guardrail( inputs=inputs, request_data=request_data, @@ -707,3 +664,144 @@ class TestOvalixGuardrail: for k in _ovalix_env(): if k in os.environ: del os.environ[k] + + +_REGEX = r"^\s*\[([^\]]+)\]" +_ROUTING_BODY = { + "application_id": "app-9", + "checkpoint_id_pre": "pre", + "checkpoint_id_post": "post", + "checkpoint_id_pre_file": None, + "checkpoint_id_post_file": None, +} + + +def _discovery_guardrail(enable_cache=True): + return OvalixGuardrail( + tracker_api_base="https://tracker.test", + tracker_api_key="key", + enable_routing_cache=enable_cache, + guardrail_name="ovalix-test", + event_hook="pre_call", + default_on=True, + ) + + +def _alias_request_data(alias="[Weather App] prod"): + return {"metadata": {"user_api_key_alias": alias, "user_api_key_user_email": "u@x.com"}} + + +def _mock_handler(g, routing=None): + get_resp = MagicMock() + get_resp.json.return_value = {"regex": _REGEX} + get_resp.raise_for_status = MagicMock() + post_resp = MagicMock() + post_resp.json.return_value = routing or _ROUTING_BODY + post_resp.raise_for_status = MagicMock() + g._async_handler.get = AsyncMock(return_value=get_resp) + g._async_handler.post = AsyncMock(return_value=post_resp) + return g._async_handler.get, g._async_handler.post + + +@pytest.mark.asyncio +async def test_static_mode_uses_config_routing(): + g = OvalixGuardrail( + tracker_api_base="https://t", + tracker_api_key="k", + application_id="app-1", + pre_checkpoint_id="pre-1", + post_checkpoint_id="post-1", + file_checkpoint_id="file-1", + guardrail_name="o", + event_hook="pre_call", + default_on=True, + ) + routing = await g._resolve_routing({}) + assert routing == ResolvedRouting("app-1", "pre-1", "post-1", "file-1", None) + + +@pytest.mark.asyncio +async def test_discovery_extracts_name_and_resolves(): + g = _discovery_guardrail(enable_cache=False) + mock_get, mock_post = _mock_handler(g) + routing = await g._resolve_routing(_alias_request_data("[Weather App] prod")) + assert routing.application_id == "app-9" + assert mock_post.call_args.args[0].endswith("/tracking/custom_application/resolve_litellm_application") + assert mock_post.call_args.kwargs["json"] == {"application_name": "Weather App"} + assert mock_get.call_args.args[0].endswith("/tracking/custom_application/litellm_app_name_regex") + + +@pytest.mark.asyncio +async def test_regex_fetched_once_even_with_cache_off(): + g = _discovery_guardrail(enable_cache=False) + mock_get, mock_post = _mock_handler(g) + await g._resolve_routing(_alias_request_data()) + await g._resolve_routing(_alias_request_data()) + assert mock_get.call_count == 1 + assert mock_post.call_count == 2 + + +@pytest.mark.asyncio +async def test_no_bracket_alias_fails_closed(): + g = _discovery_guardrail() + _mock_handler(g) + with pytest.raises(GuardrailRaisedException): + await g._resolve_routing(_alias_request_data("no brackets here")) + + +@pytest.mark.asyncio +async def test_discovery_missing_alias_raises(): + g = _discovery_guardrail() + _mock_handler(g) + with pytest.raises(GuardrailRaisedException): + await g._resolve_routing({"metadata": {"user_api_key_user_email": "u@x.com"}}) + + +@pytest.mark.asyncio +async def test_regex_fetch_failure_raises_guardrail_exception(): + g = _discovery_guardrail(enable_cache=False) + g._async_handler.get = AsyncMock(side_effect=httpx.ConnectError("boom")) + with pytest.raises(GuardrailRaisedException): + await g._resolve_routing(_alias_request_data()) + + +@pytest.mark.asyncio +async def test_resolve_endpoint_failure_raises_guardrail_exception(): + g = _discovery_guardrail(enable_cache=False) + _mock_handler(g) + g._async_handler.post = AsyncMock(side_effect=httpx.ConnectError("boom")) + with pytest.raises(GuardrailRaisedException): + await g._resolve_routing(_alias_request_data()) + + +@pytest.mark.asyncio +async def test_resolve_missing_application_id_raises_guardrail_exception(): + g = _discovery_guardrail(enable_cache=False) + _mock_handler(g, routing={"checkpoint_id_pre": "pre", "checkpoint_id_post": "post"}) + with pytest.raises(GuardrailRaisedException): + await g._resolve_routing(_alias_request_data()) + + +@pytest.mark.asyncio +async def test_routing_cache_hit_and_ttl_expiry(monkeypatch): + g = _discovery_guardrail(enable_cache=True) + mock_get, mock_post = _mock_handler(g) + clock = [1000.0] + monkeypatch.setattr("litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix.time.monotonic", lambda: clock[0]) + await g._resolve_routing(_alias_request_data()) + clock[0] = 1000.0 + 3599 + await g._resolve_routing(_alias_request_data()) + assert mock_post.call_count == 1 + clock[0] = 1000.0 + 3601 + await g._resolve_routing(_alias_request_data()) + assert mock_post.call_count == 2 + + +@pytest.mark.asyncio +async def test_routing_cache_lru_eviction(monkeypatch): + monkeypatch.setattr("litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix._ROUTING_CACHE_MAX_SIZE", 2) + g = _discovery_guardrail(enable_cache=True) + _mock_handler(g) + for name in ("[App A] x", "[App B] x", "[App C] x"): + await g._resolve_routing(_alias_request_data(name)) + assert "App A" not in g._routing_cache and len(g._routing_cache) == 2