fixing PR comments

This commit is contained in:
Shalom Jamil 2026-08-03 12:32:33 +03:00
parent 3ec57cedb3
commit 51855a20f5
4 changed files with 435 additions and 192 deletions

View file

@ -14,7 +14,9 @@ import os
import re
import time
from collections import OrderedDict
from typing import TYPE_CHECKING, Any, List, Literal, NamedTuple, Optional, Type
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Literal, NamedTuple
import httpx
@ -35,6 +37,7 @@ from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix_extraction import (
extract_tool_results,
make_tool_data,
tool_call_to_tool_data,
tool_result_text_indices,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
@ -51,6 +54,7 @@ _ROUTING_CACHE_TTL_SECONDS = 3600
_ROUTING_CACHE_NEGATIVE_TTL_SECONDS = 300
_ROUTING_CACHE_MAX_SIZE = 1000
_DEFAULT_FILE_SIZE_LIMIT = 64 * 1024 * 1024
_NO_METADATA: Mapping[str, Any] = MappingProxyType({})
_FILE_BLOCK_ESCALATION_REASON = (
"This message was blocked by Ovalix because file content anonymization isn't possible via LiteLLM"
)
@ -159,12 +163,16 @@ class OvalixGuardrail(CustomGuardrail):
self._validate_config(kwargs["supported_event_hooks"])
self._tracker_headers = httpx.Headers(
{
"Authorization": f"Bearer {self._tracker_api_key}",
"Content-Type": "application/json",
},
encoding="utf-8",
self._tracker_headers = dict(
httpx.Headers(
MappingProxyType(
{
"Authorization": f"Bearer {self._tracker_api_key}",
"Content-Type": "application/json",
}
),
encoding="utf-8",
)
)
self._async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
@ -180,14 +188,18 @@ class OvalixGuardrail(CustomGuardrail):
def _validate_config(self, supported_event_hooks: list[GuardrailEventHooks]) -> None:
"""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:
errors.append("Tracker API base, set OVALIX_TRACKER_API_BASE or pass tracker_api_base")
if not self._tracker_api_key:
errors.append("Tracker API key, set OVALIX_TRACKER_API_KEY or pass tracker_api_key")
if self._application_id and not self._pre_checkpoint_id and not self._post_checkpoint_id:
errors.append("With application_id set, provide OVALIX_PRE_CHECKPOINT_ID and/or OVALIX_POST_CHECKPOINT_ID")
errors = tuple(
message
for present, message in (
(not self._tracker_api_base, "Tracker API base, set OVALIX_TRACKER_API_BASE or pass tracker_api_base"),
(not self._tracker_api_key, "Tracker API key, set OVALIX_TRACKER_API_KEY or pass tracker_api_key"),
(
bool(self._application_id) and not self._pre_checkpoint_id and not self._post_checkpoint_id,
"With application_id set, provide OVALIX_PRE_CHECKPOINT_ID and/or OVALIX_POST_CHECKPOINT_ID",
),
)
if present
)
if errors:
raise OvalixGuardrailMissingSecrets("Missing Ovalix guardrail configuration errors: " + ". ".join(errors))
@ -199,16 +211,16 @@ class OvalixGuardrail(CustomGuardrail):
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:
def _get_actor(self, data: Mapping[str, Any]) -> 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 {}
metadata = data.get("metadata") or data.get("litellm_metadata") or _NO_METADATA
if metadata.get("user_api_key_user_email"):
return metadata["user_api_key_user_email"]
if metadata.get("user_api_key_user_id"):
return metadata["user_api_key_user_id"]
return ""
def _get_tracker_actor_id(self, data: dict) -> str:
def _get_tracker_actor_id(self, data: Mapping[str, Any]) -> str:
"""Normalize the actor string into a short, stable id for Tracker API payloads."""
# NOTE: this hash is purely for normalization — it collapses an arbitrary actor
# string (email, user id, or empty) into a compact, fixed-length, consistent
@ -218,24 +230,28 @@ class OvalixGuardrail(CustomGuardrail):
normalized_actor_id = hashlib.sha256(actor_id).hexdigest()[:8]
return normalized_actor_id
def _get_session_id(self, data: dict) -> str:
def _get_session_id(self, data: Mapping[str, Any]) -> str:
"""Return a unique identifier for the chat/session (actor + date + application_id)."""
return self._get_session_id_for_application(data, self._application_id)
async def _call_checkpoint(
self,
data_type: str,
data: dict[str, Any],
data: Mapping[str, Any],
checkpoint_id: str,
actor: str,
session_id: str,
application_id: str,
) -> dict[str, Any]:
) -> Mapping[str, Any]:
"""Call the Ovalix Tracker checkpoint API and return the JSON response."""
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"
url = (
f"{self._tracker_api_base}/tracking/litellm/file_checkpoint"
if data_type == "FILE"
else f"{self._tracker_api_base}/tracking/custom_application/checkpoint"
)
payload = {
"application_id": application_id,
"checkpoint_id": checkpoint_id,
@ -245,17 +261,17 @@ class OvalixGuardrail(CustomGuardrail):
"data": data,
"tool": "LiteLLM",
}
response = await self._async_handler.post(url, headers=dict(self._tracker_headers), json=payload)
response = await self._async_handler.post(url, headers=self._tracker_headers, json=payload)
response.raise_for_status()
return response.json()
def _verdict(self, resp: dict[str, Any]) -> tuple[str, str | None]:
def _verdict(self, resp: Mapping[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],
data: Mapping[str, Any],
checkpoint_id: str,
actor: str,
session_id: str,
@ -280,7 +296,7 @@ class OvalixGuardrail(CustomGuardrail):
async def _check_items_block_only(
self,
items: list[tuple[str, dict[str, Any]]],
items: Sequence[tuple[str, Mapping[str, Any]]],
checkpoint_id: str,
actor: str,
session_id: str,
@ -297,7 +313,7 @@ class OvalixGuardrail(CustomGuardrail):
async def _check_files_for_block(
self,
file_parts: list[FilePart],
file_parts: Sequence[FilePart],
checkpoint_id: str,
actor: str,
session_id: str,
@ -316,7 +332,7 @@ class OvalixGuardrail(CustomGuardrail):
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
request_data: Mapping[str, Any],
input_type: Literal["request", "response"],
logging_obj: Any | None = None,
) -> GenericGuardrailAPIInputs:
@ -338,7 +354,7 @@ class OvalixGuardrail(CustomGuardrail):
should_wrap_with_default_message=False,
)
structured_messages = inputs.get("structured_messages") or []
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
@ -350,9 +366,9 @@ class OvalixGuardrail(CustomGuardrail):
if file_block is not None:
self._block_current_message(file_block)
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_call_items = tuple(
("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,
@ -365,7 +381,7 @@ class OvalixGuardrail(CustomGuardrail):
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_items = tuple(("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,
@ -377,15 +393,22 @@ class OvalixGuardrail(CustomGuardrail):
if tool_result_block is not None:
self._block_current_message(tool_result_block)
texts = inputs.get("texts") or []
texts = inputs.get("texts") or ()
if not texts or not isinstance(texts, list):
return inputs
output_texts = await self._check_texts(texts, prompt_checkpoint, actor, session_id, routing.application_id)
output_texts = await self._check_texts(
texts,
prompt_checkpoint,
actor,
session_id,
routing.application_id,
tool_result_text_indices(structured_messages, texts),
)
if output_texts is None:
return inputs
return {**inputs, "texts": output_texts}
async def _file_part_to_data(self, part: FilePart) -> dict[str, Any]:
async def _file_part_to_data(self, part: FilePart) -> Mapping[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 = (
@ -397,17 +420,20 @@ class OvalixGuardrail(CustomGuardrail):
async def _check_texts(
self,
texts: list[str],
texts: Sequence[str],
checkpoint_id: str,
actor: str,
session_id: str,
application_id: str,
skip_indices: frozenset[int],
) -> list[str] | None:
output = list(texts)
changed = False
count = len(texts)
for reversed_index in range(count):
original_index = count - 1 - reversed_index
if original_index in skip_indices:
continue
is_newest = reversed_index == 0
content = texts[original_index]
try:
@ -435,7 +461,7 @@ class OvalixGuardrail(CustomGuardrail):
output[original_index] = corrected
return output if changed else None
def _get_session_id_for_application(self, data: dict, application_id: str | None) -> str:
def _get_session_id_for_application(self, data: Mapping[str, Any], 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}"
@ -448,23 +474,29 @@ class OvalixGuardrail(CustomGuardrail):
should_wrap_with_default_message=False,
)
def _get_trackers_corrected_message(self, resp: dict) -> str | None:
def _get_trackers_corrected_message(self, resp: Mapping[str, Any]) -> str | None:
"""Extract corrected/blocking message content from Tracker checkpoint response."""
modified = resp.get("modified_data")
if isinstance(modified, dict) and "content" in modified:
return modified["content"]
return None
def _get_key_alias(self, request_data: dict) -> str | None:
metadata = {**(request_data.get("metadata") or {}), **(request_data.get("litellm_metadata") or {})}
return metadata.get("user_api_key_alias") or metadata.get("user_api_key_key_alias")
def _get_key_alias(self, request_data: Mapping[str, Any]) -> str | None:
litellm_metadata = request_data.get("litellm_metadata") or _NO_METADATA
metadata = request_data.get("metadata") or _NO_METADATA
def _merged(key: str) -> object:
return litellm_metadata.get(key) if key in litellm_metadata else metadata.get(key)
alias = _merged("user_api_key_alias") or _merged("user_api_key_key_alias")
return alias if isinstance(alias, str) else None
async def _get_app_name_regex(self) -> re.Pattern[str]:
if self._app_name_regex is not None:
return self._app_name_regex
url = f"{self._tracker_api_base}/tracking/custom_application/litellm_app_name_regex"
url = f"{self._tracker_api_base}/tracking/litellm/app_name_regex"
try:
response = await self._async_handler.get(url, headers=dict(self._tracker_headers))
response = await self._async_handler.get(url, headers=self._tracker_headers)
response.raise_for_status()
compiled = re.compile(response.json()["regex"])
except Exception as e:
@ -512,7 +544,6 @@ class OvalixGuardrail(CustomGuardrail):
verbose_proxy_logger.warning(
"Ovalix guardrail passing the call through unguarded (fail_if_no_application=false): %s", reason
)
return None
def _routing_error(self, error: Exception) -> GuardrailRaisedException:
verbose_proxy_logger.exception("Ovalix routing resolution failed: %s", error)
@ -522,7 +553,7 @@ class OvalixGuardrail(CustomGuardrail):
should_wrap_with_default_message=False,
)
async def _resolve_routing(self, request_data: dict) -> ResolvedRouting | None:
async def _resolve_routing(self, request_data: Mapping[str, Any]) -> ResolvedRouting | None:
if self._application_id:
return ResolvedRouting(
self._application_id,
@ -550,10 +581,10 @@ class OvalixGuardrail(CustomGuardrail):
return routing
async def _resolve_via_tracker(self, application_name: str) -> ResolvedRouting | None:
url = f"{self._tracker_api_base}/tracking/custom_application/resolve_litellm_application"
url = f"{self._tracker_api_base}/tracking/litellm/resolve_application"
try:
response = await self._async_handler.post(
url, headers=dict(self._tracker_headers), json={"application_name": application_name}
url, headers=self._tracker_headers, json={"application_name": application_name}
)
response.raise_for_status()
body = response.json()

View file

@ -2,12 +2,14 @@ import base64
import json
import posixpath
import re
from collections.abc import Callable
from typing import Any, NamedTuple
from collections.abc import Callable, Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import NamedTuple
from urllib.parse import unquote, urlparse
_TOOL_NAME_MAX_LENGTH = 100
_DEFAULT_TOOL_RESULT_NAME = "tool_result"
_NO_TOOL_INPUT: Mapping[str, object] = MappingProxyType({})
_DATA_URL_RE = re.compile(r"^data:(?P<mime>[^;,]+)?(?P<params>(?:;[^;,]+)*?)(?P<b64>;base64)?,", re.IGNORECASE)
_URLSAFE_TO_STANDARD_B64 = str.maketrans("-_", "+/")
@ -63,7 +65,7 @@ def _name_from_url(url: str) -> str | None:
return None
def _part_from_file_block(block: dict[str, Any], size_limit: int | None, message_index: int) -> FilePart | None:
def _part_from_file_block(block: Mapping[str, object], size_limit: int | None, message_index: int) -> FilePart | None:
file_obj = block.get("file")
if not isinstance(file_obj, dict):
return None
@ -77,7 +79,9 @@ def _part_from_file_block(block: dict[str, Any], size_limit: int | None, message
return FilePart(name, None, None, False, False, message_index)
def _part_from_image_url_block(block: dict[str, Any], size_limit: int | None, message_index: int) -> FilePart | None:
def _part_from_image_url_block(
block: Mapping[str, object], size_limit: int | None, message_index: int
) -> FilePart | None:
image_url = block.get("image_url")
url = image_url.get("url") if isinstance(image_url, dict) else image_url
if not isinstance(url, str) or not url:
@ -91,7 +95,9 @@ def _part_from_image_url_block(block: dict[str, Any], size_limit: int | None, me
return FilePart(_name_from_url(url), None, None, False, False, message_index)
def _part_from_input_file_block(block: dict[str, Any], size_limit: int | None, message_index: int) -> FilePart | None:
def _part_from_input_file_block(
block: Mapping[str, object], size_limit: int | None, message_index: int
) -> FilePart | None:
name = block.get("filename") or block.get("file_id") or None
file_data = block.get("file_data")
if isinstance(file_data, str) and file_data:
@ -105,7 +111,9 @@ def _part_from_input_file_block(block: dict[str, Any], size_limit: int | None, m
return FilePart(name, None, None, False, False, message_index)
def _part_from_input_audio_block(block: dict[str, Any], size_limit: int | None, message_index: int) -> FilePart | None:
def _part_from_input_audio_block(
block: Mapping[str, object], size_limit: int | None, message_index: int
) -> FilePart | None:
audio = block.get("input_audio")
if not isinstance(audio, dict):
return None
@ -119,149 +127,215 @@ def _part_from_input_audio_block(block: dict[str, Any], size_limit: int | None,
return FilePart(name, data, None, True, oversize, message_index)
_BLOCK_PARSERS: dict[str, Callable[[dict[str, Any], int | None, int], FilePart | None]] = {
"file": _part_from_file_block,
"image_url": _part_from_image_url_block,
"input_image": _part_from_image_url_block,
"input_file": _part_from_input_file_block,
"input_audio": _part_from_input_audio_block,
}
_BLOCK_PARSERS: Mapping[str, Callable[[Mapping[str, object], int | None, int], FilePart | None]] = MappingProxyType(
{
"file": _part_from_file_block,
"image_url": _part_from_image_url_block,
"input_image": _part_from_image_url_block,
"input_file": _part_from_input_file_block,
"input_audio": _part_from_input_audio_block,
}
)
def _file_parts_of_message(
message: Mapping[str, object], size_limit: int | None, message_index: int
) -> Iterator[FilePart]:
content = message.get("content")
if not isinstance(content, list):
return
for block in content:
if not isinstance(block, Mapping):
continue
block_type = block.get("type")
if not isinstance(block_type, str):
continue
parser = _BLOCK_PARSERS.get(block_type)
if parser is None:
continue
try:
part = parser(block, size_limit, message_index)
except (TypeError, ValueError, AttributeError, KeyError):
continue
if part is not None and (part.inline or part.name):
yield part
def extract_file_parts_from_messages(
structured_messages: list[dict[str, Any]] | None, size_limit: int | None = None
) -> list[FilePart]:
parts: list[FilePart] = []
for message_index, message in enumerate(structured_messages or []):
if not isinstance(message, dict):
continue
content = message.get("content")
if not isinstance(content, list):
continue
for block in content:
if not isinstance(block, dict):
continue
block_type = block.get("type")
if not isinstance(block_type, str):
continue
parser = _BLOCK_PARSERS.get(block_type)
if parser is None:
continue
try:
part = parser(block, size_limit, message_index)
except (TypeError, ValueError, AttributeError, KeyError):
continue
if part is not None and (part.inline or part.name):
parts.append(part)
return parts
structured_messages: Sequence[Mapping[str, object]] | None, size_limit: int | None = None
) -> tuple[FilePart, ...]:
return tuple(
part
for message_index, message in enumerate(structured_messages or ())
if isinstance(message, Mapping)
for part in _file_parts_of_message(message, size_limit, message_index)
)
def extract_file_parts_from_images(images: list[str] | None, size_limit: int | None = None) -> list[FilePart]:
parts: list[FilePart] = []
for index, value in enumerate(images or []):
if not isinstance(value, str) or not value:
continue
if value.startswith(("http://", "https://")):
name = _name_from_url(value)
if name:
parts.append(FilePart(name, None, None, False, False, index))
continue
mime_hint, payload = _split_data_url(value)
data, oversize = _decode_base64_with_limit(payload, size_limit) if payload else (None, False)
if data is not None or oversize:
parts.append(FilePart(None, data, mime_hint, True, oversize, index))
return parts
def _file_part_of_image(value: str, size_limit: int | None, index: int) -> FilePart | None:
if value.startswith(("http://", "https://")):
name = _name_from_url(value)
return FilePart(name, None, None, False, False, index) if name else None
mime_hint, payload = _split_data_url(value)
data, oversize = _decode_base64_with_limit(payload, size_limit) if payload else (None, False)
if data is None and not oversize:
return None
return FilePart(None, data, mime_hint, True, oversize, index)
def make_tool_data(name: str, content: str | None, tool_input: dict[str, Any] | None = None) -> dict[str, Any]:
def extract_file_parts_from_images(images: Sequence[str] | None, size_limit: int | None = None) -> tuple[FilePart, ...]:
candidates = (
_file_part_of_image(value, size_limit, index)
for index, value in enumerate(images or ())
if isinstance(value, str) and value
)
return tuple(part for part in candidates if part is not None)
def make_tool_data(
name: str, content: str | None, tool_input: Mapping[str, object] | None = None
) -> Mapping[str, object]:
action_name = str(name) if str(name).strip() else _DEFAULT_TOOL_RESULT_NAME
tool_name = action_name[:_TOOL_NAME_MAX_LENGTH]
if not tool_name.strip():
tool_name = _DEFAULT_TOOL_RESULT_NAME
return {"content": content, "tool_name": tool_name, "action_name": action_name, "tool_input": tool_input or {}}
return {
"content": content,
"tool_name": tool_name,
"action_name": action_name,
"tool_input": dict(tool_input or ()),
}
def _tool_call_field(tool_call: Any, key: str) -> Any:
def _tool_call_field(tool_call: object, key: str) -> object:
if isinstance(tool_call, dict):
return tool_call.get(key)
return getattr(tool_call, key, None)
def tool_call_to_tool_data(tool_call: Any) -> dict[str, Any] | None:
def _json_or_str(value: object) -> str:
try:
return json.dumps(value)
except (TypeError, ValueError):
return str(value)
def _parsed_tool_input(raw_arguments: str) -> Mapping[str, object]:
if not raw_arguments:
return _NO_TOOL_INPUT
try:
parsed = json.loads(raw_arguments)
except (ValueError, TypeError):
return _NO_TOOL_INPUT
return parsed if isinstance(parsed, dict) else _NO_TOOL_INPUT
def _tool_content_and_input(raw_arguments: object) -> tuple[str, Mapping[str, object]]:
if isinstance(raw_arguments, str):
return raw_arguments, _parsed_tool_input(raw_arguments)
if raw_arguments is None:
return "", _NO_TOOL_INPUT
if isinstance(raw_arguments, dict):
return _json_or_str(raw_arguments), raw_arguments
return _json_or_str(raw_arguments), _NO_TOOL_INPUT
def tool_call_to_tool_data(tool_call: object) -> Mapping[str, object] | None:
function = _tool_call_field(tool_call, "function")
name = function.get("name") if isinstance(function, dict) else getattr(function, "name", None)
if not name or not str(name).strip():
return None
raw_arguments = function.get("arguments") if isinstance(function, dict) else getattr(function, "arguments", None)
tool_input: dict[str, Any] = {}
if isinstance(raw_arguments, str):
content = raw_arguments
if content:
try:
parsed = json.loads(content)
if isinstance(parsed, dict):
tool_input = parsed
except (ValueError, TypeError):
tool_input = {}
elif raw_arguments is None:
content = ""
elif isinstance(raw_arguments, dict):
try:
content = json.dumps(raw_arguments)
except (TypeError, ValueError):
content = str(raw_arguments)
tool_input = raw_arguments
else:
try:
content = json.dumps(raw_arguments)
except (TypeError, ValueError):
content = str(raw_arguments)
content, tool_input = _tool_content_and_input(raw_arguments)
return make_tool_data(name, content, tool_input)
def _extract_tool_content(content: Any) -> str | None:
def _tool_content_blocks(content: Sequence[object]) -> Iterator[str]:
for block in content:
if isinstance(block, Mapping):
text = block.get("text")
if isinstance(text, str) and text:
yield text
elif isinstance(block, str):
yield block
def _extract_tool_content(content: object) -> str | None:
if isinstance(content, list):
parts: list[str] = []
for block in content:
if isinstance(block, dict):
text = block.get("text")
if isinstance(text, str) and text:
parts.append(text)
elif isinstance(block, str):
parts.append(block)
content = "\n".join(parts)
content = "\n".join(_tool_content_blocks(content))
elif isinstance(content, dict):
try:
content = json.dumps(content)
except (TypeError, ValueError):
content = str(content)
content = _json_or_str(content)
if not isinstance(content, str) or not content.strip():
return None
return content
def extract_tool_results(structured_messages: list[dict[str, Any]] | None) -> list[tuple[str, str, str | None]]:
id_to_name: dict[str, str] = {}
results: list[tuple[str, str, str | None]] = []
for message in structured_messages or []:
if not isinstance(message, dict):
def _declared_names_for_call(message: Mapping[str, object], call_id: str) -> Iterator[str]:
for tool_call in message.get("tool_calls") or ():
if not isinstance(tool_call, Mapping) or tool_call.get("id") != call_id:
continue
role = message.get("role")
if role == "assistant":
for tool_call in message.get("tool_calls") or []:
if not isinstance(tool_call, dict):
continue
call_id = tool_call.get("id")
function = tool_call.get("function")
name = function.get("name") if isinstance(function, dict) else None
if isinstance(call_id, str) and call_id and name and str(name).strip():
id_to_name[call_id] = name
elif role == "tool":
function = tool_call.get("function")
name = function.get("name") if isinstance(function, Mapping) else None
if name and str(name).strip():
yield name
def _resolve_tool_name(messages: Sequence[Mapping[str, object]], tool_index: int, tool_call_id: object) -> str:
if not isinstance(tool_call_id, str) or not tool_call_id:
return _DEFAULT_TOOL_RESULT_NAME
declared = tuple(
name
for message in messages[:tool_index]
if isinstance(message, Mapping) and message.get("role") == "assistant"
for name in _declared_names_for_call(message, tool_call_id)
)
return declared[-1] if declared else _DEFAULT_TOOL_RESULT_NAME
def extract_tool_results(
structured_messages: Sequence[Mapping[str, object]] | None,
) -> tuple[tuple[str, str, str | None], ...]:
messages = tuple(structured_messages or ())
def _results() -> Iterator[tuple[str, str, str | None]]:
for index, message in enumerate(messages):
if not isinstance(message, Mapping) or message.get("role") != "tool":
continue
content = _extract_tool_content(message.get("content"))
if content is None:
continue
tool_call_id = message.get("tool_call_id")
resolved_name = id_to_name.get(tool_call_id) if isinstance(tool_call_id, str) else None
name = resolved_name or _DEFAULT_TOOL_RESULT_NAME
results.append((name, content, tool_call_id))
return results
yield _resolve_tool_name(messages, index, tool_call_id), content, tool_call_id
return tuple(_results())
def _message_text_origins(structured_messages: Sequence[Mapping[str, object]] | None) -> Iterator[tuple[str, bool]]:
for message in structured_messages or ():
if not isinstance(message, Mapping):
continue
content = message.get("content")
from_tool_result = message.get("role") == "tool" and _extract_tool_content(content) is not None
if isinstance(content, str):
yield content, from_tool_result
elif isinstance(content, list):
for block in content:
if isinstance(block, Mapping) and block.get("text") is not None:
yield block["text"], from_tool_result
def tool_result_text_indices(
structured_messages: Sequence[Mapping[str, object]] | None, texts: Sequence[str]
) -> frozenset[int]:
"""Positions in ``texts`` that hold content already submitted under the TOOL policy.
The chat-completions guardrail flow builds ``texts`` and ``structured_messages`` from the
same message list, so tool-role content lands in both and would otherwise be checked twice.
Other surfaces (e.g. Anthropic messages) build ``texts`` from a differently shaped payload,
so the mapping is only trusted when replaying it reproduces ``texts`` exactly; anything else
falls back to checking every text.
"""
origins = tuple(_message_text_origins(structured_messages))
if tuple(text for text, _ in origins) != tuple(texts):
return frozenset()
return frozenset(index for index, (_, from_tool_result) in enumerate(origins) if from_tool_result)

View file

@ -5,6 +5,7 @@ with mocked Tracker service responses (allow, anonymize, block).
import base64
import gzip
import json as json_lib
import os
from typing import Any, List
from unittest.mock import AsyncMock, MagicMock, patch
@ -770,9 +771,9 @@ async def test_discovery_extracts_name_and_resolves():
mock_get, mock_post = _mock_handler(g)
routing = await g._resolve_routing(_alias_request_data("[Weather App] prod"))
assert routing.application_id == "app-9"
assert mock_post.call_args.args[0].endswith("/tracking/custom_application/resolve_litellm_application")
assert mock_post.call_args.args[0].endswith("/tracking/litellm/resolve_application")
assert mock_post.call_args.kwargs["json"] == {"application_name": "Weather App"}
assert mock_get.call_args.args[0].endswith("/tracking/custom_application/litellm_app_name_regex")
assert mock_get.call_args.args[0].endswith("/tracking/litellm/app_name_regex")
@pytest.mark.asyncio
@ -946,6 +947,23 @@ async def test_response_side_file_uses_file_checkpoint():
assert seen["last"]["checkpoint_id"] == "file-1"
@pytest.mark.asyncio
async def test_file_checkpoint_call_routes_to_litellm_file_endpoint():
g = _static_guardrail()
seen = {}
async def _post(url, headers=None, json=None):
seen["url"] = url
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._call_checkpoint("FILE", {"name": "f.txt", "content": "x"}, "file-1", "a", "s", "app-1")
assert seen["url"] == "https://t/tracking/litellm/file_checkpoint"
@pytest.mark.asyncio
async def test_tool_call_block_raises():
g = _static_guardrail()
@ -1072,8 +1090,70 @@ async def test_empty_user_sends_empty_actor_matching_reference():
assert seen["last"]["actor"] == ""
def _recording_post(mapping_fn=lambda body: _ALLOW):
calls = []
async def _post(url, headers=None, json=None):
calls.append((json["data_type"], json["data"].get("content")))
r = MagicMock()
r.json.return_value = mapping_fn(json)
r.raise_for_status = MagicMock()
return r
return _post, calls
@pytest.mark.asyncio
async def test_text_equal_to_tool_result_is_still_inspected():
async def test_every_checkpoint_payload_is_json_serializable():
"""httpx json-encodes the checkpoint body, so a non-dict mapping in it would 500 at runtime."""
g = _static_guardrail()
data_url = "data:text/plain;base64," + base64.b64encode(b"secret").decode()
inputs = GenericGuardrailAPIInputs(
texts=["hello", "sunny"],
tool_calls=[{"id": "c1", "type": "function", "function": {"name": "noop", "arguments": None}}],
structured_messages=[
{"role": "user", "content": [{"type": "file", "file": {"filename": "s.txt", "file_data": data_url}}]},
{"role": "assistant", "tool_calls": [{"id": "c2", "function": {"name": "get_weather"}}]},
{"role": "tool", "tool_call_id": "c2", "content": "sunny"},
],
)
encoded = []
async def _post(url, headers=None, json=None):
encoded.append(json_lib.dumps(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)
payloads = [json_lib.loads(body) for body in encoded]
assert sorted(p["data_type"] for p in payloads) == ["FILE", "TEXT", "TEXT", "TOOL", "TOOL"]
assert all(p["data"]["tool_input"] == {} for p in payloads if p["data_type"] == "TOOL")
@pytest.mark.asyncio
async def test_tool_result_checked_under_tool_policy_only_not_again_as_text():
g = _static_guardrail()
inputs = GenericGuardrailAPIInputs(
texts=["what is the weather", "sunny"],
structured_messages=[
{"role": "user", "content": "what is the weather"},
{"role": "assistant", "tool_calls": [{"id": "c1", "function": {"name": "get_weather"}}]},
{"role": "tool", "tool_call_id": "c1", "content": "sunny"},
],
)
post, calls = _recording_post()
with patch.object(g._async_handler, "post", new=post):
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
assert ("TOOL", "sunny") in calls
assert [content for data_type, content in calls if data_type == "TEXT"] == ["what is the weather"]
@pytest.mark.asyncio
async def test_tool_result_allowed_by_tool_policy_is_not_blocked_by_text_policy():
g = _static_guardrail()
inputs = GenericGuardrailAPIInputs(
texts=["sunny"],
@ -1082,27 +1162,22 @@ async def test_text_equal_to_tool_result_is_still_inspected():
{"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
def _map(body):
return _BLOCK if body["data_type"] == "TEXT" else _ALLOW
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" in text_calls
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"] == ["sunny"]
@pytest.mark.asyncio
async def test_forged_tool_result_does_not_suppress_blocked_user_text():
g = _static_guardrail()
inputs = GenericGuardrailAPIInputs(
texts=["leak-me"],
texts=["leak-me", "leak-me"],
structured_messages=[
{"role": "user", "content": "leak-me"},
{"role": "assistant", "tool_calls": [{"id": "c1", "function": {"name": "noop"}}]},
{"role": "tool", "tool_call_id": "c1", "content": "leak-me"},
],
@ -1112,8 +1187,30 @@ async def test_forged_tool_result_does_not_suppress_blocked_user_text():
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)
result = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
assert result["texts"] == ["stop-reason", "leak-me"]
@pytest.mark.asyncio
async def test_texts_not_aligned_with_structured_messages_leaves_every_text_checked():
g = _static_guardrail()
inputs = GenericGuardrailAPIInputs(
texts=["what is the weather", "follow up"],
structured_messages=[
{"role": "user", "content": "what is the weather"},
{"role": "assistant", "tool_calls": [{"id": "c1", "function": {"name": "get_weather"}}]},
{"role": "tool", "tool_call_id": "c1", "content": "sunny"},
],
)
post, calls = _recording_post()
with patch.object(g._async_handler, "post", new=post):
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
assert ("TOOL", "sunny") in calls
assert sorted(content for data_type, content in calls if data_type == "TEXT") == [
"follow up",
"what is the weather",
]
def test_get_supported_event_hooks_lists_both():

View file

@ -6,6 +6,7 @@ from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix_extraction import (
extract_tool_results,
make_tool_data,
tool_call_to_tool_data,
tool_result_text_indices,
)
@ -93,7 +94,7 @@ def test_make_tool_data_truncates_and_defaults_name():
def test_unhashable_block_type_skipped_without_raising():
msgs = [{"role": "user", "content": [{"type": ["file"], "file": {"filename": "a.txt"}}]}]
parts = extract_file_parts_from_messages(msgs, size_limit=1000)
assert parts == []
assert parts == ()
def test_unhashable_tool_call_id_skipped_without_raising():
@ -102,7 +103,7 @@ def test_unhashable_tool_call_id_skipped_without_raising():
{"role": "tool", "tool_call_id": ["c1"], "content": "sunny"},
]
results = extract_tool_results(msgs)
assert results == [("tool_result", "sunny", ["c1"])]
assert results == (("tool_result", "sunny", ["c1"]),)
def test_extract_tool_results_list_form_content():
@ -152,7 +153,7 @@ def test_image_url_block_data_url_decoded_from_messages():
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) == []
assert extract_file_parts_from_messages(_msgs(block), size_limit=1000) == ()
def test_input_image_block_data_url_decoded():
@ -163,7 +164,7 @@ def test_input_image_block_data_url_decoded():
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) == []
assert extract_file_parts_from_messages(_msgs(block), size_limit=1000) == ()
def test_file_block_reference_without_bytes_is_name_only():
@ -216,7 +217,7 @@ def test_input_audio_block_undecodable_is_name_only():
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) == []
assert extract_file_parts_from_messages(_msgs(block), size_limit=1000) == ()
def test_tool_call_dict_arguments_serialized_and_parsed():
@ -242,7 +243,7 @@ def test_tool_call_invalid_json_string_arguments_kept_as_content():
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"])]
assert results == (("tool_result", "orphan", ["c1"]),)
def test_images_field_http_url_is_name_only_reference():
@ -268,12 +269,12 @@ def test_messages_skip_non_dict_and_unknown_blocks():
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) == []
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) == []
assert extract_tool_results(msgs) == ()
def test_extract_tool_results_skips_non_dict_messages_and_tool_calls():
@ -282,7 +283,7 @@ def test_extract_tool_results_skips_non_dict_messages_and_tool_calls():
{"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")]
assert extract_tool_results(msgs) == (("f", "ok", "c1"),)
def test_malformed_data_url_yields_no_bytes():
@ -293,11 +294,11 @@ def test_malformed_data_url_yields_no_bytes():
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) == []
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) == []
assert extract_file_parts_from_images([123, None, ""], size_limit=1000) == ()
def test_make_tool_data_whitespace_after_truncation_defaults_name():
@ -337,14 +338,14 @@ def test_file_oversize_detected_after_decode_when_estimate_passes():
def test_image_http_url_that_urlparse_rejects_is_dropped():
parts = extract_file_parts_from_images(["http://["], size_limit=1000)
assert parts == []
assert parts == ()
def test_message_content_not_a_list_is_skipped():
parts = extract_file_parts_from_messages(
[{"role": "user", "content": "just a plain string prompt"}], size_limit=1000
)
assert parts == []
assert parts == ()
def test_tool_call_dict_arguments_non_serializable_falls_back_to_str():
@ -360,3 +361,43 @@ def test_tool_call_non_serializable_other_arguments_falls_back_to_str():
def test_tool_result_dict_content_non_serializable_falls_back_to_str():
results = extract_tool_results([{"role": "tool", "tool_call_id": "c1", "content": {"x": {1, 2}}}])
assert len(results) == 1 and isinstance(results[0][1], str) and results[0][1].strip()
def test_tool_result_text_indices_marks_only_tool_role_positions():
messages = [
{"role": "user", "content": "ask"},
{"role": "assistant", "tool_calls": [{"id": "c1", "function": {"name": "f"}}]},
{"role": "tool", "tool_call_id": "c1", "content": "result"},
]
assert tool_result_text_indices(messages, ["ask", "result"]) == frozenset({1})
def test_tool_result_text_indices_covers_every_text_block_of_a_tool_message():
messages = [
{"role": "user", "content": "ask"},
{
"role": "tool",
"tool_call_id": "c1",
"content": [{"type": "text", "text": "a"}, {"type": "text", "text": "b"}],
},
]
assert tool_result_text_indices(messages, ["ask", "a", "b"]) == frozenset({1, 2})
def test_tool_result_text_indices_empty_when_texts_do_not_replay_messages():
messages = [
{"role": "user", "content": "ask"},
{"role": "tool", "tool_call_id": "c1", "content": "result"},
]
assert tool_result_text_indices(messages, ["result"]) == frozenset()
assert tool_result_text_indices(messages, ["ask", "tampered"]) == frozenset()
def test_tool_result_text_indices_skips_blank_tool_content_never_submitted_as_tool():
messages = [{"role": "tool", "tool_call_id": "c1", "content": " "}]
assert extract_tool_results(messages) == ()
assert tool_result_text_indices(messages, [" "]) == frozenset()
def test_tool_result_text_indices_empty_without_structured_messages():
assert tool_result_text_indices(None, ["ask"]) == frozenset()