diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py index 65c068162e9..b22b8a41dae 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py @@ -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() diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix_extraction.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix_extraction.py index f3f91a85a56..c614daf6182 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix_extraction.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix_extraction.py @@ -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[^;,]+)?(?P(?:;[^;,]+)*?)(?P;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) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py index 42fd6c19ad0..83f01091038 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix.py @@ -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(): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix_extraction.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix_extraction.py index 25d67d8dda9..d9f27f4f56f 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix_extraction.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_ovalix_extraction.py @@ -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()