mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
greptileai comments fixes
This commit is contained in:
parent
3718682021
commit
f3815406c0
2 changed files with 134 additions and 194 deletions
|
|
@ -5,12 +5,12 @@ post_call (model output) checkpoints with optional correction/blocking.
|
|||
"""
|
||||
|
||||
import datetime
|
||||
import hashlib
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
|
|
@ -22,7 +22,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import AllMessageValues, GenericGuardrailAPIInputs
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
|
@ -30,7 +30,6 @@ if TYPE_CHECKING:
|
|||
|
||||
BLOCKED_BY_OVALIX_FALLBACK_MESSAGE = "This message was blocked by Ovalix"
|
||||
BLOCKED_ACTION_TYPE = "block"
|
||||
USER_MESSAGE_ROLE = "user"
|
||||
|
||||
|
||||
class OvalixGuardrailMissingSecrets(Exception):
|
||||
|
|
@ -181,7 +180,7 @@ class OvalixGuardrail(CustomGuardrail):
|
|||
|
||||
def _get_session_id(self, data: dict) -> str:
|
||||
"""Return a unique identifier for the chat/session (actor + date + application_id)."""
|
||||
actor = hash(self._get_actor(data))
|
||||
actor = hashlib.sha256(self._get_actor(data).encode()).hexdigest()[:8]
|
||||
today = datetime.datetime.now().strftime("%Y-%m-%d")
|
||||
return f"{actor}_{today}_{self._application_id}"
|
||||
|
||||
|
|
@ -239,124 +238,36 @@ class OvalixGuardrail(CustomGuardrail):
|
|||
|
||||
actor = self._get_actor(request_data)
|
||||
session_id = self._get_session_id(request_data)
|
||||
texts = inputs.get("texts") or []
|
||||
if not texts or not isinstance(texts, list):
|
||||
return inputs
|
||||
|
||||
if input_type == "response":
|
||||
llm_response = self._get_llm_response_text(
|
||||
request_data.get("response", None)
|
||||
if not self._post_checkpoint_id:
|
||||
return inputs
|
||||
corrected_llm_responses = await self._generate_post_guardrail_llm_texts(
|
||||
texts, actor, session_id, self._post_checkpoint_id
|
||||
)
|
||||
if llm_response:
|
||||
(
|
||||
corrected_llm_response,
|
||||
is_blocked,
|
||||
) = await self._handle_post_llm_response(
|
||||
llm_response, actor, session_id
|
||||
)
|
||||
# TODO: set the llm response text to `corrected_llm_response`. will be addressed later.
|
||||
return inputs
|
||||
|
||||
messages = inputs.get("structured_messages") or []
|
||||
if not messages:
|
||||
return inputs
|
||||
return {**inputs, "texts": corrected_llm_responses}
|
||||
|
||||
if self._pre_checkpoint_id:
|
||||
post_guardrail_texts = await self._generate_post_guardrail_text(
|
||||
messages, actor, session_id
|
||||
post_guardrail_texts = await self._generate_post_guardrail_llm_texts(
|
||||
texts, actor, session_id, self._pre_checkpoint_id
|
||||
)
|
||||
return {**inputs, "texts": post_guardrail_texts}
|
||||
return inputs
|
||||
|
||||
def _block_current_message(self, blocking_message: str) -> None:
|
||||
"""Raise OvalixGuardrailBlockedException with the given message (no default wrapper)."""
|
||||
raise OvalixGuardrailBlockedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=blocking_message,
|
||||
should_wrap_with_default_message=False,
|
||||
)
|
||||
|
||||
def _get_llm_response_text(
|
||||
self, response: Optional[litellm.ModelResponse]
|
||||
) -> Optional[str]:
|
||||
"""Extract the first assistant text content from a ModelResponse, or None."""
|
||||
if not response:
|
||||
return None
|
||||
if isinstance(response, litellm.ModelResponse):
|
||||
for choice in response.choices:
|
||||
if isinstance(choice, litellm.Choices):
|
||||
if choice.message.content and isinstance(
|
||||
choice.message.content, str
|
||||
):
|
||||
return choice.message.content
|
||||
return None
|
||||
|
||||
async def _handle_post_llm_response(
|
||||
self, llm_response: str, actor: str, session_id: str
|
||||
) -> tuple[str, bool]:
|
||||
"""Run post-call checkpoint on model output; return corrected text or raise if blocked."""
|
||||
if not self._post_checkpoint_id:
|
||||
raise ValueError(
|
||||
"Ovalix: post-checkpoint ID is required for post_call handling."
|
||||
)
|
||||
|
||||
try:
|
||||
resp = await self._call_checkpoint(
|
||||
llm_response, self._post_checkpoint_id, actor, session_id
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Ovalix apply_guardrail checkpoint call failed: %s", e
|
||||
)
|
||||
raise GuardrailRaisedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=f"Ovalix guardrail error: {e!s}",
|
||||
should_wrap_with_default_message=False,
|
||||
) from e
|
||||
|
||||
action_type = (resp.get("action_type") or "").lower()
|
||||
if action_type == BLOCKED_ACTION_TYPE:
|
||||
blocking_message = (
|
||||
self._get_trackers_corrected_message(resp)
|
||||
or BLOCKED_BY_OVALIX_FALLBACK_MESSAGE
|
||||
)
|
||||
return blocking_message, True
|
||||
return self._get_trackers_corrected_message(resp) or llm_response, False
|
||||
|
||||
async def _generate_post_guardrail_text(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
actor: str,
|
||||
session_id: str,
|
||||
async def _generate_post_guardrail_llm_texts(
|
||||
self, texts: List[str], actor: str, session_id: str, checkpoint_id: str
|
||||
) -> List[str]:
|
||||
"""
|
||||
Generate post-guardrail text for the given messages.
|
||||
|
||||
Args:
|
||||
messages: List of messages
|
||||
actor: Actor
|
||||
session_id: Session ID
|
||||
request_data: Request data
|
||||
|
||||
Returns:
|
||||
List of post-guardrail texts
|
||||
"""
|
||||
is_last_prompt = True
|
||||
"""Generate post-guardrail LLM responses for the given LLM responses."""
|
||||
post_guardrail_texts: List[str] = []
|
||||
|
||||
if not self._pre_checkpoint_id:
|
||||
# should not happen - if it does, the guardrail is not configured correctly and self._validate_config did not raise an error
|
||||
raise ValueError("Ovalix: pre-checkpoint ID is required")
|
||||
|
||||
for message in reversed(messages):
|
||||
content = message.get("content", None) or ""
|
||||
if not isinstance(content, str):
|
||||
continue
|
||||
message_role = message.get("role", None)
|
||||
if message_role and message_role != USER_MESSAGE_ROLE:
|
||||
# we are not scanning the assistant/system/developer past responses, only the responses that the user sent
|
||||
post_guardrail_texts.insert(0, content)
|
||||
continue
|
||||
is_first_response = True
|
||||
for llm_response in reversed(texts):
|
||||
try:
|
||||
resp = await self._call_checkpoint(
|
||||
content, self._pre_checkpoint_id, actor, session_id
|
||||
llm_response, checkpoint_id, actor, session_id
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
|
|
@ -369,21 +280,30 @@ class OvalixGuardrail(CustomGuardrail):
|
|||
) from e
|
||||
|
||||
action_type = (resp.get("action_type") or "").lower()
|
||||
if action_type == BLOCKED_ACTION_TYPE:
|
||||
blocking_message = (
|
||||
self._get_trackers_corrected_message(resp)
|
||||
or BLOCKED_BY_OVALIX_FALLBACK_MESSAGE
|
||||
)
|
||||
if is_last_prompt:
|
||||
self._block_current_message(blocking_message)
|
||||
else:
|
||||
post_guardrail_texts.insert(0, blocking_message)
|
||||
blocking_message = (
|
||||
self._get_trackers_corrected_message(resp)
|
||||
or BLOCKED_BY_OVALIX_FALLBACK_MESSAGE
|
||||
)
|
||||
if action_type == BLOCKED_ACTION_TYPE and is_first_response:
|
||||
self._block_current_message(blocking_message)
|
||||
elif action_type == BLOCKED_ACTION_TYPE:
|
||||
post_guardrail_texts.insert(0, blocking_message)
|
||||
else:
|
||||
new_content = self._get_trackers_corrected_message(resp) or content
|
||||
post_guardrail_texts.insert(0, new_content)
|
||||
is_last_prompt = False
|
||||
corrected_text = (
|
||||
self._get_trackers_corrected_message(resp) or llm_response
|
||||
)
|
||||
post_guardrail_texts.insert(0, corrected_text)
|
||||
is_first_response = False
|
||||
return post_guardrail_texts
|
||||
|
||||
def _block_current_message(self, blocking_message: str) -> None:
|
||||
"""Raise OvalixGuardrailBlockedException with the given message (no default wrapper)."""
|
||||
raise OvalixGuardrailBlockedException(
|
||||
guardrail_name=self.guardrail_name,
|
||||
message=blocking_message,
|
||||
should_wrap_with_default_message=False,
|
||||
)
|
||||
|
||||
def _get_trackers_corrected_message(self, resp: dict) -> Optional[str]:
|
||||
"""Extract corrected/blocking message content from Tracker checkpoint response."""
|
||||
modified = resp.get("modified_data")
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ Unit tests for Ovalix guardrail: config resolution and apply_guardrail behavior
|
|||
with mocked Tracker service responses (allow, anonymize, block).
|
||||
"""
|
||||
import os
|
||||
from typing import Any, List
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -196,7 +197,8 @@ class TestOvalixGuardrail:
|
|||
try:
|
||||
guardrail = OvalixGuardrail(**_guardrail_kwargs())
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[{"role": "user", "content": "how are you?"}]
|
||||
structured_messages=[{"role": "user", "content": "how are you?"}],
|
||||
texts=["how are you?"],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
|
|
@ -232,7 +234,8 @@ class TestOvalixGuardrail:
|
|||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[
|
||||
{"role": "user", "content": "Hello, my name is David."}
|
||||
]
|
||||
],
|
||||
texts=["Hello, my name is David."],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
|
|
@ -266,7 +269,8 @@ class TestOvalixGuardrail:
|
|||
try:
|
||||
guardrail = OvalixGuardrail(**_guardrail_kwargs())
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[{"role": "user", "content": "I am 15 YO"}]
|
||||
structured_messages=[{"role": "user", "content": "I am 15 YO"}],
|
||||
texts=["I am 15 YO"],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
|
|
@ -305,7 +309,8 @@ class TestOvalixGuardrail:
|
|||
structured_messages=[
|
||||
{"role": "user", "content": "I am 15 YO"},
|
||||
{"role": "user", "content": "how are you?"},
|
||||
]
|
||||
],
|
||||
texts=["I am 15 YO", "how are you?"],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
|
|
@ -343,13 +348,18 @@ class TestOvalixGuardrail:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_allow_returns_inputs(self):
|
||||
"""When input_type is response and Tracker allows, apply_guardrail returns inputs unchanged."""
|
||||
"""When input_type is response and Tracker allows, apply_guardrail returns inputs with texts updated from Tracker."""
|
||||
for k, v in _ovalix_env().items():
|
||||
os.environ[k] = v
|
||||
try:
|
||||
guardrail = OvalixGuardrail(**_guardrail_kwargs())
|
||||
inputs = GenericGuardrailAPIInputs()
|
||||
request_data = {"response": None}
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[
|
||||
{"role": "assistant", "content": "Safe assistant reply"}
|
||||
],
|
||||
texts=["Safe assistant reply"],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = TRACKER_RESPONSE_ALLOW
|
||||
|
|
@ -357,11 +367,7 @@ class TestOvalixGuardrail:
|
|||
|
||||
with patch.object(
|
||||
guardrail._async_handler, "post", new_callable=AsyncMock
|
||||
) as mock_post, patch.object(
|
||||
guardrail,
|
||||
"_get_llm_response_text",
|
||||
return_value="Safe assistant reply",
|
||||
):
|
||||
) as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
|
|
@ -370,7 +376,7 @@ class TestOvalixGuardrail:
|
|||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result == inputs
|
||||
assert result.get("texts") == ["how are you?"]
|
||||
assert mock_post.call_count == 1
|
||||
finally:
|
||||
for k in _ovalix_env():
|
||||
|
|
@ -378,74 +384,33 @@ class TestOvalixGuardrail:
|
|||
del os.environ[k]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_block_returns_inputs(
|
||||
self, guardrail_with_env
|
||||
):
|
||||
"""When Tracker blocks on response, apply_guardrail still returns inputs (no raise)."""
|
||||
async def test_apply_guardrail_response_block_raises(self, guardrail_with_env):
|
||||
"""When Tracker blocks on response, apply_guardrail raises OvalixGuardrailBlockedException."""
|
||||
guardrail = guardrail_with_env
|
||||
inputs = GenericGuardrailAPIInputs()
|
||||
request_data = {"response": None}
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[{"role": "user", "content": "I am 15 YO"}],
|
||||
texts=["I am 15 YO"],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = TRACKER_RESPONSE_BLOCK
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
guardrail._async_handler, "post", new_callable=AsyncMock
|
||||
) as mock_post, patch.object(
|
||||
guardrail,
|
||||
"_get_llm_response_text",
|
||||
return_value="I am 15 YO",
|
||||
):
|
||||
mock_post.return_value = mock_response
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result == inputs
|
||||
assert mock_post.call_count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_non_user_messages_not_sent_to_tracker(
|
||||
self, guardrail_with_env
|
||||
):
|
||||
"""Only user messages are sent to Tracker; system/assistant content is passed through."""
|
||||
guardrail = guardrail_with_env
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[
|
||||
{"role": "system", "content": "You are helpful."},
|
||||
{"role": "user", "content": "hello"},
|
||||
]
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
# Tracker allows and returns same content for the user message
|
||||
allow_hello = {
|
||||
"action_type": "allow",
|
||||
"data_type": "TEXT",
|
||||
"original_data": {"content": "hello"},
|
||||
"modified_data": {"content": "hello"},
|
||||
"alerts": [],
|
||||
}
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = allow_hello
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
guardrail._async_handler, "post", new_callable=AsyncMock
|
||||
) as mock_post:
|
||||
mock_post.return_value = mock_response
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
logging_obj=None,
|
||||
)
|
||||
with pytest.raises(OvalixGuardrailBlockedException) as exc_info:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result.get("texts") == ["You are helpful.", "hello"]
|
||||
assert "This message was blocked by Ovalix" in str(exc_info.value.message)
|
||||
assert exc_info.value.status_code == 400
|
||||
assert mock_post.call_count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -455,7 +420,8 @@ class TestOvalixGuardrail:
|
|||
"""When Tracker response has no modified_data.content, original content is used."""
|
||||
guardrail = guardrail_with_env
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[{"role": "user", "content": "original text"}]
|
||||
structured_messages=[{"role": "user", "content": "original text"}],
|
||||
texts=["original text"],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
|
|
@ -490,7 +456,8 @@ class TestOvalixGuardrail:
|
|||
"""When Tracker returns HTTP error (e.g. 400), GuardrailRaisedException is raised."""
|
||||
guardrail = guardrail_with_env
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[{"role": "user", "content": "hello"}]
|
||||
structured_messages=[{"role": "user", "content": "hello"}],
|
||||
texts=["hello"],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
|
|
@ -523,7 +490,8 @@ class TestOvalixGuardrail:
|
|||
try:
|
||||
guardrail = OvalixGuardrail(**_guardrail_kwargs())
|
||||
inputs = GenericGuardrailAPIInputs(
|
||||
structured_messages=[{"role": "user", "content": "hello"}]
|
||||
structured_messages=[{"role": "user", "content": "hello"}],
|
||||
texts=["hello"],
|
||||
)
|
||||
request_data = {}
|
||||
|
||||
|
|
@ -552,7 +520,7 @@ class TestOvalixGuardrail:
|
|||
os.environ[k] = v
|
||||
try:
|
||||
guardrail = OvalixGuardrail(**_guardrail_kwargs())
|
||||
inputs = GenericGuardrailAPIInputs(structured_messages=[])
|
||||
inputs = GenericGuardrailAPIInputs(structured_messages=[], texts=[])
|
||||
request_data = {}
|
||||
|
||||
with patch.object(
|
||||
|
|
@ -619,3 +587,55 @@ class TestOvalixGuardrail:
|
|||
session_id_2 = guardrail._get_session_id(data)
|
||||
assert session_id_1 == session_id_2
|
||||
assert "app-1" in session_id_1
|
||||
|
||||
def test_block_current_message_raises_ovalix_blocked_exception(
|
||||
self, guardrail_with_env
|
||||
):
|
||||
"""_block_current_message raises OvalixGuardrailBlockedException with status_code 400."""
|
||||
guardrail = guardrail_with_env
|
||||
with pytest.raises(OvalixGuardrailBlockedException) as exc_info:
|
||||
guardrail._block_current_message("Custom block reason")
|
||||
assert "Custom block reason" in str(exc_info.value.message)
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
def test_get_trackers_corrected_message(self, guardrail_with_env):
|
||||
"""_get_trackers_corrected_message returns modified_data.content or None."""
|
||||
guardrail = guardrail_with_env
|
||||
assert (
|
||||
guardrail._get_trackers_corrected_message(
|
||||
{"modified_data": {"content": "corrected text"}}
|
||||
)
|
||||
== "corrected text"
|
||||
)
|
||||
assert guardrail._get_trackers_corrected_message({"modified_data": {}}) is None
|
||||
assert (
|
||||
guardrail._get_trackers_corrected_message({"modified_data": "not-a-dict"})
|
||||
is None
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_no_texts_returns_unchanged(self):
|
||||
"""When input_type is response and inputs have no texts, apply_guardrail returns inputs without calling Tracker."""
|
||||
for k, v in _ovalix_env().items():
|
||||
os.environ[k] = v
|
||||
try:
|
||||
guardrail = OvalixGuardrail(**_guardrail_kwargs())
|
||||
inputs = GenericGuardrailAPIInputs()
|
||||
request_data = {}
|
||||
|
||||
with patch.object(
|
||||
guardrail._async_handler, "post", new_callable=AsyncMock
|
||||
) as mock_post:
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=None,
|
||||
)
|
||||
|
||||
assert result == inputs
|
||||
mock_post.assert_not_called()
|
||||
finally:
|
||||
for k in _ovalix_env():
|
||||
if k in os.environ:
|
||||
del os.environ[k]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue