fix(guardrails): treat null tool_calls in a TrustGuard transform as untouched

An echoing TrustGuard serialises an assistant message without tool
calls as tool_calls: null, which the hook rejected as a malformed
transform. Null now means the same as an omitted key, so the original
tool calls are kept, while any other non-list value still fails closed.
A dead branch in the last-user-message rewrite is gone: the caller never
passes an empty list

Tests now cover every branch of the hook and its initializer: the 401
and 403 passthrough, a 200 whose body is not JSON under both fallback
modes, non-object entries on both transform paths, a transform with no
text to rewrite, and the initializer wiring its params and registering
the callback
This commit is contained in:
albertbausili 2026-09-14 09:03:15 +02:00
parent fee2ff525d
commit 976d1466c7
2 changed files with 167 additions and 4 deletions

View file

@ -91,9 +91,9 @@ def _texts_from_messages(messages: Sequence[Mapping[str, object]]) -> tuple[str,
def _tool_calls_in_message(message: Mapping[str, object]) -> tuple[object, ...] | None:
if "tool_calls" not in message:
raw: Final = message.get("tool_calls")
if raw is None:
return None
raw: Final = message["tool_calls"]
if not isinstance(raw, list):
raise HTTPException(status_code=400, detail=TRANSFORM_MISSING)
return tuple(raw)
@ -112,8 +112,6 @@ def _rewrite_last_user_message(
) -> tuple[Mapping[str, object], ...]:
user_indices: Final = tuple(index for index, message in enumerate(messages) if message.get("role") == "user")
target: Final = user_indices[-1] if user_indices else len(messages) - 1
if target < 0:
return ({"role": "user", "content": redacted},) # mutable-ok: write-back message
return tuple(
{**message, "content": redacted} if index == target else dict(message) # mutable-ok: write-back message
for index, message in enumerate(messages)

View file

@ -3,6 +3,7 @@ from typing import Literal
from unittest.mock import AsyncMock, patch
import httpx
import litellm
import pytest
from fastapi import HTTPException
from httpx import Request, Response
@ -13,6 +14,7 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTra
from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_endpoints import get_provider_specific_params
from litellm.proxy.guardrails.guardrail_hooks.neuraltrust import initialize_guardrail
from litellm.proxy.guardrails.guardrail_hooks.neuraltrust.neuraltrust import (
NeuralTrustGuardrail,
)
@ -671,6 +673,98 @@ class TestNeuralTrustGuardrail:
)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_transform_messages_with_non_object_entry_fail_closed(self) -> None:
guardrail = _guardrail()
mock_post = AsyncMock(
return_value=_response({"status": "transform", "transformed_payload": {"messages": ["REDACTED"]}})
)
with patch.object(guardrail.async_handler, "post", mock_post):
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["secret"], "structured_messages": [{"role": "user", "content": "secret"}]},
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_transform_input_with_non_object_original_message_fail_closed(self) -> None:
guardrail = _guardrail()
mock_post = AsyncMock(
return_value=_response({"status": "transform", "transformed_payload": {"input": "[REDACTED]"}})
)
with patch.object(guardrail.async_handler, "post", mock_post):
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["secret"], "structured_messages": ["secret"]}, # pyright: ignore[reportArgumentType] # malformed on purpose
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_transform_input_without_any_text_fail_closed(self) -> None:
guardrail = _guardrail()
mock_post = AsyncMock(
return_value=_response({"status": "transform", "transformed_payload": {"input": "[REDACTED]"}})
)
with patch.object(guardrail.async_handler, "post", mock_post):
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": [], "tool_calls": [{"id": "call_1", "type": "function", "function": {}}]},
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert exc_info.value.status_code == 400
assert "transform missing payload" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_transform_null_tool_calls_keeps_the_original_ones(self) -> None:
guardrail = _guardrail(event_hook="post_call")
original_tool_calls = [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}]
mock_post = AsyncMock(
return_value=_response(
{
"status": "transform",
"transformed_payload": {"messages": [{"role": "assistant", "content": "ok", "tool_calls": None}]},
}
)
)
with patch.object(guardrail.async_handler, "post", mock_post):
result = await guardrail.apply_guardrail(
inputs={"texts": ["secret"], "tool_calls": original_tool_calls},
request_data={},
input_type="response",
logging_obj=_logging(),
)
assert result["texts"] == ["ok"]
assert result["tool_calls"] is original_tool_calls
@pytest.mark.asyncio
async def test_transform_non_list_tool_calls_fail_closed(self) -> None:
guardrail = _guardrail(event_hook="post_call")
mock_post = AsyncMock(
return_value=_response(
{
"status": "transform",
"transformed_payload": {"messages": [{"role": "assistant", "content": "ok", "tool_calls": {}}]},
}
)
)
with patch.object(guardrail.async_handler, "post", mock_post):
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["secret"], "tool_calls": [{"id": "call_1", "type": "function", "function": {}}]},
request_data={},
input_type="response",
logging_obj=_logging(),
)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_forwards_tools(self) -> None:
guardrail = _guardrail()
@ -766,6 +860,53 @@ class TestNeuralTrustGuardrail:
assert exc_info.value.status_code == 503
assert "request failed" in str(exc_info.value.detail)
@pytest.mark.asyncio
@pytest.mark.parametrize("status_code", [401, 403])
async def test_auth_failures_return_their_status_even_if_fail_open(self, status_code: int) -> None:
guardrail = _guardrail(unreachable_fallback="fail_open")
mock_post = AsyncMock(return_value=_response({"error": "bad key"}, status_code=status_code))
with patch.object(guardrail.async_handler, "post", mock_post):
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert exc_info.value.status_code == status_code
assert "authentication failed" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_non_json_200_fail_closed(self) -> None:
guardrail = _guardrail()
request = Request("POST", "https://trustguard.neuraltrust.ai/v1/evaluate")
mock_post = AsyncMock(return_value=Response(200, request=request, text="<html>captive portal</html>"))
with patch.object(guardrail.async_handler, "post", mock_post):
with pytest.raises(HTTPException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert exc_info.value.status_code == 503
assert "unreachable" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_non_json_200_follows_fail_open(self) -> None:
guardrail = _guardrail(unreachable_fallback="fail_open")
inputs: GenericGuardrailAPIInputs = {"texts": ["hello"]}
request = Request("POST", "https://trustguard.neuraltrust.ai/v1/evaluate")
mock_post = AsyncMock(return_value=Response(200, request=request, text="<html>captive portal</html>"))
with patch.object(guardrail.async_handler, "post", mock_post):
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert result == inputs
@pytest.mark.asyncio
async def test_http_502_follows_fail_open(self) -> None:
guardrail = _guardrail(unreachable_fallback="fail_open")
@ -885,6 +1026,30 @@ class TestNeuralTrustGuardrail:
with pytest.raises(ValueError, match="positive"):
_guardrail(timeout=timeout)
def test_initializer_wires_params_and_registers_the_callback(self) -> None:
params = LitellmParams(
guardrail="neuraltrust",
mode="post_call",
api_key="tgk_from_params",
api_base="https://trustguard.example.test/",
collector_key="tgcol_from_params",
unreachable_fallback="fail_open",
timeout=2,
default_on=True,
)
hook = initialize_guardrail(params, {"guardrail_name": "tg-prod"})
try:
assert hook.api_key == "tgk_from_params"
assert hook.api_base == "https://trustguard.example.test"
assert hook.collector_key == "tgcol_from_params"
assert hook.unreachable_fallback == "fail_open"
assert hook.timeout == 2.0
assert hook.guardrail_name == "tg-prod"
assert hook.default_on is True
assert hook in litellm.callbacks
finally:
litellm.logging_callback_manager.remove_callback_from_all_lists(hook)
def test_registry_contains_neuraltrust(self) -> None:
from litellm.proxy.guardrails.guardrail_hooks.neuraltrust import (
NeuralTrustGuardrail as Registered,