greptileai comments fixes

This commit is contained in:
Shalom Jamil 2026-02-23 12:59:55 +02:00
parent 3718682021
commit f3815406c0
2 changed files with 134 additions and 194 deletions

View file

@ -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")

View file

@ -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]