mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Feat(guardrail): Adding support for custom Ovalix guardrail
This commit is contained in:
parent
9cee51abb9
commit
c93a8cc17e
5 changed files with 682 additions and 0 deletions
46
litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py
Normal file
46
litellm/proxy/guardrails/guardrail_hooks/ovalix/__init__.py
Normal 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,
|
||||
}
|
||||
397
litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py
Normal file
397
litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
37
litellm/types/proxy/guardrails/guardrail_hooks/ovalix.py
Normal file
37
litellm/types/proxy/guardrails/guardrail_hooks/ovalix.py
Normal 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"
|
||||
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue