mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(ovalix): scan tool calls carried on messages, tolerate optional regex groups
Tool-call blocking only read the top-level tool_calls input. The Anthropic request path fills structured_messages but leaves that input empty, so calls made in prior assistant turns were never checkpointed. Merge both sources and dedupe on the tool payload, since the OpenAI path populates both and would otherwise scan every call twice. _extract_application_name called strip() on group(1) whenever the regex had any groups. An optional group that did not participate is None, which raised AttributeError instead of falling through to the no-application handling.
This commit is contained in:
parent
da9e61431b
commit
3641581823
3 changed files with 84 additions and 5 deletions
|
|
@ -34,9 +34,11 @@ from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix_extraction import (
|
|||
FilePart,
|
||||
extract_file_parts_from_images,
|
||||
extract_file_parts_from_messages,
|
||||
extract_tool_calls_from_messages,
|
||||
extract_tool_results,
|
||||
make_tool_data,
|
||||
tool_call_to_tool_data,
|
||||
tool_data_key,
|
||||
tool_result_text_indices,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
|
@ -418,9 +420,14 @@ class OvalixGuardrail(CustomGuardrail):
|
|||
)
|
||||
return inputs
|
||||
|
||||
tool_call_items: Final = tuple(
|
||||
("TOOL", td) for td in (tool_call_to_tool_data(tc) for tc in (inputs.get("tool_calls") or ())) if td
|
||||
tool_calls: Final = (
|
||||
*(inputs.get("tool_calls") or ()),
|
||||
*extract_tool_calls_from_messages(structured_messages),
|
||||
)
|
||||
unique_tool_data: Final = {
|
||||
tool_data_key(data): data for data in (tool_call_to_tool_data(tc) for tc in tool_calls) if data
|
||||
}
|
||||
tool_call_items: Final = tuple(("TOOL", data) for data in unique_tool_data.values())
|
||||
tool_block: Final = await self._check_items_block_only(
|
||||
tool_call_items,
|
||||
prompt_checkpoint,
|
||||
|
|
@ -562,8 +569,8 @@ class OvalixGuardrail(CustomGuardrail):
|
|||
match: Final = regex.search(alias)
|
||||
if not match:
|
||||
return None
|
||||
name: Final = (match.group(1) if match.groups() else match.group(0)).strip()
|
||||
return name or None
|
||||
captured: Final = (match.group(1) if match.groups() else match.group(0)) or ""
|
||||
return captured.strip() or None
|
||||
|
||||
def _routing_cache_get(self, name: str) -> tuple[bool, ResolvedRouting | None]:
|
||||
entry: Final = self._routing_cache.get(name)
|
||||
|
|
|
|||
|
|
@ -261,6 +261,30 @@ def tool_call_to_tool_data(tool_call: object) -> Mapping[str, object] | None:
|
|||
return make_tool_data(name, content, tool_input)
|
||||
|
||||
|
||||
def _message_tool_calls(message: Mapping[str, object]) -> Sequence[object]:
|
||||
tool_calls: Final = message.get("tool_calls")
|
||||
return tool_calls if isinstance(tool_calls, list) else ()
|
||||
|
||||
|
||||
def extract_tool_calls_from_messages(structured_messages: Sequence[object] | None) -> tuple[object, ...]:
|
||||
"""Tool calls declared on the messages themselves.
|
||||
|
||||
Surfaces such as the Anthropic request path populate ``structured_messages`` but leave the
|
||||
top-level ``tool_calls`` input empty, so calls made in prior assistant turns are only visible here.
|
||||
"""
|
||||
return tuple(
|
||||
tool_call
|
||||
for message in structured_messages or ()
|
||||
if isinstance(message, Mapping)
|
||||
for tool_call in _message_tool_calls(message)
|
||||
)
|
||||
|
||||
|
||||
def tool_data_key(tool_data: Mapping[str, object]) -> str:
|
||||
"""Stable identity for a tool payload, so a call reached from two sources is only scanned once."""
|
||||
return json.dumps(tool_data, sort_keys=True, default=str)
|
||||
|
||||
|
||||
def _tool_content_blocks(content: Sequence[object]) -> Iterator[str]:
|
||||
for block in content:
|
||||
if isinstance(block, Mapping):
|
||||
|
|
|
|||
|
|
@ -1159,7 +1159,7 @@ async def test_every_checkpoint_payload_is_json_serializable():
|
|||
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 sorted(p["data_type"] for p in payloads) == ["FILE", "TEXT", "TEXT", "TOOL", "TOOL", "TOOL"]
|
||||
assert all(p["data"]["tool_input"] == {} for p in payloads if p["data_type"] == "TOOL")
|
||||
|
||||
|
||||
|
|
@ -1745,3 +1745,51 @@ async def test_apply_guardrail_resolved_by_alias_routes_checkpoints_by_name():
|
|||
assert seen["body"]["application_name"] == "Weather App"
|
||||
assert seen["body"]["input_type"] == "request"
|
||||
assert "application_id" not in seen["body"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_message_tool_calls_are_scanned_without_top_level_tool_calls():
|
||||
g = _static_guardrail()
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=[],
|
||||
structured_messages=[
|
||||
{
|
||||
"role": "assistant",
|
||||
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "exfil", "arguments": "{}"}}],
|
||||
}
|
||||
],
|
||||
)
|
||||
with patch.object(
|
||||
g._async_handler, "post", new=_post_returning(lambda body: _BLOCK if body["data_type"] == "TOOL" else _ALLOW)
|
||||
):
|
||||
with pytest.raises(OvalixGuardrailBlockedException):
|
||||
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_present_in_both_sources_is_scanned_once():
|
||||
g = _static_guardrail()
|
||||
call = {"id": "c1", "type": "function", "function": {"name": "get_weather", "arguments": '{"city": "sf"}'}}
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
texts=[],
|
||||
tool_calls=[call],
|
||||
structured_messages=[{"role": "assistant", "tool_calls": [call]}],
|
||||
)
|
||||
tool_bodies = []
|
||||
|
||||
async def _post(url, headers=None, json=None):
|
||||
if json["data_type"] == "TOOL":
|
||||
tool_bodies.append(json)
|
||||
r = MagicMock()
|
||||
r.json.return_value = _ALLOW
|
||||
r.raise_for_status = MagicMock()
|
||||
return r
|
||||
|
||||
with patch.object(g._async_handler, "post", new=_post):
|
||||
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request", logging_obj=None)
|
||||
assert len(tool_bodies) == 1
|
||||
|
||||
|
||||
def test_extract_application_name_tolerates_unmatched_optional_group():
|
||||
g = _static_guardrail()
|
||||
assert g._extract_application_name("app-", re.compile(r"app-(\w+)?")) is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue