addressing PR comments

This commit is contained in:
Shalom Jamil 2026-07-23 11:06:53 +03:00
parent 8b7a8bf73b
commit 05c5669ba9
3 changed files with 323 additions and 12 deletions

View file

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

View file

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

View file

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