mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(ovalix): scan files, tool calls, tool results via custom_application checkpoint
This commit is contained in:
parent
c023ec98b1
commit
fa0d3ea14c
2 changed files with 424 additions and 68 deletions
|
|
@ -4,13 +4,17 @@ Use Ovalix Guardrails for your LLM calls. Supports pre_call (user input) and
|
|||
post_call (model output) checkpoints with optional correction/blocking.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import datetime
|
||||
import gzip
|
||||
import hashlib
|
||||
import mimetypes
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, NamedTuple, Optional, Type
|
||||
from typing import TYPE_CHECKING, Any, List, Literal, NamedTuple, Optional, Type
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -24,6 +28,14 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix_extraction import (
|
||||
FilePart,
|
||||
extract_file_parts_from_images,
|
||||
extract_file_parts_from_messages,
|
||||
extract_tool_results,
|
||||
make_tool_data,
|
||||
tool_call_to_tool_data,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
|
@ -33,8 +45,23 @@ if TYPE_CHECKING:
|
|||
|
||||
BLOCKED_BY_OVALIX_FALLBACK_MESSAGE = "This message was blocked by Ovalix"
|
||||
BLOCKED_ACTION_TYPE = "block"
|
||||
_MODIFY_ACTION_TYPES = ("anonymize", "sanitize")
|
||||
_ROUTING_CACHE_TTL_SECONDS = 3600
|
||||
_ROUTING_CACHE_MAX_SIZE = 1000
|
||||
_DEFAULT_FILE_SIZE_LIMIT = 64 * 1024 * 1024
|
||||
_FILE_BLOCK_ESCALATION_REASON = (
|
||||
"This message was blocked by Ovalix because file content anonymization isn't possible via LiteLLM"
|
||||
)
|
||||
_TOOL_BLOCK_ESCALATION_REASON = (
|
||||
"This message was blocked by Ovalix because tool call anonymization isn't possible via LiteLLM"
|
||||
)
|
||||
_TOOL_RESULT_BLOCK_ESCALATION_REASON = (
|
||||
"This message was blocked by Ovalix because tool result anonymization isn't possible via LiteLLM"
|
||||
)
|
||||
|
||||
|
||||
def _encode_file_wire_format(raw: bytes) -> str:
|
||||
return base64.b64encode(gzip.compress(raw)).decode()
|
||||
|
||||
|
||||
class ResolvedRouting(NamedTuple):
|
||||
|
|
@ -183,36 +210,98 @@ class OvalixGuardrail(CustomGuardrail):
|
|||
|
||||
def _get_session_id(self, data: dict) -> str:
|
||||
"""Return a unique identifier for the chat/session (actor + date + application_id)."""
|
||||
actor_hash = self._get_tracker_actor_id(data)
|
||||
today = datetime.datetime.now(datetime.timezone.utc).strftime("%Y-%m-%d")
|
||||
return f"{actor_hash}_{today}_{self._application_id}"
|
||||
return self._get_session_id_for_application(data, self._application_id)
|
||||
|
||||
async def _call_checkpoint(
|
||||
self,
|
||||
content: str,
|
||||
data_type: str,
|
||||
data: dict[str, Any],
|
||||
checkpoint_id: str,
|
||||
actor: str,
|
||||
session_id: str,
|
||||
) -> Dict[str, Any]:
|
||||
application_id: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Call the Ovalix Tracker checkpoint API and return the JSON response."""
|
||||
application_id = self._application_id
|
||||
if not application_id or not checkpoint_id:
|
||||
raise ValueError("Ovalix: application_id or checkpoint_id not resolved")
|
||||
|
||||
url = f"{self._tracker_api_base}/tracking/custom_application/checkpoint"
|
||||
headers = dict(self._tracker_headers)
|
||||
payload = {
|
||||
"application_id": application_id,
|
||||
"checkpoint_id": checkpoint_id,
|
||||
"actor": actor,
|
||||
"session_id": session_id,
|
||||
"data_type": "TEXT",
|
||||
"data": {"content": content},
|
||||
"data_type": data_type,
|
||||
"data": data,
|
||||
"tool": "LiteLLM",
|
||||
}
|
||||
response = await self._async_handler.post(url, headers=headers, json=payload)
|
||||
response = await self._async_handler.post(url, headers=dict(self._tracker_headers), json=payload)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
def _verdict(self, resp: dict[str, Any]) -> tuple[str, str | None]:
|
||||
return (resp.get("action_type") or "").lower(), self._get_trackers_corrected_message(resp)
|
||||
|
||||
async def _block_reason_for_item(
|
||||
self,
|
||||
data_type: str,
|
||||
data: dict[str, Any],
|
||||
checkpoint_id: str,
|
||||
actor: str,
|
||||
session_id: str,
|
||||
application_id: str,
|
||||
escalation_reason: str,
|
||||
) -> str | None:
|
||||
try:
|
||||
resp = await self._call_checkpoint(data_type, data, checkpoint_id, actor, session_id, application_id)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Ovalix checkpoint call failed: %s", e)
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=f"Ovalix guardrail error: {e!s}",
|
||||
should_wrap_with_default_message=False,
|
||||
) from e
|
||||
action, corrected = self._verdict(resp)
|
||||
if action == BLOCKED_ACTION_TYPE:
|
||||
return corrected or BLOCKED_BY_OVALIX_FALLBACK_MESSAGE
|
||||
if action in _MODIFY_ACTION_TYPES:
|
||||
return escalation_reason
|
||||
return None
|
||||
|
||||
async def _check_items_block_only(
|
||||
self,
|
||||
items: list[tuple[str, dict[str, Any]]],
|
||||
checkpoint_id: str,
|
||||
actor: str,
|
||||
session_id: str,
|
||||
application_id: str,
|
||||
escalation_reason: str,
|
||||
) -> str | None:
|
||||
for data_type, data in items:
|
||||
reason = await self._block_reason_for_item(
|
||||
data_type, data, checkpoint_id, actor, session_id, application_id, escalation_reason
|
||||
)
|
||||
if reason is not None:
|
||||
return reason
|
||||
return None
|
||||
|
||||
async def _check_files_for_block(
|
||||
self,
|
||||
file_parts: list[FilePart],
|
||||
checkpoint_id: str,
|
||||
actor: str,
|
||||
session_id: str,
|
||||
application_id: str,
|
||||
) -> str | None:
|
||||
for part in sorted(file_parts, key=lambda p: p.message_index, reverse=True):
|
||||
data = await self._file_part_to_data(part)
|
||||
reason = await self._block_reason_for_item(
|
||||
"FILE", data, checkpoint_id, actor, session_id, application_id, _FILE_BLOCK_ESCALATION_REASON
|
||||
)
|
||||
if reason is not None:
|
||||
return reason
|
||||
return None
|
||||
|
||||
@log_guardrail_information
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
|
|
@ -221,74 +310,129 @@ class OvalixGuardrail(CustomGuardrail):
|
|||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""
|
||||
Apply Ovalix guardrail to the given inputs (request or response text).
|
||||
routing = await self._resolve_routing(request_data)
|
||||
actor = self._get_tracker_actor_id(request_data)
|
||||
session_id = self._get_session_id_for_application(request_data, routing.application_id)
|
||||
is_response = input_type == "response"
|
||||
|
||||
Used by the unified guardrail flow and the /apply_guardrail API.
|
||||
For "request", uses the pre-checkpoint; for "response", uses the post-checkpoint.
|
||||
prompt_checkpoint = routing.checkpoint_id_post if is_response else routing.checkpoint_id_pre
|
||||
file_checkpoint = (
|
||||
routing.checkpoint_id_post_file if is_response else routing.checkpoint_id_pre_file
|
||||
) or prompt_checkpoint
|
||||
if not prompt_checkpoint:
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message="Ovalix guardrail error: no checkpoint resolved for input_type",
|
||||
should_wrap_with_default_message=False,
|
||||
)
|
||||
|
||||
Args:
|
||||
inputs: Guardrail API inputs (e.g. texts to check).
|
||||
request_data: Full request payload (messages, metadata, response).
|
||||
input_type: "request" (pre_call) or "response" (post_call).
|
||||
logging_obj: Optional logging context.
|
||||
structured_messages = inputs.get("structured_messages") or []
|
||||
file_parts = (
|
||||
extract_file_parts_from_images(inputs.get("images"), size_limit=_DEFAULT_FILE_SIZE_LIMIT)
|
||||
if is_response
|
||||
else extract_file_parts_from_messages(structured_messages, size_limit=_DEFAULT_FILE_SIZE_LIMIT)
|
||||
)
|
||||
file_block = await self._check_files_for_block(
|
||||
file_parts, file_checkpoint, actor, session_id, routing.application_id
|
||||
)
|
||||
if file_block is not None:
|
||||
self._block_current_message(file_block)
|
||||
|
||||
Returns:
|
||||
Updated inputs (e.g. with replaced/corrected texts, or unchanged).
|
||||
"""
|
||||
if not self._pre_checkpoint_id and not self._post_checkpoint_id:
|
||||
return inputs
|
||||
tool_call_items = [
|
||||
("TOOL", td) for td in (tool_call_to_tool_data(tc) for tc in (inputs.get("tool_calls") or [])) if td
|
||||
]
|
||||
tool_block = await self._check_items_block_only(
|
||||
tool_call_items,
|
||||
prompt_checkpoint,
|
||||
actor,
|
||||
session_id,
|
||||
routing.application_id,
|
||||
_TOOL_BLOCK_ESCALATION_REASON,
|
||||
)
|
||||
if tool_block is not None:
|
||||
self._block_current_message(tool_block)
|
||||
|
||||
tool_results = extract_tool_results(structured_messages)
|
||||
tool_result_items = [("TOOL", make_tool_data(name, content)) for name, content, _ in tool_results]
|
||||
tool_result_block = await self._check_items_block_only(
|
||||
tool_result_items,
|
||||
prompt_checkpoint,
|
||||
actor,
|
||||
session_id,
|
||||
routing.application_id,
|
||||
_TOOL_RESULT_BLOCK_ESCALATION_REASON,
|
||||
)
|
||||
if tool_result_block is not None:
|
||||
self._block_current_message(tool_result_block)
|
||||
|
||||
tracker_actor_id = self._get_tracker_actor_id(request_data)
|
||||
session_id = self._get_session_id(request_data)
|
||||
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
|
||||
)
|
||||
if output_texts is None:
|
||||
return inputs
|
||||
return {**inputs, "texts": output_texts}
|
||||
|
||||
if input_type == "response":
|
||||
if not self._post_checkpoint_id:
|
||||
return inputs
|
||||
corrected_llm_responses = await self._generate_post_guardrail_llm_texts(
|
||||
texts, tracker_actor_id, session_id, self._post_checkpoint_id
|
||||
)
|
||||
return {**inputs, "texts": corrected_llm_responses}
|
||||
async def _file_part_to_data(self, part: FilePart) -> dict[str, Any]:
|
||||
extension = mimetypes.guess_extension(part.mime_hint) if part.mime_hint else None
|
||||
name = part.name or (f"file{extension}" if extension else "file")
|
||||
content = (
|
||||
await asyncio.get_event_loop().run_in_executor(None, _encode_file_wire_format, part.data)
|
||||
if part.data
|
||||
else None
|
||||
)
|
||||
return {"name": name, "content": content}
|
||||
|
||||
if self._pre_checkpoint_id:
|
||||
post_guardrail_texts = await self._generate_post_guardrail_llm_texts(
|
||||
texts, tracker_actor_id, session_id, self._pre_checkpoint_id
|
||||
)
|
||||
return {**inputs, "texts": post_guardrail_texts}
|
||||
return inputs
|
||||
|
||||
async def _generate_post_guardrail_llm_texts(
|
||||
self, texts: List[str], actor: str, session_id: str, checkpoint_id: str
|
||||
) -> List[str]:
|
||||
"""Generate post-guardrail LLM responses for the given LLM responses."""
|
||||
post_guardrail_texts: List[str] = []
|
||||
|
||||
is_first_response = True
|
||||
for llm_response in reversed(texts):
|
||||
async def _check_texts(
|
||||
self,
|
||||
texts: list[str],
|
||||
checkpoint_id: str,
|
||||
actor: str,
|
||||
session_id: str,
|
||||
application_id: str,
|
||||
skip_contents: set[str],
|
||||
) -> list[str] | None:
|
||||
output = list(texts)
|
||||
changed = False
|
||||
count = len(texts)
|
||||
for reversed_index in range(count):
|
||||
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(llm_response, checkpoint_id, actor, session_id)
|
||||
resp = await self._call_checkpoint(
|
||||
"TEXT", {"content": content}, checkpoint_id, actor, session_id, application_id
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Ovalix apply_guardrail checkpoint call failed: %s", e)
|
||||
verbose_proxy_logger.exception("Ovalix checkpoint call failed: %s", e)
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=f"Ovalix guardrail error: {e!s}",
|
||||
should_wrap_with_default_message=False,
|
||||
) from e
|
||||
action, corrected = self._verdict(resp)
|
||||
if action == BLOCKED_ACTION_TYPE:
|
||||
block_message = corrected or BLOCKED_BY_OVALIX_FALLBACK_MESSAGE
|
||||
if is_newest:
|
||||
self._block_current_message(block_message)
|
||||
if output[original_index] != block_message:
|
||||
changed = True
|
||||
output[original_index] = block_message
|
||||
continue
|
||||
if action in _MODIFY_ACTION_TYPES and corrected is not None and corrected != content:
|
||||
changed = True
|
||||
output[original_index] = corrected
|
||||
return output if changed else None
|
||||
|
||||
action_type = (resp.get("action_type") or "").lower()
|
||||
blocking_message = self._get_trackers_corrected_message(resp) or BLOCKED_BY_OVALIX_FALLBACK_MESSAGE
|
||||
if action_type == BLOCKED_ACTION_TYPE and is_first_response:
|
||||
self._block_current_message(blocking_message)
|
||||
elif action_type == BLOCKED_ACTION_TYPE:
|
||||
post_guardrail_texts.insert(0, blocking_message)
|
||||
else:
|
||||
corrected_text = self._get_trackers_corrected_message(resp) or llm_response
|
||||
post_guardrail_texts.insert(0, corrected_text)
|
||||
is_first_response = False
|
||||
return post_guardrail_texts
|
||||
def _get_session_id_for_application(self, data: dict, application_id: str | None) -> str:
|
||||
actor_hash = self._get_tracker_actor_id(data)
|
||||
today = datetime.datetime.now(datetime.timezone.utc).strftime("%Y-%m-%d")
|
||||
return f"{actor_hash}_{today}_{application_id}"
|
||||
|
||||
def _block_current_message(self, blocking_message: str) -> None:
|
||||
"""Raise OvalixGuardrailBlockedException with the given message (no default wrapper)."""
|
||||
|
|
@ -358,7 +502,7 @@ class OvalixGuardrail(CustomGuardrail):
|
|||
self._pre_checkpoint_id,
|
||||
self._post_checkpoint_id,
|
||||
self._file_checkpoint_id,
|
||||
None,
|
||||
self._file_checkpoint_id,
|
||||
)
|
||||
alias = self._get_key_alias(request_data)
|
||||
if not alias:
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ Unit tests for Ovalix guardrail: config resolution and apply_guardrail behavior
|
|||
with mocked Tracker service responses (allow, anonymize, block).
|
||||
"""
|
||||
|
||||
import base64
|
||||
import gzip
|
||||
import os
|
||||
from typing import Any, List
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -217,7 +219,7 @@ class TestOvalixGuardrail:
|
|||
|
||||
@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."""
|
||||
"""_call_checkpoint POSTs to tracker with application_id, checkpoint_id, actor, session_id, data, tool."""
|
||||
for k, v in _ovalix_env().items():
|
||||
os.environ[k] = v
|
||||
try:
|
||||
|
|
@ -229,10 +231,12 @@ class TestOvalixGuardrail:
|
|||
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",
|
||||
data_type="TEXT",
|
||||
data={"content": "hello"},
|
||||
checkpoint_id="pre-1",
|
||||
actor="a1b2c3d4",
|
||||
session_id="session-1",
|
||||
application_id="app-1",
|
||||
)
|
||||
|
||||
assert result == TRACKER_RESPONSE_ALLOW
|
||||
|
|
@ -246,6 +250,7 @@ class TestOvalixGuardrail:
|
|||
assert body["session_id"] == "session-1"
|
||||
assert body["data_type"] == "TEXT"
|
||||
assert body["data"] == {"content": "hello"}
|
||||
assert body["tool"] == "LiteLLM"
|
||||
finally:
|
||||
for k in _ovalix_env():
|
||||
if k in os.environ:
|
||||
|
|
@ -400,7 +405,7 @@ class TestOvalixGuardrail:
|
|||
|
||||
@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 with texts updated from Tracker."""
|
||||
"""When input_type is response and Tracker allows, apply_guardrail leaves texts unchanged (allow never rewrites)."""
|
||||
for k, v in _ovalix_env().items():
|
||||
os.environ[k] = v
|
||||
try:
|
||||
|
|
@ -424,7 +429,7 @@ class TestOvalixGuardrail:
|
|||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.get("texts") == ["how are you?"]
|
||||
assert result.get("texts") == ["Safe assistant reply"]
|
||||
assert mock_post.call_count == 1
|
||||
finally:
|
||||
for k in _ovalix_env():
|
||||
|
|
@ -717,7 +722,7 @@ async def test_static_mode_uses_config_routing():
|
|||
default_on=True,
|
||||
)
|
||||
routing = await g._resolve_routing({})
|
||||
assert routing == ResolvedRouting("app-1", "pre-1", "post-1", "file-1", None)
|
||||
assert routing == ResolvedRouting("app-1", "pre-1", "post-1", "file-1", "file-1")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -805,3 +810,210 @@ async def test_routing_cache_lru_eviction(monkeypatch):
|
|||
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
|
||||
|
||||
|
||||
_ALLOW = {"action_type": "allow", "modified_data": {"content": "x"}}
|
||||
_BLOCK = {"action_type": "block", "modified_data": {"content": "stop-reason"}}
|
||||
_ANON = {"action_type": "anonymize", "modified_data": {"content": "redacted"}}
|
||||
|
||||
|
||||
def _static_guardrail():
|
||||
return 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,
|
||||
)
|
||||
|
||||
|
||||
def _post_returning(mapping_fn):
|
||||
resp_factory = mapping_fn
|
||||
|
||||
async def _post(url, headers=None, json=None):
|
||||
r = MagicMock()
|
||||
r.json.return_value = resp_factory(json)
|
||||
r.raise_for_status = MagicMock()
|
||||
return r
|
||||
|
||||
return _post
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_block_raises():
|
||||
g = _static_guardrail()
|
||||
data_url = "data:text/plain;base64," + base64.b64encode(b"secret").decode()
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["hi"],
|
||||
structured_messages=[
|
||||
{"role": "user", "content": [{"type": "file", "file": {"filename": "s.txt", "file_data": data_url}}]}
|
||||
],
|
||||
)
|
||||
with patch.object(
|
||||
g._async_handler, "post", new=_post_returning(lambda body: _BLOCK if body["data_type"] == "FILE" else _ALLOW)
|
||||
):
|
||||
with pytest.raises(OvalixGuardrailBlockedException) as exc:
|
||||
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
|
||||
assert "stop-reason" in str(exc.value.message)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_uses_file_checkpoint_and_gzip_wire():
|
||||
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}}]}
|
||||
],
|
||||
)
|
||||
seen = {}
|
||||
|
||||
async def _post(url, headers=None, json=None):
|
||||
seen["last"] = json
|
||||
r = MagicMock()
|
||||
r.json.return_value = _ALLOW
|
||||
r.raise_for_status = MagicMock()
|
||||
return r
|
||||
|
||||
with patch.object(g._async_handler, "post", new=_post):
|
||||
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
|
||||
assert seen["last"]["data_type"] == "FILE"
|
||||
assert seen["last"]["checkpoint_id"] == "file-1"
|
||||
assert gzip.decompress(base64.b64decode(seen["last"]["data"]["content"])) == b"secret"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_side_file_uses_file_checkpoint():
|
||||
g = _static_guardrail()
|
||||
data_url = "data:image/png;base64," + base64.b64encode(b"img").decode()
|
||||
inputs = GenericGuardrailAPIInputs(texts=[], images=[data_url])
|
||||
seen = {}
|
||||
|
||||
async def _post(url, headers=None, json=None):
|
||||
seen["last"] = json
|
||||
r = MagicMock()
|
||||
r.json.return_value = _ALLOW
|
||||
r.raise_for_status = MagicMock()
|
||||
return r
|
||||
|
||||
with patch.object(g._async_handler, "post", new=_post):
|
||||
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="response", logging_obj=None)
|
||||
assert seen["last"]["data_type"] == "FILE"
|
||||
assert seen["last"]["checkpoint_id"] == "file-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_block_raises():
|
||||
g = _static_guardrail()
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=[], tool_calls=[{"id": "c1", "type": "function", "function": {"name": "exfil", "arguments": "{}"}}]
|
||||
)
|
||||
with patch.object(
|
||||
g._async_handler, "post", new=_post_returning(lambda body: _BLOCK if body["data_type"] == "TOOL" else _ALLOW)
|
||||
):
|
||||
with pytest.raises(OvalixGuardrailBlockedException):
|
||||
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="response", logging_obj=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_anonymize_escalates_to_block():
|
||||
g = _static_guardrail()
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=[], tool_calls=[{"id": "c1", "type": "function", "function": {"name": "exfil", "arguments": "{}"}}]
|
||||
)
|
||||
with patch.object(
|
||||
g._async_handler, "post", new=_post_returning(lambda body: _ANON if body["data_type"] == "TOOL" else _ALLOW)
|
||||
):
|
||||
with pytest.raises(OvalixGuardrailBlockedException) as exc:
|
||||
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="response", logging_obj=None)
|
||||
assert "tool call anonymization isn't possible" in str(exc.value.message)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_result_block_raises():
|
||||
g = _static_guardrail()
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["sunny"],
|
||||
structured_messages=[
|
||||
{"role": "assistant", "tool_calls": [{"id": "c1", "function": {"name": "get_weather"}}]},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "sunny"},
|
||||
],
|
||||
)
|
||||
with patch.object(
|
||||
g._async_handler, "post", new=_post_returning(lambda body: _BLOCK if body["data_type"] == "TOOL" else _ALLOW)
|
||||
):
|
||||
with pytest.raises(OvalixGuardrailBlockedException):
|
||||
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_newest_text_block_raises_older_anonymized():
|
||||
g = _static_guardrail()
|
||||
inputs = GenericGuardrailAPIInputs(texts=["old", "new"])
|
||||
|
||||
def _map(body):
|
||||
if body["data_type"] != "TEXT":
|
||||
return _ALLOW
|
||||
return _BLOCK if body["data"]["content"] == "new" 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)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_older_text_anonymized_returns_modified_texts():
|
||||
g = _static_guardrail()
|
||||
inputs = GenericGuardrailAPIInputs(texts=["make me anon", "safe newest"])
|
||||
|
||||
def _map(body):
|
||||
if body["data_type"] != "TEXT":
|
||||
return _ALLOW
|
||||
return (
|
||||
{"action_type": "anonymize", "modified_data": {"content": "ANON"}}
|
||||
if body["data"]["content"] == "make me anon"
|
||||
else _ALLOW
|
||||
)
|
||||
|
||||
with patch.object(g._async_handler, "post", new=_post_returning(_map)):
|
||||
result = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
|
||||
assert result["texts"] == ["ANON", "safe newest"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_all_allow_passes_through():
|
||||
g = _static_guardrail()
|
||||
inputs = GenericGuardrailAPIInputs(texts=["hi"])
|
||||
with patch.object(g._async_handler, "post", new=_post_returning(lambda body: _ALLOW)):
|
||||
result = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
|
||||
assert result["texts"] == ["hi"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_result_content_skipped_on_text_path():
|
||||
g = _static_guardrail()
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=["sunny"],
|
||||
structured_messages=[
|
||||
{"role": "assistant", "tool_calls": [{"id": "c1", "function": {"name": "get_weather"}}]},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "sunny"},
|
||||
],
|
||||
)
|
||||
text_calls = []
|
||||
|
||||
async def _post(url, headers=None, json=None):
|
||||
if json["data_type"] == "TEXT":
|
||||
text_calls.append(json["data"]["content"])
|
||||
r = MagicMock()
|
||||
r.json.return_value = _ALLOW
|
||||
r.raise_for_status = MagicMock()
|
||||
return r
|
||||
|
||||
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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue