diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py index 585f27f8517..d962057617e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py @@ -22,7 +22,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import GenericGuardrailAPIInputs +from litellm.types.utils import AllMessageValues, GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -120,50 +120,56 @@ class OvalixGuardrail(CustomGuardrail): self, supported_event_hooks: List[GuardrailEventHooks] ) -> None: """Ensure required secrets and checkpoint IDs are set; auto-add hooks when IDs are present.""" + errors: List[str] = [] + if not self._tracker_api_base: - raise OvalixGuardrailMissingSecrets( - "Ovalix Tracker API base required. Set OVALIX_TRACKER_API_BASE or pass tracker_api_base in litellm_params." + errors.append( + "Tracker API base, set OVALIX_TRACKER_API_BASE or pass tracker_api_base" ) if not self._tracker_api_key: - raise OvalixGuardrailMissingSecrets( - "Ovalix Tracker API key required. Set OVALIX_TRACKER_API_KEY or pass tracker_api_key in litellm_params." + errors.append( + "Tracker API key, set OVALIX_TRACKER_API_KEY or pass tracker_api_key" ) if not self._application_id: - raise OvalixGuardrailMissingSecrets( - "Ovalix Application ID required. Set OVALIX_APPLICATION_ID or pass application_id in litellm_params." + errors.append( + "Application ID, set OVALIX_APPLICATION_ID or pass application_id" ) - if ( not self._pre_checkpoint_id and GuardrailEventHooks.pre_call in supported_event_hooks ): - raise OvalixGuardrailMissingSecrets( - "Ovalix Pre-checkpoint ID required. Set OVALIX_PRE_CHECKPOINT_ID or pass pre_checkpoint_id in litellm_params." + errors.append( + "Pre-checkpoint ID, set OVALIX_PRE_CHECKPOINT_ID or pass pre_checkpoint_id" ) - elif ( - self._pre_checkpoint_id - and GuardrailEventHooks.pre_call not in supported_event_hooks - ): - supported_event_hooks.append(GuardrailEventHooks.pre_call) - if ( not self._post_checkpoint_id and GuardrailEventHooks.post_call in supported_event_hooks ): - raise OvalixGuardrailMissingSecrets( - "Ovalix Post-checkpoint ID required. Set OVALIX_POST_CHECKPOINT_ID or pass post_checkpoint_id in litellm_params." + errors.append( + "Post-checkpoint ID, set OVALIX_POST_CHECKPOINT_ID or pass post_checkpoint_id" ) - elif ( + if not self._pre_checkpoint_id and not self._post_checkpoint_id: + errors.append( + "Pre-checkpoint ID or Post-checkpoint ID, set OVALIX_PRE_CHECKPOINT_ID or OVALIX_POST_CHECKPOINT_ID or pass pre_checkpoint_id or post_checkpoint_id" + ) + + if errors: + raise OvalixGuardrailMissingSecrets( + "Missing Ovalix guardrail configuration errors: " + ". ".join(errors) + ) + + # auto-add hooks when checkpoint IDs are present + if ( + self._pre_checkpoint_id + and GuardrailEventHooks.pre_call not in supported_event_hooks + ): + supported_event_hooks.append(GuardrailEventHooks.pre_call) + if ( self._post_checkpoint_id and GuardrailEventHooks.post_call not in supported_event_hooks ): supported_event_hooks.append(GuardrailEventHooks.post_call) - if not self._pre_checkpoint_id and not self._post_checkpoint_id: - raise OvalixGuardrailMissingSecrets( - "Ovalix Pre-checkpoint ID or Post-checkpoint ID required. Set OVALIX_PRE_CHECKPOINT_ID or OVALIX_POST_CHECKPOINT_ID or pass pre_checkpoint_id or post_checkpoint_id in litellm_params." - ) - def _get_actor(self, data: dict) -> str: """Return a stable actor identifier from request metadata (e.g. user email or id).""" metadata = data.get("metadata") or data.get("litellm_metadata") or {} @@ -201,9 +207,7 @@ class OvalixGuardrail(CustomGuardrail): "data_type": "TEXT", "data": {"content": content}, } - response = await self._async_handler.post( - url, headers=headers, json=payload - ) + response = await self._async_handler.post(url, headers=headers, json=payload) response.raise_for_status() return response.json() @@ -250,13 +254,13 @@ class OvalixGuardrail(CustomGuardrail): # TODO: set the llm response text to `corrected_llm_response`. will be addressed later. return inputs - messages = request_data.get("messages") or [] + messages = inputs.get("structured_messages") or [] if not messages: return inputs if self._pre_checkpoint_id: post_guardrail_texts = await self._generate_post_guardrail_text( - messages, actor, session_id, request_data + messages, actor, session_id ) return {**inputs, "texts": post_guardrail_texts} return inputs @@ -318,10 +322,9 @@ class OvalixGuardrail(CustomGuardrail): async def _generate_post_guardrail_text( self, - messages: List[Dict[str, Any]], + messages: List[AllMessageValues], actor: str, session_id: str, - request_data: dict, ) -> List[str]: """ Generate post-guardrail text for the given messages. @@ -348,7 +351,7 @@ class OvalixGuardrail(CustomGuardrail): continue message_role = message.get("role", None) if message_role and message_role != USER_MESSAGE_ROLE: - # we are not scanning the LLM/system/developer past responses, only the response that the user sent + # we are not scanning the assistant/system/developer past responses, only the responses that the user sent post_guardrail_texts.insert(0, content) continue try: 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 89fcf4e8c3f..a76267e808e 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py @@ -1,197 +1,621 @@ """ -Unit tests for Ovalix guardrail types (OvalixGuardrailConfigModel) and config model resolution. +Unit tests for Ovalix guardrail: config resolution and apply_guardrail behavior +with mocked Tracker service responses (allow, anonymize, block). """ -from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix import OvalixGuardrail +import os +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix import ( + OvalixGuardrail, + OvalixGuardrailBlockedException, + OvalixGuardrailMissingSecrets, +) +from litellm.types.utils import GenericGuardrailAPIInputs + +# Example Tracker responses (as returned by the checkpoint API) +TRACKER_RESPONSE_ALLOW = { + "action_type": "allow", + "data_type": "TEXT", + "original_data": {"content": "how are you?"}, + "modified_data": {"content": "how are you?"}, + "alerts": [], +} + +TRACKER_RESPONSE_ANONYMIZE = { + "action_type": "anonymize", + "data_type": "TEXT", + "original_data": {"content": "Hello, my name is David."}, + "modified_data": {"content": "Hello, my name is {Name}. How are you?"}, + "alerts": [ + { + "title": "Sensitive Data Alert", + "subtitle": "We've identified that you were trying to share sensitive information", + "alerts": ["Name:\tDavid\nRedacted to:\t{Name}"], + } + ], +} + +TRACKER_RESPONSE_BLOCK = { + "action_type": "block", + "data_type": "TEXT", + "original_data": {"content": "I am 15 YO"}, + "modified_data": {"content": "This message was blocked by Ovalix"}, + "alerts": [ + { + "title": "Sensitive Data Alert", + "subtitle": "We've identified that you were trying to share sensitive information", + "alerts": ["Age:\t15\nBlocked"], + } + ], +} + + +def _ovalix_env(): + return { + "OVALIX_TRACKER_API_BASE": "https://tracker.test", + "OVALIX_TRACKER_API_KEY": "key", + "OVALIX_APPLICATION_ID": "app-1", + "OVALIX_PRE_CHECKPOINT_ID": "pre-1", + "OVALIX_POST_CHECKPOINT_ID": "post-1", + } + + +def _guardrail_kwargs(): + return { + "guardrail_name": "ovalix-test", + "event_hook": "pre_call", + "default_on": True, + } class TestOvalixGuardrailConfigModel: - """Tests for OvalixGuardrailConfigModel from litellm.types.proxy.guardrails.guardrail_hooks.ovalix.""" + """Minimal config model tests: wiring only.""" - def test_config_model_ui_friendly_name(self): - """Test that config model has correct UI friendly name.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, - ) - - assert OvalixGuardrailConfigModel.ui_friendly_name() == "Ovalix Guardrail" - - def test_config_model_fields(self): - """Test that config model has expected fields and default values.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, - ) - - model = OvalixGuardrailConfigModel() - - assert model.tracker_api_base is None - assert model.tracker_api_key is None - assert model.application_id is None - assert model.pre_checkpoint_id is None - assert model.post_checkpoint_id is None - - def test_get_config_model(self): - """Test get_config_model returns OvalixGuardrailConfigModel.""" + def test_get_config_model_returns_ovalix_config_model(self): + """get_config_model returns OvalixGuardrailConfigModel for proxy/config wiring.""" config_model = OvalixGuardrail.get_config_model() assert config_model is not None assert config_model.__name__ == "OvalixGuardrailConfigModel" - assert hasattr(config_model, "ui_friendly_name") assert config_model.ui_friendly_name() == "Ovalix Guardrail" - def test_config_model_with_all_fields_set(self): - """Test construction with all optional fields set to explicit values.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, + +class TestOvalixGuardrail: + """Behavioral tests with mocked Tracker checkpoint API.""" + + def setup_method(self): + for key in list(os.environ.keys()): + if key.startswith("OVALIX_"): + del os.environ[key] + + def teardown_method(self): + for key in list(os.environ.keys()): + if key.startswith("OVALIX_"): + del os.environ[key] + + @pytest.fixture + def guardrail_with_env(self): + """Guardrail with OVALIX_* env set; cleans up in teardown.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + yield OvalixGuardrail(**_guardrail_kwargs()) + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + def test_initialization_requires_secrets(self): + """Initialization raises when required Tracker/application/checkpoint config is missing.""" + with pytest.raises(OvalixGuardrailMissingSecrets): + OvalixGuardrail( + guardrail_name="ovalix-test", + event_hook="pre_call", + default_on=True, + ) + + def test_initialization_with_explicit_params(self): + """Guardrail initializes with explicit tracker base, key, app and checkpoint IDs.""" + guardrail = OvalixGuardrail( + tracker_api_base="https://tracker.example", + tracker_api_key="secret", + application_id="app-x", + pre_checkpoint_id="pre-x", + post_checkpoint_id="post-x", + **_guardrail_kwargs(), ) + assert guardrail._tracker_api_base == "https://tracker.example" + assert guardrail._application_id == "app-x" + assert guardrail._pre_checkpoint_id == "pre-x" + assert guardrail._post_checkpoint_id == "post-x" - model = OvalixGuardrailConfigModel( - tracker_api_base="https://tracker.ovalix.example", - tracker_api_key="key-123", - application_id="app-456", - pre_checkpoint_id="pre-cp-1", - post_checkpoint_id="post-cp-1", + def test_initialization_with_env_vars(self): + """Guardrail picks up OVALIX_* env vars when params not passed.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + assert guardrail._tracker_api_base == "https://tracker.test" + assert guardrail._tracker_api_key == "key" + assert guardrail._application_id == "app-1" + assert guardrail._pre_checkpoint_id == "pre-1" + assert guardrail._post_checkpoint_id == "post-1" + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_call_checkpoint_sends_correct_payload_and_returns_json(self): + """_call_checkpoint POSTs to tracker with application_id, checkpoint_id, actor, session_id, data.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + mock_response = MagicMock() + 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: + mock_post.return_value = mock_response + result = await guardrail._call_checkpoint( + content="hello", + checkpoint_id="pre-1", + actor="user@test.com", + session_id="session-1", + ) + + 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" + ) + body = call_args.kwargs["json"] + assert body["application_id"] == "app-1" + assert body["checkpoint_id"] == "pre-1" + assert body["actor"] == "user@test.com" + assert body["session_id"] == "session-1" + assert body["data_type"] == "TEXT" + assert body["data"] == {"content": "hello"} + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_allow_passes_through(self): + """When Tracker returns allow, apply_guardrail returns inputs with texts set to modified_data content.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "how are you?"}] + ) + request_data = {} + + mock_response = MagicMock() + 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: + mock_post.return_value = mock_response + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result.get("texts") == ["how are you?"] + assert mock_post.call_count == 1 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_anonymize_returns_modified_text(self): + """When Tracker returns anonymize, apply_guardrail returns texts with modified_data content.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[ + {"role": "user", "content": "Hello, my name is David."} + ] + ) + request_data = {} + + mock_response = MagicMock() + 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: + mock_post.return_value = mock_response + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result.get("texts") == ["Hello, my name is {Name}. How are you?"] + assert mock_post.call_count == 1 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_block_raises_with_tracker_message(self): + """When Tracker returns block on the (chronologically) last user message, OvalixGuardrailBlockedException is raised.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "I am 15 YO"}] + ) + request_data = {} + + mock_response = MagicMock() + 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: + mock_post.return_value = mock_response + with pytest.raises(OvalixGuardrailBlockedException) as exc_info: + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert "This message was blocked by Ovalix" in str(exc_info.value.message) + assert exc_info.value.status_code == 400 + assert mock_post.call_count == 1 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_block_non_last_replaced_in_texts(self): + """When Tracker returns block on a non-last user message, that message is replaced in texts and no exception is raised.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[ + {"role": "user", "content": "I am 15 YO"}, + {"role": "user", "content": "how are you?"}, + ] + ) + request_data = {} + + def side_effect(*args, **kwargs): + body = kwargs.get("json", {}) + content = (body.get("data") or {}).get("content", "") + resp = MagicMock() + if "15" in content: + resp.json.return_value = TRACKER_RESPONSE_BLOCK + else: + resp.json.return_value = TRACKER_RESPONSE_ALLOW + resp.raise_for_status = MagicMock() + return resp + + 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, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result.get("texts") == [ + "This message was blocked by Ovalix", + "how are you?", + ] + assert mock_post.call_count == 2 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_response_allow_returns_inputs(self): + """When input_type is response and Tracker allows, apply_guardrail returns inputs unchanged.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs() + request_data = {"response": None} + + mock_response = MagicMock() + 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, patch.object( + guardrail, + "_get_llm_response_text", + return_value="Safe assistant reply", + ): + mock_post.return_value = mock_response + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + logging_obj=None, + ) + + assert result == inputs + assert mock_post.call_count == 1 + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_response_block_returns_inputs( + self, guardrail_with_env + ): + """When Tracker blocks on response, apply_guardrail still returns inputs (no raise).""" + guardrail = guardrail_with_env + inputs = GenericGuardrailAPIInputs() + request_data = {"response": None} + + mock_response = MagicMock() + 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, patch.object( + guardrail, + "_get_llm_response_text", + return_value="I am 15 YO", + ): + mock_post.return_value = mock_response + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="response", + logging_obj=None, + ) + + assert result == inputs + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_apply_guardrail_request_non_user_messages_not_sent_to_tracker( + self, guardrail_with_env + ): + """Only user messages are sent to Tracker; system/assistant content is passed through.""" + guardrail = guardrail_with_env + inputs = GenericGuardrailAPIInputs( + structured_messages=[ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "hello"}, + ] ) + request_data = {} - assert model.tracker_api_base == "https://tracker.ovalix.example" - assert model.tracker_api_key == "key-123" - assert model.application_id == "app-456" - assert model.pre_checkpoint_id == "pre-cp-1" - assert model.post_checkpoint_id == "post-cp-1" - - def test_config_model_with_partial_fields(self): - """Test construction with only a subset of fields set.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, - ) - - model = OvalixGuardrailConfigModel( - tracker_api_base="https://custom.tracker", - application_id="app-only", - ) - - assert model.tracker_api_base == "https://custom.tracker" - assert model.application_id == "app-only" - assert model.tracker_api_key is None - assert model.pre_checkpoint_id is None - assert model.post_checkpoint_id is None - - def test_config_model_inherits_base_optional_params(self): - """Test that model has optional_params from GuardrailConfigModel and defaults to None.""" - from litellm.types.proxy.guardrails.guardrail_hooks.base import ( - GuardrailConfigModel, - ) - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, - ) - - model = OvalixGuardrailConfigModel() - assert hasattr(model, "optional_params") - assert model.optional_params is None - assert isinstance(model, GuardrailConfigModel) - - def test_config_model_serialization_dump(self): - """Test model_dump produces expected keys and values.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, - ) - - model = OvalixGuardrailConfigModel( - tracker_api_base="https://tracker.test", - pre_checkpoint_id="pre-1", - ) - data = model.model_dump() - - assert "tracker_api_base" in data - assert "tracker_api_key" in data - assert "application_id" in data - assert "pre_checkpoint_id" in data - assert "post_checkpoint_id" in data - assert data["tracker_api_base"] == "https://tracker.test" - assert data["pre_checkpoint_id"] == "pre-1" - assert data["tracker_api_key"] is None - - def test_config_model_deserialization_from_dict(self): - """Test model_validate builds instance from dict.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, - ) - - payload = { - "tracker_api_base": "https://from-dict.example", - "tracker_api_key": "secret", - "application_id": "app-dict", - "pre_checkpoint_id": "pre-d", - "post_checkpoint_id": "post-d", + # Tracker allows and returns same content for the user message + allow_hello = { + "action_type": "allow", + "data_type": "TEXT", + "original_data": {"content": "hello"}, + "modified_data": {"content": "hello"}, + "alerts": [], } - model = OvalixGuardrailConfigModel.model_validate(payload) + mock_response = MagicMock() + mock_response.json.return_value = allow_hello + mock_response.raise_for_status = MagicMock() - assert model.tracker_api_base == "https://from-dict.example" - assert model.tracker_api_key == "secret" - assert model.application_id == "app-dict" - assert model.pre_checkpoint_id == "pre-d" - assert model.post_checkpoint_id == "post-d" + 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, + request_data=request_data, + input_type="request", + logging_obj=None, + ) - def test_config_model_deserialization_empty_dict(self): - """Test model_validate with empty dict yields defaults.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, + assert result.get("texts") == ["You are helpful.", "hello"] + 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 + ): + """When Tracker response has no modified_data.content, original content is used.""" + guardrail = guardrail_with_env + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "original text"}] ) + request_data = {} - model = OvalixGuardrailConfigModel.model_validate({}) - - assert model.tracker_api_base is None - assert model.tracker_api_key is None - assert model.application_id is None - assert model.pre_checkpoint_id is None - assert model.post_checkpoint_id is None - - def test_config_model_round_trip(self): - """Test model_dump then model_validate preserves data.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, - ) - - original = OvalixGuardrailConfigModel( - tracker_api_base="https://round.trip", - post_checkpoint_id="post-rt", - ) - data = original.model_dump() - restored = OvalixGuardrailConfigModel.model_validate(data) - - assert restored.tracker_api_base == original.tracker_api_base - assert restored.tracker_api_key == original.tracker_api_key - assert restored.application_id == original.application_id - assert restored.pre_checkpoint_id == original.pre_checkpoint_id - assert restored.post_checkpoint_id == original.post_checkpoint_id - - def test_config_model_has_expected_field_names(self): - """Test that model defines all expected Ovalix config field names.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, - ) - - expected = { - "tracker_api_base", - "tracker_api_key", - "application_id", - "pre_checkpoint_id", - "post_checkpoint_id", + mock_response = MagicMock() + mock_response.json.return_value = { + "action_type": "allow", + "data_type": "TEXT", + "original_data": {"content": "original text"}, + "modified_data": {}, + "alerts": [], } - assert expected.issubset(OvalixGuardrailConfigModel.model_fields.keys()) + mock_response.raise_for_status = MagicMock() - def test_config_model_field_descriptions_present(self): - """Test that Ovalix-specific fields have non-empty descriptions.""" - from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import ( - OvalixGuardrailConfigModel, + 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, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result.get("texts") == ["original text"] + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + 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( + structured_messages=[{"role": "user", "content": "hello"}] + ) + request_data = {} + + mock_response = MagicMock() + mock_response.raise_for_status.side_effect = httpx.HTTPStatusError( + "Bad Request", + request=MagicMock(), + response=MagicMock(status_code=400), ) - ovalix_fields = [ - "tracker_api_base", - "tracker_api_key", - "application_id", - "pre_checkpoint_id", - "post_checkpoint_id", - ] - for name in ovalix_fields: - assert name in OvalixGuardrailConfigModel.model_fields - desc = OvalixGuardrailConfigModel.model_fields[name].description - assert desc is not None and len(desc) > 0 + 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( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_apply_guardrail_checkpoint_error_raises_guardrail_exception(self): + """When Tracker checkpoint call fails, GuardrailRaisedException is raised.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "hello"}] + ) + request_data = {} + + with patch.object( + guardrail._async_handler, + "post", + new_callable=AsyncMock, + side_effect=httpx.ConnectError("Connection refused"), + ): + with pytest.raises(GuardrailRaisedException): + await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + @pytest.mark.asyncio + async def test_apply_guardrail_request_empty_messages_returns_inputs(self): + """When request has no messages, apply_guardrail returns inputs without calling Tracker.""" + for k, v in _ovalix_env().items(): + os.environ[k] = v + try: + guardrail = OvalixGuardrail(**_guardrail_kwargs()) + inputs = GenericGuardrailAPIInputs(structured_messages=[]) + request_data = {} + + with patch.object( + guardrail._async_handler, "post", new_callable=AsyncMock + ) as mock_post: + result = await guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + logging_obj=None, + ) + + assert result == inputs + mock_post.assert_not_called() + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + def test_get_actor_from_metadata(self): + """Actor is taken from metadata.user_api_key_user_email or user_api_key_user_id.""" + for k, v in _ovalix_env().items(): + 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({}) == "unknown" + finally: + for k in _ovalix_env(): + if k in os.environ: + del os.environ[k] + + def test_get_actor_prefers_email_over_id(self, guardrail_with_env): + """When both user_api_key_user_email and user_api_key_user_id exist, email is used.""" + guardrail = guardrail_with_env + data = { + "metadata": { + "user_api_key_user_email": "primary@test.com", + "user_api_key_user_id": "uid-99", + } + } + assert guardrail._get_actor(data) == "primary@test.com" + + def test_get_session_id_deterministic_and_includes_app_id(self, guardrail_with_env): + """Session ID is stable for same actor/day and includes application_id.""" + guardrail = guardrail_with_env + data = {"metadata": {"user_api_key_user_id": "user-1"}} + session_id_1 = guardrail._get_session_id(data) + session_id_2 = guardrail._get_session_id(data) + assert session_id_1 == session_id_2 + assert "app-1" in session_id_1