mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(guardrails): resolve comments
This commit is contained in:
parent
3f0edf6013
commit
e2d5312a8a
3 changed files with 89 additions and 62 deletions
|
|
@ -134,26 +134,6 @@ class SingulrGuardrail(CustomGuardrail):
|
|||
def _build_user_message(text: str) -> Mapping[str, Any]:
|
||||
return {"role": "user", "content": text} # mutable-ok: short-lived JSON payload dict
|
||||
|
||||
@staticmethod
|
||||
def _extract_content_text(content: str | Sequence[Mapping[str, Any]] | None) -> str | None:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
text: Final = "\n".join(block.get("text", "") for block in content if block.get("type") == "text")
|
||||
return text or None
|
||||
return None
|
||||
|
||||
def _extract_completion_text(self, response: Mapping[str, Any]) -> str | None:
|
||||
choices: Final = response.get("choices") or ()
|
||||
for choice in choices:
|
||||
if choice.get("finish_reason") != "stop":
|
||||
continue
|
||||
message = choice.get("message") or _EMPTY_MAPPING
|
||||
text = self._extract_content_text(message.get("content"))
|
||||
if text:
|
||||
return text
|
||||
return None
|
||||
|
||||
def _build_headers(self) -> Mapping[str, str]:
|
||||
all_headers: Final = MappingProxyType(
|
||||
{
|
||||
|
|
@ -226,9 +206,10 @@ class SingulrGuardrail(CustomGuardrail):
|
|||
)
|
||||
|
||||
images: Final = inputs.get("images")
|
||||
tools: Final = inputs.get("tools")
|
||||
|
||||
if not messages and not images:
|
||||
verbose_proxy_logger.debug("Singulr: No messages or images to check after filtering")
|
||||
if not messages and not images and not tools:
|
||||
verbose_proxy_logger.debug("Singulr: No messages, images, or tools to check after filtering")
|
||||
return inputs
|
||||
|
||||
metadata: Final = self._build_metadata(request_data=request_data)
|
||||
|
|
@ -239,6 +220,7 @@ class SingulrGuardrail(CustomGuardrail):
|
|||
guardrail_scope="request",
|
||||
messages=messages,
|
||||
images=images,
|
||||
tools=tools,
|
||||
metadata=metadata,
|
||||
)
|
||||
payload = singulr_req_obj.model_dump(mode="json")
|
||||
|
|
@ -387,18 +369,20 @@ class SingulrGuardrail(CustomGuardrail):
|
|||
await self._call_api(payload_req)
|
||||
|
||||
if result:
|
||||
completion_text = self._extract_completion_text(result)
|
||||
assistant_message = AssistantMessage(
|
||||
role="assistant",
|
||||
content=completion_text,
|
||||
tool_calls=(),
|
||||
)
|
||||
singulr_res_obj = SingulrGuardrailPayload(
|
||||
correlation_id=kwargs.get("litellm_call_id"),
|
||||
guardrail_scope="response",
|
||||
response=assistant_message,
|
||||
response=result,
|
||||
)
|
||||
payload = singulr_res_obj.model_dump(mode="json")
|
||||
try:
|
||||
payload = singulr_res_obj.model_dump(mode="json")
|
||||
except Exception as exc: # noqa: BLE001 # result can be any callback shape; fall back to a stringified report
|
||||
verbose_proxy_logger.debug("Singulr: could not JSON-serialize response, falling back: %s", exc)
|
||||
payload = {
|
||||
"correlation_id": kwargs.get("litellm_call_id"),
|
||||
"guardrail_scope": "response",
|
||||
"response": str(result),
|
||||
}
|
||||
await self._call_api(payload)
|
||||
except GuardrailRaisedException:
|
||||
guardrail_status = "guardrail_intervened"
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ from typing import Any, Literal
|
|||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from litellm.types.llms.openai import ChatCompletionToolParam
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
|
||||
|
|
@ -35,7 +37,8 @@ class SingulrGuardrailPayload(BaseModel):
|
|||
guardrail_scope: str | None = None
|
||||
messages: Sequence[Any] | None = None
|
||||
images: Sequence[str] | None = None
|
||||
response: AssistantMessage | None = None
|
||||
tools: Sequence[ChatCompletionToolParam] | None = None
|
||||
response: Any = None # pyright: ignore[reportExplicitAny] # logging_only reports raw litellm callback results (ModelResponse, EmbeddingResponse, etc.)
|
||||
metadata: Mapping[str, Any] | None = None
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm.types.guardrails import GuardrailEventHooks
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.singulr import (
|
||||
SingulrGuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -167,6 +168,41 @@ class TestSingulrRequestPayload:
|
|||
mock_post.assert_not_called()
|
||||
assert result == {"texts": []}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tools_are_forwarded(self, singulr_guardrail):
|
||||
"""Regression: tool/function definitions are client-controlled and can
|
||||
carry prompt-injection content, so they must reach Singulr for
|
||||
inspection instead of only messages and images."""
|
||||
resp = _make_response({"should_block": False})
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": "search_docs", "description": "Search internal docs", "parameters": {}},
|
||||
}
|
||||
]
|
||||
with patch.object(singulr_guardrail.async_handler, "post", return_value=resp) as mock_post:
|
||||
await singulr_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["How do I reset my password?"], "tools": tools},
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
sent_payload = mock_post.call_args.kwargs["json"]
|
||||
assert sent_payload["tools"] == tools
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tools_alone_still_triggers_the_api_call(self, singulr_guardrail):
|
||||
"""Regression: a request with tool definitions but no text or images
|
||||
must still be checked, not skipped for lack of a message."""
|
||||
resp = _make_response({"should_block": False})
|
||||
tools = [{"type": "function", "function": {"name": "delete_file", "description": "", "parameters": {}}}]
|
||||
with patch.object(singulr_guardrail.async_handler, "post", return_value=resp) as mock_post:
|
||||
await singulr_guardrail.apply_guardrail(
|
||||
inputs={"texts": [], "tools": tools},
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
mock_post.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_images_alone_still_triggers_the_api_call(self, singulr_guardrail):
|
||||
"""Regression: an image-only request (no text) must still be checked,
|
||||
|
|
@ -493,36 +529,6 @@ class TestSingulrApplyGuardrailDispatch:
|
|||
assert result is inputs
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Content extraction helpers (used by the logging_only hook)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSingulrContentExtraction:
|
||||
def test_extract_content_text_from_plain_string(self, singulr_guardrail):
|
||||
assert singulr_guardrail._extract_content_text("hello") == "hello"
|
||||
|
||||
def test_extract_content_text_from_content_blocks(self, singulr_guardrail):
|
||||
content = [{"type": "text", "text": "hello"}, {"type": "image_url", "image_url": {}}]
|
||||
assert singulr_guardrail._extract_content_text(content) == "hello"
|
||||
|
||||
def test_extract_content_text_returns_none_for_no_text_blocks(self, singulr_guardrail):
|
||||
content = [{"type": "image_url", "image_url": {}}]
|
||||
assert singulr_guardrail._extract_content_text(content) is None
|
||||
|
||||
def test_extract_completion_text_skips_non_stop_choices(self, singulr_guardrail):
|
||||
response = {
|
||||
"choices": [
|
||||
{"finish_reason": "tool_calls", "message": {"content": "should be skipped"}},
|
||||
{"finish_reason": "stop", "message": {"content": "final answer"}},
|
||||
]
|
||||
}
|
||||
assert singulr_guardrail._extract_completion_text(response) == "final answer"
|
||||
|
||||
def test_extract_completion_text_returns_none_for_no_choices(self, singulr_guardrail):
|
||||
assert singulr_guardrail._extract_completion_text({}) is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# logging_only hook
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -542,7 +548,41 @@ class TestSingulrLoggingHook:
|
|||
assert request_payload["guardrail_scope"] == "request"
|
||||
assert request_payload["messages"] == kwargs["messages"]
|
||||
assert response_payload["guardrail_scope"] == "response"
|
||||
assert response_payload["response"]["content"] == "hello there"
|
||||
assert response_payload["response"] == result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_a_real_model_response_without_swallowing_it(self, singulr_guardrail):
|
||||
"""Regression: a normal completion callback passes a ModelResponse, not a
|
||||
dict. The response payload must carry its actual serialized content instead
|
||||
of silently dropping it because ModelResponse isn't a Mapping."""
|
||||
resp = _make_response({"should_block": False})
|
||||
result = ModelResponse(
|
||||
choices=[{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "hello there"}}]
|
||||
)
|
||||
with patch.object(singulr_guardrail.async_handler, "post", return_value=resp) as mock_post:
|
||||
await singulr_guardrail.async_logging_hook(kwargs={}, result=result, call_type="acompletion")
|
||||
|
||||
response_payload = mock_post.call_args.kwargs["json"]
|
||||
assert response_payload["guardrail_scope"] == "response"
|
||||
assert response_payload["response"]["choices"][0]["message"]["content"] == "hello there"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_serializable_result_falls_back_to_string_report(self, singulr_guardrail):
|
||||
"""A result that pydantic can't serialize to JSON must still get reported,
|
||||
as a stringified fallback, instead of raising out of the logging_only hook."""
|
||||
resp = _make_response({"should_block": False})
|
||||
|
||||
class Unserializable:
|
||||
def __repr__(self) -> str:
|
||||
return "<Unserializable>"
|
||||
|
||||
with patch.object(singulr_guardrail.async_handler, "post", return_value=resp) as mock_post:
|
||||
await singulr_guardrail.async_logging_hook(
|
||||
kwargs={}, result=Unserializable(), call_type="acompletion"
|
||||
)
|
||||
|
||||
response_payload = mock_post.call_args.kwargs["json"]
|
||||
assert response_payload["response"] == "<Unserializable>"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_messages_and_no_result_skips_both_api_calls(self, singulr_guardrail):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue