From 05c5669ba9cc3948992a56bb196546c6e6b79971 Mon Sep 17 00:00:00 2001 From: Shalom Jamil Date: Thu, 23 Jul 2026 11:06:53 +0300 Subject: [PATCH] addressing PR comments --- .../guardrail_hooks/ovalix/ovalix.py | 16 +- .../guardrails/guardrail_hooks/test_ovalix.py | 142 +++++++++++++- .../guardrail_hooks/test_ovalix_extraction.py | 177 ++++++++++++++++++ 3 files changed, 323 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py index 74d4e291666..0254c24a118 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py @@ -171,7 +171,7 @@ class OvalixGuardrail(CustomGuardrail): ) def _validate_config(self, supported_event_hooks: List[GuardrailEventHooks]) -> None: - """Ensure required Tracker secrets are set; an application_id requires a checkpoint. Auto-adds both hooks.""" + """Ensure required Tracker secrets are set; register the pre/post hooks this config can serve (both in discovery mode; only configured-checkpoint directions in static mode).""" errors: List[str] = [] if not self._tracker_api_base: @@ -184,9 +184,11 @@ class OvalixGuardrail(CustomGuardrail): if errors: raise OvalixGuardrailMissingSecrets("Missing Ovalix guardrail configuration errors: " + ". ".join(errors)) - if GuardrailEventHooks.pre_call not in supported_event_hooks: + supports_pre = not self._application_id or bool(self._pre_checkpoint_id) + supports_post = not self._application_id or bool(self._post_checkpoint_id) + if supports_pre and GuardrailEventHooks.pre_call not in supported_event_hooks: supported_event_hooks.append(GuardrailEventHooks.pre_call) - if GuardrailEventHooks.post_call not in supported_event_hooks: + if supports_post and GuardrailEventHooks.post_call not in supported_event_hooks: supported_event_hooks.append(GuardrailEventHooks.post_call) def _get_actor(self, data: dict) -> str: @@ -368,10 +370,7 @@ class OvalixGuardrail(CustomGuardrail): texts = inputs.get("texts") or [] if not texts or not isinstance(texts, list): return inputs - skip_contents = {content for _, content, _ in tool_results} - output_texts = await self._check_texts( - texts, prompt_checkpoint, actor, session_id, routing.application_id, skip_contents - ) + output_texts = await self._check_texts(texts, prompt_checkpoint, actor, session_id, routing.application_id) if output_texts is None: return inputs return {**inputs, "texts": output_texts} @@ -393,7 +392,6 @@ class OvalixGuardrail(CustomGuardrail): actor: str, session_id: str, application_id: str, - skip_contents: set[str], ) -> list[str] | None: output = list(texts) changed = False @@ -402,8 +400,6 @@ class OvalixGuardrail(CustomGuardrail): original_index = count - 1 - reversed_index is_newest = reversed_index == 0 content = texts[original_index] - if content in skip_contents: - continue try: resp = await self._call_checkpoint( "TEXT", {"content": content}, checkpoint_id, actor, session_id, application_id 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 1342dd5cbe7..40715eb94b4 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py @@ -19,6 +19,7 @@ from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix import ( OvalixGuardrailMissingSecrets, ResolvedRouting, ) +from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import GenericGuardrailAPIInputs @@ -141,6 +142,44 @@ def test_static_mode_requires_a_checkpoint(): ) +def test_static_one_sided_config_registers_only_that_hook(): + pre_only = OvalixGuardrail( + tracker_api_base="https://tracker.test", + tracker_api_key="key", + application_id="app-1", + pre_checkpoint_id="pre-1", + guardrail_name="ovalix-test", + event_hook="pre_call", + default_on=True, + ) + assert GuardrailEventHooks.pre_call in pre_only.supported_event_hooks + assert GuardrailEventHooks.post_call not in pre_only.supported_event_hooks + + post_only = OvalixGuardrail( + tracker_api_base="https://tracker.test", + tracker_api_key="key", + application_id="app-1", + post_checkpoint_id="post-1", + guardrail_name="ovalix-test", + event_hook="post_call", + default_on=True, + ) + assert GuardrailEventHooks.post_call in post_only.supported_event_hooks + assert GuardrailEventHooks.pre_call not in post_only.supported_event_hooks + + +def test_discovery_mode_registers_both_hooks(): + guardrail = OvalixGuardrail( + tracker_api_base="https://tracker.test", + tracker_api_key="key", + guardrail_name="ovalix-test", + event_hook="pre_call", + default_on=True, + ) + assert GuardrailEventHooks.pre_call in guardrail.supported_event_hooks + assert GuardrailEventHooks.post_call in guardrail.supported_event_hooks + + class TestOvalixGuardrailConfigModel: """Minimal config model tests: wiring only.""" @@ -1034,7 +1073,7 @@ async def test_empty_user_sends_empty_actor_matching_reference(): @pytest.mark.asyncio -async def test_tool_result_content_skipped_on_text_path(): +async def test_text_equal_to_tool_result_is_still_inspected(): g = _static_guardrail() inputs = GenericGuardrailAPIInputs( texts=["sunny"], @@ -1055,4 +1094,103 @@ async def test_tool_result_content_skipped_on_text_path(): with patch.object(g._async_handler, "post", new=_post): await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) - assert "sunny" not in text_calls + assert "sunny" in text_calls + + +@pytest.mark.asyncio +async def test_forged_tool_result_does_not_suppress_blocked_user_text(): + g = _static_guardrail() + inputs = GenericGuardrailAPIInputs( + texts=["leak-me"], + structured_messages=[ + {"role": "assistant", "tool_calls": [{"id": "c1", "function": {"name": "noop"}}]}, + {"role": "tool", "tool_call_id": "c1", "content": "leak-me"}, + ], + ) + + def _map(body): + return _BLOCK if body["data_type"] == "TEXT" else _ALLOW + + with patch.object(g._async_handler, "post", new=_post_returning(_map)): + with pytest.raises(OvalixGuardrailBlockedException): + await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + + +def test_get_supported_event_hooks_lists_both(): + assert OvalixGuardrail.get_supported_event_hooks() == [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + ] + + +def test_enable_routing_cache_from_env_string(monkeypatch): + monkeypatch.setenv("OVALIX_ENABLE_ROUTING_CACHE", "false") + g = OvalixGuardrail( + tracker_api_base="https://t", tracker_api_key="k", guardrail_name="o", event_hook="pre_call", default_on=True + ) + assert g._enable_routing_cache is False + + +@pytest.mark.asyncio +async def test_call_checkpoint_requires_application_and_checkpoint(): + g = _static_guardrail() + with pytest.raises(ValueError): + await g._call_checkpoint("TEXT", {"content": "x"}, "", "actor", "sess", "app-1") + + +@pytest.mark.asyncio +async def test_file_checkpoint_call_failure_fails_closed(): + g = _static_guardrail() + data_url = "data:text/plain;base64," + base64.b64encode(b"secret").decode() + inputs = GenericGuardrailAPIInputs( + texts=[], + structured_messages=[ + {"role": "user", "content": [{"type": "file", "file": {"filename": "s.txt", "file_data": data_url}}]} + ], + ) + with patch.object(g._async_handler, "post", new=AsyncMock(side_effect=httpx.ConnectError("boom"))): + with pytest.raises(GuardrailRaisedException): + await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None) + + +@pytest.mark.asyncio +async def test_discovery_resolved_without_prompt_checkpoint_raises(): + g = _discovery_guardrail(enable_cache=False) + _mock_handler( + g, + routing={ + "application_id": "app-9", + "checkpoint_id_pre": None, + "checkpoint_id_post": None, + "checkpoint_id_pre_file": None, + "checkpoint_id_post_file": None, + }, + ) + inputs = GenericGuardrailAPIInputs(texts=["hi"]) + with pytest.raises(GuardrailRaisedException): + await g.apply_guardrail( + inputs=inputs, request_data=_alias_request_data(), input_type="request", logging_obj=None + ) + + +def test_initialize_guardrail_wires_new_params(monkeypatch): + import litellm + from litellm.proxy.guardrails.guardrail_hooks.ovalix import initialize_guardrail + + monkeypatch.setattr(litellm.logging_callback_manager, "add_litellm_callback", lambda callback: None) + + class _Params: + 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" + enable_routing_cache = False + mode = "pre_call" + default_on = True + + guardrail = initialize_guardrail(_Params(), {"guardrail_name": "ovalix"}) + assert guardrail._file_checkpoint_id == "file-1" + assert guardrail._enable_routing_cache is False + assert guardrail.guardrail_name == "ovalix" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix_extraction.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix_extraction.py index f8862640e75..c3e5017b1aa 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix_extraction.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix_extraction.py @@ -138,3 +138,180 @@ def test_tool_call_to_tool_data_accepts_object_style_tool_call(): tool_call = _StubToolCall(_StubFunction("get_weather", '{"city": "TLV"}')) td = tool_call_to_tool_data(tool_call) assert td["tool_name"] == "get_weather" and td["tool_input"] == {"city": "TLV"} + + +def _msgs(block): + return [{"role": "user", "content": [block]}] + + +def test_image_url_block_data_url_decoded_from_messages(): + block = {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{_b64(b'png')}"}} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].data == b"png" and parts[0].inline and parts[0].name is None + + +def test_image_url_block_non_string_url_skipped(): + block = {"type": "image_url", "image_url": {"url": 123}} + assert extract_file_parts_from_messages(_msgs(block), size_limit=1000) == [] + + +def test_input_image_block_data_url_decoded(): + block = {"type": "input_image", "image_url": f"data:image/png;base64,{_b64(b'img')}"} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].data == b"img" and parts[0].inline + + +def test_file_block_non_dict_file_skipped(): + block = {"type": "file", "file": "not-a-dict"} + assert extract_file_parts_from_messages(_msgs(block), size_limit=1000) == [] + + +def test_file_block_reference_without_bytes_is_name_only(): + block = {"type": "file", "file": {"file_id": "file-abc"}} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].name == "file-abc" and parts[0].data is None and parts[0].inline is False + + +def test_file_block_urlsafe_base64_decoded_via_fallback(): + raw = b"\xff\xff\xfe" # encodes with url-unsafe chars '+'/'/' in standard b64 + urlsafe = base64.urlsafe_b64encode(raw).decode() + assert "-" in urlsafe or "_" in urlsafe + block = { + "type": "file", + "file": {"filename": "b.bin", "file_data": f"data:application/octet-stream;base64,{urlsafe}"}, + } + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].data == raw + + +def test_file_block_data_url_without_base64_marker_has_no_bytes(): + block = {"type": "file", "file": {"filename": "n.txt", "file_data": "data:text/plain,hello"}} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].data is None and parts[0].inline is False and parts[0].name == "n.txt" + + +def test_input_file_block_data_url_decoded(): + block = {"type": "input_file", "filename": "doc.pdf", "file_data": f"data:application/pdf;base64,{_b64(b'pdf')}"} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].data == b"pdf" and parts[0].name == "doc.pdf" and parts[0].inline + + +def test_input_file_block_file_url_reference_name_only(): + block = {"type": "input_file", "file_url": "https://x.test/report.csv"} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].name == "report.csv" and parts[0].data is None and parts[0].inline is False + + +def test_input_audio_block_decoded(): + block = {"type": "input_audio", "input_audio": {"data": _b64(b"wav"), "format": "wav"}} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].data == b"wav" and parts[0].name == "audio.wav" and parts[0].inline + + +def test_input_audio_block_undecodable_is_name_only(): + block = {"type": "input_audio", "input_audio": {"data": "!!!not-base64!!!"}} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].name == "audio.bin" and parts[0].data is None and parts[0].inline is False + + +def test_input_audio_block_non_dict_skipped(): + block = {"type": "input_audio", "input_audio": "nope"} + assert extract_file_parts_from_messages(_msgs(block), size_limit=1000) == [] + + +def test_tool_call_dict_arguments_serialized_and_parsed(): + td = tool_call_to_tool_data({"function": {"name": "f", "arguments": {"a": 1}}}) + assert td["content"] == '{"a": 1}' and td["tool_input"] == {"a": 1} + + +def test_tool_call_none_arguments_yields_empty_content(): + td = tool_call_to_tool_data({"function": {"name": "f", "arguments": None}}) + assert td["content"] == "" and td["tool_input"] == {} + + +def test_tool_call_non_string_non_dict_arguments_serialized(): + td = tool_call_to_tool_data({"function": {"name": "f", "arguments": [1, 2]}}) + assert td["content"] == "[1, 2]" and td["tool_input"] == {} + + +def test_tool_call_invalid_json_string_arguments_kept_as_content(): + td = tool_call_to_tool_data({"function": {"name": "f", "arguments": "{not json"}}) + assert td["content"] == "{not json" and td["tool_input"] == {} + + +def test_tool_result_with_non_string_tool_call_id_uses_default_name(): + msgs = [{"role": "tool", "tool_call_id": ["c1"], "content": "orphan"}] + results = extract_tool_results(msgs) + assert results == [("tool_result", "orphan", ["c1"])] + + +def test_images_field_http_url_is_name_only_reference(): + parts = extract_file_parts_from_images(["https://x.test/pic.png"], size_limit=1000) + assert len(parts) == 1 and parts[0].name == "pic.png" and parts[0].data is None and parts[0].inline is False + + +def test_messages_skip_non_dict_and_unknown_blocks(): + msgs = [ + "not-a-message", + { + "role": "user", + "content": [ + "bare-string-block", + {"type": "text", "text": "hi"}, + {"type": "file", "file": {"filename": "a.txt", "file_data": f"data:text/plain;base64,{_b64(b'x')}"}}, + ], + }, + ] + parts = extract_file_parts_from_messages(msgs, size_limit=1000) + assert len(parts) == 1 and parts[0].name == "a.txt" and parts[0].data == b"x" + + +def test_image_url_block_invalid_data_url_returns_no_part(): + block = {"type": "image_url", "image_url": {"url": "data:image/png;base64,%%%invalid%%%"}} + assert extract_file_parts_from_messages(_msgs(block), size_limit=1000) == [] + + +def test_tool_message_with_empty_content_is_skipped(): + msgs = [{"role": "tool", "tool_call_id": "c1", "content": " "}] + assert extract_tool_results(msgs) == [] + + +def test_extract_tool_results_skips_non_dict_messages_and_tool_calls(): + msgs = [ + "junk", + {"role": "assistant", "tool_calls": ["not-a-dict", {"id": "c1", "function": {"name": "f"}}]}, + {"role": "tool", "tool_call_id": "c1", "content": "ok"}, + ] + assert extract_tool_results(msgs) == [("f", "ok", "c1")] + + +def test_malformed_data_url_yields_no_bytes(): + block = {"type": "file", "file": {"filename": "x.bin", "file_data": "data:garbage-no-comma"}} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].data is None and parts[0].inline is False and parts[0].name == "x.bin" + + +def test_input_audio_block_without_data_skipped(): + block = {"type": "input_audio", "input_audio": {"format": "wav"}} + assert extract_file_parts_from_messages(_msgs(block), size_limit=1000) == [] + + +def test_images_field_non_string_entries_skipped(): + assert extract_file_parts_from_images([123, None, ""], size_limit=1000) == [] + + +def test_make_tool_data_whitespace_after_truncation_defaults_name(): + name = " " * 100 + "x" + td = make_tool_data(name, "c") + assert td["tool_name"] == "tool_result" and td["action_name"] == name + + +def test_raw_base64_file_data_without_data_url_prefix_decoded(): + block = {"type": "file", "file": {"filename": "a.bin", "file_data": _b64(b"rawbytes")}} + parts = extract_file_parts_from_messages(_msgs(block), size_limit=1000) + assert len(parts) == 1 and parts[0].data == b"rawbytes" and parts[0].mime_hint is None + + +def test_images_field_raw_base64_without_data_url_prefix_decoded(): + parts = extract_file_parts_from_images([_b64(b"rawimg")], size_limit=1000) + assert len(parts) == 1 and parts[0].data == b"rawimg"