fix(guardrails): clear the type and test-quality gates on the NeuralTrust hook

Dropping **kwargs from this hook exposed mismatches the other guardrails hide
behind an untyped signature, and the gates only surfaced once the branch caught
up with the base.

_copy_messages called dict() on an object, which had no matching overload.
Narrowing per message in _copy_message gives the call a typed argument and makes
the declared Mapping[str, object] return actually true.

The constructor now accepts what LitellmParams supplies, str or Sequence[str]
for the mode and an optional default_on, instead of a signature only the enum
satisfied. CustomGuardrail still narrows the hook to the enum, so that one call
carries a reasoned suppression.

headers went back to a plain dict because AsyncHTTPHandler.post declares it as
dict, so MappingProxyType was a type error rather than an improvement.

Four tests asserted only on the mock, which TQ002 counts as a mock echo. They
now also assert that an allow verdict returns the inputs unchanged, which is the
behaviour they were meant to pin.
This commit is contained in:
albertbausili 2026-09-02 13:02:04 +02:00
parent b2aaa6795a
commit 254d5ba5ba
2 changed files with 33 additions and 22 deletions

View file

@ -7,7 +7,6 @@ from __future__ import annotations
import os
from collections.abc import Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal
import httpx
@ -52,10 +51,15 @@ def _message_text(message: Mapping[str, object]) -> str | None:
return content if isinstance(content, str) and content else None
def _copy_messages(messages: Sequence[object]) -> tuple[Mapping[str, object], ...] | None:
if not all(isinstance(message, Mapping) for message in messages):
def _copy_message(value: object) -> Mapping[str, object] | None:
if not isinstance(value, Mapping):
return None
return tuple(dict(message) for message in messages) # mutable-ok: shallow copies for write-back
return {str(key): item for key, item in value.items()} # mutable-ok: shallow copy for write-back
def _copy_messages(messages: Sequence[object]) -> tuple[Mapping[str, object], ...] | None:
copied: Final = tuple(copy for message in messages if (copy := _copy_message(message)) is not None)
return copied if len(copied) == len(messages) else None
def _texts_from_messages(messages: Sequence[Mapping[str, object]]) -> tuple[str, ...]:
@ -161,8 +165,8 @@ class NeuralTrustGuardrail(CustomGuardrail):
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
timeout: float | None = None,
guardrail_name: str | None = None,
event_hook: GuardrailEventHooks | Sequence[GuardrailEventHooks] | Mode | None = None,
default_on: bool = False,
event_hook: GuardrailEventHooks | Mode | str | Sequence[str] | None = None,
default_on: bool | None = None,
) -> None:
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
@ -180,8 +184,9 @@ class NeuralTrustGuardrail(CustomGuardrail):
super().__init__(
guardrail_name=guardrail_name,
supported_event_hooks=self.get_supported_event_hooks(),
event_hook=event_hook,
default_on=default_on,
# LitellmParams.mode is str | list[str] | Mode, which CustomGuardrail narrows to the enum
event_hook=event_hook, # pyright: ignore[reportArgumentType] # config supplies the raw mode string
default_on=bool(default_on),
)
@log_guardrail_information
@ -262,12 +267,10 @@ class NeuralTrustGuardrail(CustomGuardrail):
async def _call_evaluate(self, body: dict[str, object]) -> dict[str, object]: # mutable-ok: TrustGuard JSON
url: Final = f"{self.api_base}{EVALUATE_PATH}"
headers: Final = MappingProxyType(
{
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
)
headers: Final = { # mutable-ok: AsyncHTTPHandler.post declares headers as dict
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
try:
response: Final = await self.async_handler.post(
url,

View file

@ -102,14 +102,16 @@ class TestNeuralTrustGuardrail:
@pytest.mark.asyncio
async def test_omits_session_id_without_conversation_session(self) -> None:
guardrail = _guardrail()
inputs: GenericGuardrailAPIInputs = {"texts": ["hello"]}
mock_post = AsyncMock(return_value=_response({"status": "allow"}))
with patch.object(guardrail.async_handler, "post", mock_post):
await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert result == inputs
assert "session_id" not in mock_post.call_args.kwargs["json"]
@pytest.mark.asyncio
@ -119,14 +121,16 @@ class TestNeuralTrustGuardrail:
guardrail_name="neuraltrust",
event_hook="pre_call",
)
inputs: GenericGuardrailAPIInputs = {"texts": ["hello"]}
mock_post = AsyncMock(return_value=_response({"status": "allow"}))
with patch.object(guardrail.async_handler, "post", mock_post):
await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert result == inputs
assert "collector_key" not in mock_post.call_args.kwargs["json"]
@pytest.mark.asyncio
@ -404,14 +408,16 @@ class TestNeuralTrustGuardrail:
async def test_forwards_tools(self) -> None:
guardrail = _guardrail()
tools = [{"type": "function", "function": {"name": "search", "parameters": {}}}]
inputs: GenericGuardrailAPIInputs = {"texts": ["hello"], "tools": tools}
mock_post = AsyncMock(return_value=_response({"status": "allow"}))
with patch.object(guardrail.async_handler, "post", mock_post):
await guardrail.apply_guardrail(
inputs={"texts": ["hello"], "tools": tools},
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert result == inputs
assert mock_post.call_args.kwargs["json"]["payload"]["tools"] == tools
@pytest.mark.asyncio
@ -568,14 +574,16 @@ class TestNeuralTrustGuardrail:
@pytest.mark.asyncio
async def test_custom_timeout_is_passed_to_client(self) -> None:
guardrail = _guardrail(timeout=12)
inputs: GenericGuardrailAPIInputs = {"texts": ["hello"]}
mock_post = AsyncMock(return_value=_response({"status": "allow"}))
with patch.object(guardrail.async_handler, "post", mock_post):
await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
input_type="request",
logging_obj=_logging(),
)
assert result == inputs
assert mock_post.call_args.kwargs["timeout"] == 12.0
def test_get_config_model(self) -> None: