mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
b2aaa6795a
commit
254d5ba5ba
2 changed files with 33 additions and 22 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue