Feat(guardrail): Adding support for custom Ovalix guardrail

This commit is contained in:
Shalom Jamil 2026-02-23 12:58:31 +02:00
parent 9cee51abb9
commit c93a8cc17e
5 changed files with 682 additions and 0 deletions

View file

@ -0,0 +1,46 @@
"""Ovalix guardrail hook: registration and initialization for the proxy."""
from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .ovalix import OvalixGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
"""Create and register an Ovalix guardrail callback from proxy config."""
import litellm
tracker_api_base = getattr(litellm_params, "tracker_api_base", None)
tracker_api_key = getattr(litellm_params, "tracker_api_key", None)
application_id = getattr(litellm_params, "application_id", None)
pre_checkpoint_id = getattr(litellm_params, "pre_checkpoint_id", None)
post_checkpoint_id = getattr(litellm_params, "post_checkpoint_id", None)
_ovalix_callback = OvalixGuardrail(
guardrail_name=guardrail.get("guardrail_name", ""),
tracker_api_base=tracker_api_base,
tracker_api_key=tracker_api_key,
application_id=application_id,
pre_checkpoint_id=pre_checkpoint_id,
post_checkpoint_id=post_checkpoint_id,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(_ovalix_callback)
return _ovalix_callback
# Registry of guardrail name -> initializer for proxy config loading.
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.OVALIX.value: initialize_guardrail,
}
# Registry of guardrail name -> guardrail class (e.g. for apply_guardrail API).
guardrail_class_registry = {
SupportedGuardrailIntegrations.OVALIX.value: OvalixGuardrail,
}

View file

@ -0,0 +1,397 @@
"""Ovalix guardrail integration: pre- and post-call checks via the Tracker service.
Use Ovalix Guardrails for your LLM calls. Supports pre_call (user input) and
post_call (model output) checkpoints with optional correction/blocking.
"""
import datetime
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 (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
BLOCKED_BY_OVALIX_FALLBACK_MESSAGE = "This message was blocked by Ovalix"
BLOCKED_ACTION_TYPE = "block"
USER_MESSAGE_ROLE = "user"
class OvalixGuardrailMissingSecrets(Exception):
"""Raised when required Ovalix config (API base, key, application/checkpoint IDs) is missing."""
pass
class OvalixGuardrailBlockedException(GuardrailRaisedException):
"""
Raised when Ovalix blocks a message. Sets status_code=400 so the proxy
returns 400 and HTTP clients do not retry (they retry on 5xx).
"""
status_code = 400
def __init__(
self,
guardrail_name: Optional[str] = None,
message: str = "",
should_wrap_with_default_message: bool = True,
):
super().__init__(
guardrail_name=guardrail_name,
message=message,
should_wrap_with_default_message=should_wrap_with_default_message,
)
class OvalixGuardrail(CustomGuardrail):
"""
Ovalix guardrail: pre-prompt (pre_call) and post-prompt (post_call) checks
via the Tracker service, with application and checkpoint resolution from the
Monolith backend.
"""
def __init__(
self,
tracker_api_base: Optional[str] = None,
tracker_api_key: Optional[str] = None,
application_id: Optional[str] = None,
pre_checkpoint_id: Optional[str] = None,
post_checkpoint_id: Optional[str] = None,
**kwargs: Any,
):
self._tracker_api_base = tracker_api_base or os.environ.get(
"OVALIX_TRACKER_API_BASE"
)
self._tracker_api_key = tracker_api_key or os.environ.get(
"OVALIX_TRACKER_API_KEY"
)
self._application_id = application_id or os.environ.get("OVALIX_APPLICATION_ID")
self._pre_checkpoint_id = pre_checkpoint_id or os.environ.get(
"OVALIX_PRE_CHECKPOINT_ID"
)
self._post_checkpoint_id = post_checkpoint_id or os.environ.get(
"OVALIX_POST_CHECKPOINT_ID"
)
if "supported_event_hooks" not in kwargs:
kwargs["supported_event_hooks"] = []
self._validate_config(kwargs["supported_event_hooks"])
self._tracker_headers = httpx.Headers(
{
"Authorization": f"Bearer {self._tracker_api_key}",
"Content-Type": "application/json",
},
encoding="utf-8",
)
self._async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
super().__init__(**kwargs)
verbose_proxy_logger.debug(
"Ovalix Guardrail initialized: tracker=%s, application_id=%s, pre_checkpoint_id=%s, post_checkpoint_id=%s",
self._tracker_api_base,
self._application_id,
self._pre_checkpoint_id,
self._post_checkpoint_id,
)
def _validate_config(
self, supported_event_hooks: List[GuardrailEventHooks]
) -> None:
"""Ensure required secrets and checkpoint IDs are set; auto-add hooks when IDs are present."""
if not self._tracker_api_base:
raise OvalixGuardrailMissingSecrets(
"Ovalix Tracker API base required. Set OVALIX_TRACKER_API_BASE or pass tracker_api_base in litellm_params."
)
if not self._tracker_api_key:
raise OvalixGuardrailMissingSecrets(
"Ovalix Tracker API key required. Set OVALIX_TRACKER_API_KEY or pass tracker_api_key in litellm_params."
)
if not self._application_id:
raise OvalixGuardrailMissingSecrets(
"Ovalix Application ID required. Set OVALIX_APPLICATION_ID or pass application_id in litellm_params."
)
if (
not self._pre_checkpoint_id
and GuardrailEventHooks.pre_call in supported_event_hooks
):
raise OvalixGuardrailMissingSecrets(
"Ovalix Pre-checkpoint ID required. Set OVALIX_PRE_CHECKPOINT_ID or pass pre_checkpoint_id in litellm_params."
)
elif (
self._pre_checkpoint_id
and GuardrailEventHooks.pre_call not in supported_event_hooks
):
supported_event_hooks.append(GuardrailEventHooks.pre_call)
if (
not self._post_checkpoint_id
and GuardrailEventHooks.post_call in supported_event_hooks
):
raise OvalixGuardrailMissingSecrets(
"Ovalix Post-checkpoint ID required. Set OVALIX_POST_CHECKPOINT_ID or pass post_checkpoint_id in litellm_params."
)
elif (
self._post_checkpoint_id
and GuardrailEventHooks.post_call not in supported_event_hooks
):
supported_event_hooks.append(GuardrailEventHooks.post_call)
if not self._pre_checkpoint_id and not self._post_checkpoint_id:
raise OvalixGuardrailMissingSecrets(
"Ovalix Pre-checkpoint ID or Post-checkpoint ID required. Set OVALIX_PRE_CHECKPOINT_ID or OVALIX_POST_CHECKPOINT_ID or pass pre_checkpoint_id or post_checkpoint_id in litellm_params."
)
def _get_actor(self, data: dict) -> str:
"""Return a stable actor identifier from request metadata (e.g. user email or id)."""
metadata = data.get("metadata") or data.get("litellm_metadata") or {}
if metadata.get("user_api_key_user_email"):
return metadata["user_api_key_user_email"]
if metadata.get("user_api_key_user_id"):
return metadata["user_api_key_user_id"]
return "unknown"
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))
today = datetime.datetime.now().strftime("%Y-%m-%d")
return f"{actor}_{today}_{self._application_id}"
async def _call_checkpoint(
self,
content: str,
checkpoint_id: str,
actor: str,
session_id: str,
) -> Dict[str, Any]:
"""Call the Ovalix Tracker checkpoint API and return the JSON response."""
application_id = self._application_id
if not application_id or not checkpoint_id:
raise ValueError("Ovalix: application_id or checkpoint_id not resolved")
url = f"{self._tracker_api_base}/tracking/custom_application/checkpoint"
headers = dict(self._tracker_headers)
payload = {
"application_id": application_id,
"checkpoint_id": checkpoint_id,
"actor": actor,
"session_id": session_id,
"data_type": "TEXT",
"data": {"content": content},
}
response = await self._async_handler.post(
url, headers=headers, json=payload
)
response.raise_for_status()
return response.json()
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional[Any] = None,
) -> GenericGuardrailAPIInputs:
"""
Apply Ovalix guardrail to the given inputs (request or response text).
Used by the unified guardrail flow and the /apply_guardrail API.
For "request", uses the pre-checkpoint; for "response", uses the post-checkpoint.
Args:
inputs: Guardrail API inputs (e.g. texts to check).
request_data: Full request payload (messages, metadata, response).
input_type: "request" (pre_call) or "response" (post_call).
logging_obj: Optional logging context.
Returns:
Updated inputs (e.g. with replaced/corrected texts, or unchanged).
"""
if not self._pre_checkpoint_id and not self._post_checkpoint_id:
return inputs
actor = self._get_actor(request_data)
session_id = self._get_session_id(request_data)
if input_type == "response":
llm_response = self._get_llm_response_text(
request_data.get("response", None)
)
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 = request_data.get("messages") or []
if not messages:
return inputs
if self._pre_checkpoint_id:
post_guardrail_texts = await self._generate_post_guardrail_text(
messages, actor, session_id, request_data
)
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[Dict[str, Any]],
actor: str,
session_id: str,
request_data: dict,
) -> 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
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 LLM/system/developer past responses, only the response that the user sent
post_guardrail_texts.insert(0, content)
continue
try:
resp = await self._call_checkpoint(
content, self._pre_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
)
if is_last_prompt:
self._block_current_message(blocking_message)
else:
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
return post_guardrail_texts
def _get_trackers_corrected_message(self, resp: dict) -> Optional[str]:
"""Extract corrected/blocking message content from Tracker checkpoint response."""
modified = resp.get("modified_data")
if isinstance(modified, dict) and "content" in modified:
return modified["content"]
return None
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
return OvalixGuardrailConfigModel

View file

@ -17,6 +17,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
ContentFilterCategoryConfig,
)
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import (
QualifireGuardrailConfigModel,
)
@ -69,6 +72,7 @@ class SupportedGuardrailIntegrations(Enum):
GENERIC_GUARDRAIL_API = "generic_guardrail_api"
QUALIFIRE = "qualifire"
CUSTOM_CODE = "custom_code"
OVALIX = "ovalix"
class Role(Enum):
@ -681,6 +685,7 @@ class LitellmParams(
BaseLitellmParams,
EnkryptAIGuardrailConfigs,
IBMGuardrailsBaseConfigModel,
OvalixGuardrailConfigModel,
QualifireGuardrailConfigModel,
):
guardrail: str = Field(description="The type of guardrail integration to use")

View file

@ -0,0 +1,37 @@
"""Pydantic config model for the Ovalix guardrail (Tracker API, application and checkpoint IDs)."""
from typing import Optional
from pydantic import Field
from .base import GuardrailConfigModel
class OvalixGuardrailConfigModel(GuardrailConfigModel):
"""Configuration parameters for the Ovalix guardrail (pre/post call checkpoints)."""
tracker_api_base: Optional[str] = Field(
default=None,
description="Base URL for the Ovalix Tracker service.",
)
tracker_api_key: Optional[str] = Field(
default=None,
description="API key for the Ovalix Tracker service.",
)
application_id: Optional[str] = Field(
default=None,
description="Application ID for the Ovalix Tracker service.",
)
pre_checkpoint_id: Optional[str] = Field(
default=None,
description="Pre-checkpoint ID for the Ovalix Tracker service.",
)
post_checkpoint_id: Optional[str] = Field(
default=None,
description="Post-checkpoint ID for the Ovalix Tracker service.",
)
@staticmethod
def ui_friendly_name() -> str:
"""Display name for this guardrail in the proxy UI."""
return "Ovalix Guardrail"

View file

@ -0,0 +1,197 @@
"""
Unit tests for Ovalix guardrail types (OvalixGuardrailConfigModel) and config model resolution.
"""
from litellm.proxy.guardrails.guardrail_hooks.ovalix.ovalix import OvalixGuardrail
class TestOvalixGuardrailConfigModel:
"""Tests for OvalixGuardrailConfigModel from litellm.types.proxy.guardrails.guardrail_hooks.ovalix."""
def test_config_model_ui_friendly_name(self):
"""Test that config model has correct UI friendly name."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
assert OvalixGuardrailConfigModel.ui_friendly_name() == "Ovalix Guardrail"
def test_config_model_fields(self):
"""Test that config model has expected fields and default values."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
model = OvalixGuardrailConfigModel()
assert model.tracker_api_base is None
assert model.tracker_api_key is None
assert model.application_id is None
assert model.pre_checkpoint_id is None
assert model.post_checkpoint_id is None
def test_get_config_model(self):
"""Test get_config_model returns OvalixGuardrailConfigModel."""
config_model = OvalixGuardrail.get_config_model()
assert config_model is not None
assert config_model.__name__ == "OvalixGuardrailConfigModel"
assert hasattr(config_model, "ui_friendly_name")
assert config_model.ui_friendly_name() == "Ovalix Guardrail"
def test_config_model_with_all_fields_set(self):
"""Test construction with all optional fields set to explicit values."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
model = OvalixGuardrailConfigModel(
tracker_api_base="https://tracker.ovalix.example",
tracker_api_key="key-123",
application_id="app-456",
pre_checkpoint_id="pre-cp-1",
post_checkpoint_id="post-cp-1",
)
assert model.tracker_api_base == "https://tracker.ovalix.example"
assert model.tracker_api_key == "key-123"
assert model.application_id == "app-456"
assert model.pre_checkpoint_id == "pre-cp-1"
assert model.post_checkpoint_id == "post-cp-1"
def test_config_model_with_partial_fields(self):
"""Test construction with only a subset of fields set."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
model = OvalixGuardrailConfigModel(
tracker_api_base="https://custom.tracker",
application_id="app-only",
)
assert model.tracker_api_base == "https://custom.tracker"
assert model.application_id == "app-only"
assert model.tracker_api_key is None
assert model.pre_checkpoint_id is None
assert model.post_checkpoint_id is None
def test_config_model_inherits_base_optional_params(self):
"""Test that model has optional_params from GuardrailConfigModel and defaults to None."""
from litellm.types.proxy.guardrails.guardrail_hooks.base import (
GuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
model = OvalixGuardrailConfigModel()
assert hasattr(model, "optional_params")
assert model.optional_params is None
assert isinstance(model, GuardrailConfigModel)
def test_config_model_serialization_dump(self):
"""Test model_dump produces expected keys and values."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
model = OvalixGuardrailConfigModel(
tracker_api_base="https://tracker.test",
pre_checkpoint_id="pre-1",
)
data = model.model_dump()
assert "tracker_api_base" in data
assert "tracker_api_key" in data
assert "application_id" in data
assert "pre_checkpoint_id" in data
assert "post_checkpoint_id" in data
assert data["tracker_api_base"] == "https://tracker.test"
assert data["pre_checkpoint_id"] == "pre-1"
assert data["tracker_api_key"] is None
def test_config_model_deserialization_from_dict(self):
"""Test model_validate builds instance from dict."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
payload = {
"tracker_api_base": "https://from-dict.example",
"tracker_api_key": "secret",
"application_id": "app-dict",
"pre_checkpoint_id": "pre-d",
"post_checkpoint_id": "post-d",
}
model = OvalixGuardrailConfigModel.model_validate(payload)
assert model.tracker_api_base == "https://from-dict.example"
assert model.tracker_api_key == "secret"
assert model.application_id == "app-dict"
assert model.pre_checkpoint_id == "pre-d"
assert model.post_checkpoint_id == "post-d"
def test_config_model_deserialization_empty_dict(self):
"""Test model_validate with empty dict yields defaults."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
model = OvalixGuardrailConfigModel.model_validate({})
assert model.tracker_api_base is None
assert model.tracker_api_key is None
assert model.application_id is None
assert model.pre_checkpoint_id is None
assert model.post_checkpoint_id is None
def test_config_model_round_trip(self):
"""Test model_dump then model_validate preserves data."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
original = OvalixGuardrailConfigModel(
tracker_api_base="https://round.trip",
post_checkpoint_id="post-rt",
)
data = original.model_dump()
restored = OvalixGuardrailConfigModel.model_validate(data)
assert restored.tracker_api_base == original.tracker_api_base
assert restored.tracker_api_key == original.tracker_api_key
assert restored.application_id == original.application_id
assert restored.pre_checkpoint_id == original.pre_checkpoint_id
assert restored.post_checkpoint_id == original.post_checkpoint_id
def test_config_model_has_expected_field_names(self):
"""Test that model defines all expected Ovalix config field names."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
expected = {
"tracker_api_base",
"tracker_api_key",
"application_id",
"pre_checkpoint_id",
"post_checkpoint_id",
}
assert expected.issubset(OvalixGuardrailConfigModel.model_fields.keys())
def test_config_model_field_descriptions_present(self):
"""Test that Ovalix-specific fields have non-empty descriptions."""
from litellm.types.proxy.guardrails.guardrail_hooks.ovalix import (
OvalixGuardrailConfigModel,
)
ovalix_fields = [
"tracker_api_base",
"tracker_api_key",
"application_id",
"pre_checkpoint_id",
"post_checkpoint_id",
]
for name in ovalix_fields:
assert name in OvalixGuardrailConfigModel.model_fields
desc = OvalixGuardrailConfigModel.model_fields[name].description
assert desc is not None and len(desc) > 0