mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fixing PR comments
This commit is contained in:
parent
3ec57cedb3
commit
51855a20f5
4 changed files with 435 additions and 192 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue