mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
addressing PR comments
This commit is contained in:
parent
8b7a8bf73b
commit
05c5669ba9
3 changed files with 323 additions and 12 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue