mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
[Feat] Guardrails - add logging for important status fields (#15090)
* add StandardLoggingPayloadStatusFields * add status_fields * add StandardLoggingPayloadStatusFields * noma guard: add_standard_logging_guardrail_information_to_request_data * fix: StandardLoggingPayloadStatusFields * fix tests * fix StandardLoggingPayloadStatus * get_standard_logging_object_payload * test_bedrock_guardrail_status_failure * fix: _get_status_fields * fixes new guardrail tracing * fix ruff
This commit is contained in:
parent
f205b2c0a5
commit
0ca11eefde
13 changed files with 898 additions and 44 deletions
|
|
@ -14,6 +14,7 @@ Found under `kwargs["standard_logging_object"]`. This is a standard payload, log
|
|||
| `cost_breakdown` | `Optional[CostBreakdown]` | Detailed cost breakdown object |
|
||||
| `response_cost_failure_debug_info` | `StandardLoggingModelCostFailureDebugInformation` | Debug information if cost tracking fails |
|
||||
| `status` | `StandardLoggingPayloadStatus` | Status of the payload |
|
||||
| `status_fields` | `StandardLoggingPayloadStatusFields` | Typed status fields for easy filtering and analytics |
|
||||
| `total_tokens` | `int` | Total number of tokens |
|
||||
| `prompt_tokens` | `int` | Number of prompt tokens |
|
||||
| `completion_tokens` | `int` | Number of completion tokens |
|
||||
|
|
@ -168,12 +169,83 @@ A literal type with two possible values:
|
|||
| `guardrail_mode` | `Optional[Union[GuardrailEventHooks, List[GuardrailEventHooks]]]` | Guardrail mode |
|
||||
| `guardrail_request` | `Optional[dict]` | Guardrail request |
|
||||
| `guardrail_response` | `Optional[Union[dict, str, List[dict]]]` | Guardrail response |
|
||||
| `guardrail_status` | `Literal["success", "failure"]` | Guardrail status |
|
||||
| `guardrail_status` | `Literal["success", "failure", "blocked"]` | Guardrail execution status: `success` = no violations detected, `blocked` = content blocked/modified due to policy violations, `failure` = technical error or API failure |
|
||||
| `start_time` | `Optional[float]` | Start time of the guardrail |
|
||||
| `end_time` | `Optional[float]` | End time of the guardrail |
|
||||
| `duration` | `Optional[float]` | Duration of the guardrail in seconds |
|
||||
| `masked_entity_count` | `Optional[Dict[str, int]]` | Count of masked entities |
|
||||
|
||||
## StandardLoggingPayloadStatusFields
|
||||
|
||||
Typed status fields for easy filtering and analytics.
|
||||
|
||||
| Field | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| `llm_api_status` | `StandardLoggingPayloadStatus` | Status of the LLM API call: `"success"` if completed successfully, `"failure"` if errored |
|
||||
| `guardrail_status` | `GuardrailStatus` | Status of guardrail execution (see below) |
|
||||
|
||||
### StandardLoggingPayloadStatus
|
||||
|
||||
A literal type with two possible values:
|
||||
- `"success"` - The LLM API request completed successfully
|
||||
- `"failure"` - The LLM API request failed
|
||||
|
||||
### GuardrailStatus
|
||||
|
||||
A literal type with four possible values:
|
||||
- `"success"` - Guardrail ran and allowed content through (no violations detected)
|
||||
- `"guardrail_intervened"` - Guardrail blocked or modified content due to policy violations
|
||||
- `"guardrail_failed_to_respond"` - Guardrail had a technical failure or API error
|
||||
- `"not_run"` - No guardrail was executed for this request
|
||||
|
||||
### Usage Examples
|
||||
|
||||
Filter logs for requests where guardrails intervened:
|
||||
```json
|
||||
{
|
||||
"status_fields": {
|
||||
"guardrail_status": "guardrail_intervened"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Find guardrail technical failures:
|
||||
```json
|
||||
{
|
||||
"status_fields": {
|
||||
"guardrail_status": "guardrail_failed_to_respond"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Get successful LLM requests:
|
||||
```json
|
||||
{
|
||||
"status_fields": {
|
||||
"llm_api_status": "success"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Find requests where guardrails ran successfully without intervention:
|
||||
```json
|
||||
{
|
||||
"status_fields": {
|
||||
"guardrail_status": "success",
|
||||
"llm_api_status": "success"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Find requests where no guardrail was run:
|
||||
```json
|
||||
{
|
||||
"status_fields": {
|
||||
"guardrail_status": "not_run"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## StandardLoggingPromptManagementMetadata
|
||||
|
||||
Used for tracking prompt versioning and management information.
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Literal, Optional, Type, Union, get_args
|
||||
from typing import Any, Dict, List, Optional, Type, Union, get_args
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.caching import DualCache
|
||||
|
|
@ -14,6 +14,7 @@ from litellm.types.guardrails import (
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
GuardrailStatus,
|
||||
LLMResponseTypes,
|
||||
StandardLoggingGuardrailInformation,
|
||||
)
|
||||
|
|
@ -352,7 +353,7 @@ class CustomGuardrail(CustomLogger):
|
|||
self,
|
||||
guardrail_json_response: Union[Exception, str, dict, List[dict]],
|
||||
request_data: dict,
|
||||
guardrail_status: Literal["success", "failure", "blocked"],
|
||||
guardrail_status: GuardrailStatus,
|
||||
start_time: Optional[float] = None,
|
||||
end_time: Optional[float] = None,
|
||||
duration: Optional[float] = None,
|
||||
|
|
@ -460,7 +461,7 @@ class CustomGuardrail(CustomLogger):
|
|||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response=e,
|
||||
request_data=request_data,
|
||||
guardrail_status="failure",
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
duration=duration,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
|
|
|
|||
|
|
@ -89,6 +89,7 @@ from litellm.types.utils import (
|
|||
CostResponseTypes,
|
||||
DynamicPromptManagementParamLiteral,
|
||||
EmbeddingResponse,
|
||||
GuardrailStatus,
|
||||
ImageResponse,
|
||||
LiteLLMBatch,
|
||||
LiteLLMLoggingBaseClass,
|
||||
|
|
@ -107,6 +108,7 @@ from litellm.types.utils import (
|
|||
StandardLoggingPayload,
|
||||
StandardLoggingPayloadErrorInformation,
|
||||
StandardLoggingPayloadStatus,
|
||||
StandardLoggingPayloadStatusFields,
|
||||
StandardLoggingPromptManagementMetadata,
|
||||
StandardLoggingVectorStoreRequest,
|
||||
TextCompletionResponse,
|
||||
|
|
@ -4425,6 +4427,51 @@ class StandardLoggingPayloadSetup:
|
|||
return request_tags
|
||||
|
||||
|
||||
|
||||
def _get_status_fields(
|
||||
status: StandardLoggingPayloadStatus,
|
||||
guardrail_information: Optional[dict],
|
||||
error_str: Optional[str]
|
||||
) -> "StandardLoggingPayloadStatusFields":
|
||||
"""
|
||||
Determine status fields based on request status and guardrail information.
|
||||
|
||||
Args:
|
||||
status: Overall request status ("success" or "failure")
|
||||
guardrail_information: Guardrail information from metadata
|
||||
error_str: Error string if any
|
||||
|
||||
Returns:
|
||||
StandardLoggingPayloadStatusFields with llm_api_status and guardrail_status
|
||||
"""
|
||||
# Mapping for legacy guardrail status values to new GuardrailStatus values
|
||||
GUARDRAIL_STATUS_MAP: Dict[str, GuardrailStatus] = {
|
||||
"success": "success",
|
||||
"blocked": "guardrail_intervened", # legacy
|
||||
"guardrail_intervened": "guardrail_intervened", # direct
|
||||
"failure": "guardrail_failed_to_respond", # legacy
|
||||
"guardrail_failed_to_respond": "guardrail_failed_to_respond", # direct
|
||||
"not_run": "not_run"
|
||||
}
|
||||
|
||||
# Set LLM API status
|
||||
llm_api_status: StandardLoggingPayloadStatus = status
|
||||
|
||||
|
||||
#########################################################
|
||||
# Map - guardrail_information.guardrail_status to guardrail_status
|
||||
#########################################################
|
||||
guardrail_status: GuardrailStatus = "not_run"
|
||||
if guardrail_information and isinstance(guardrail_information, dict):
|
||||
raw_status = guardrail_information.get("guardrail_status", "not_run")
|
||||
guardrail_status = GUARDRAIL_STATUS_MAP.get(raw_status, "not_run")
|
||||
|
||||
return StandardLoggingPayloadStatusFields(
|
||||
llm_api_status=llm_api_status,
|
||||
guardrail_status=guardrail_status
|
||||
)
|
||||
|
||||
|
||||
def get_standard_logging_object_payload(
|
||||
kwargs: Optional[dict],
|
||||
init_response_obj: Union[Any, BaseModel, dict],
|
||||
|
|
@ -4534,7 +4581,6 @@ def get_standard_logging_object_payload(
|
|||
start_time=start_time,
|
||||
response_id=id,
|
||||
)
|
||||
|
||||
_request_body = proxy_server_request.get("body", {})
|
||||
end_user_id = clean_metadata["user_api_key_end_user_id"] or _request_body.get(
|
||||
"user", None
|
||||
|
|
@ -4590,6 +4636,11 @@ def get_standard_logging_object_payload(
|
|||
cache_hit=cache_hit,
|
||||
stream=stream,
|
||||
status=status,
|
||||
status_fields=_get_status_fields(
|
||||
status=status,
|
||||
guardrail_information=metadata.get("standard_logging_guardrail_information", None),
|
||||
error_str=error_str
|
||||
),
|
||||
custom_llm_provider=cast(Optional[str], kwargs.get("custom_llm_provider")),
|
||||
saved_cache_cost=saved_cache_cost,
|
||||
startTime=start_time_float,
|
||||
|
|
|
|||
|
|
@ -41,6 +41,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
|
|||
)
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
GuardrailStatus,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
|
|
@ -361,11 +362,30 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
prepared_request.headers,
|
||||
)
|
||||
|
||||
httpx_response = await self.async_handler.post(
|
||||
url=prepared_request.url,
|
||||
data=prepared_request.body, # type: ignore
|
||||
headers=prepared_request.headers, # type: ignore
|
||||
)
|
||||
try:
|
||||
httpx_response = await self.async_handler.post(
|
||||
url=prepared_request.url,
|
||||
data=prepared_request.body, # type: ignore
|
||||
headers=prepared_request.headers, # type: ignore
|
||||
)
|
||||
except Exception as e:
|
||||
# Endpoint down, timeout, or other HTTP/network errors
|
||||
verbose_proxy_logger.error(
|
||||
"Bedrock AI: failed to make guardrail request: %s", str(e)
|
||||
)
|
||||
# Add guardrail information with failure status
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_provider=self.guardrail_provider,
|
||||
guardrail_json_response={"error": str(e)},
|
||||
request_data=request_data or {},
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=datetime.now().timestamp(),
|
||||
duration=(datetime.now() - start_time).total_seconds(),
|
||||
)
|
||||
# Re-raise the exception to maintain existing behavior
|
||||
raise
|
||||
|
||||
#########################################################
|
||||
# Add guardrail information to request trace
|
||||
#########################################################
|
||||
|
|
@ -437,15 +457,30 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
|
||||
def _get_bedrock_guardrail_response_status(
|
||||
self, response: httpx.Response
|
||||
) -> Literal["success", "failure"]:
|
||||
) -> GuardrailStatus:
|
||||
"""
|
||||
Get the status of the bedrock guardrail response.
|
||||
|
||||
Returns:
|
||||
"success": Content allowed through with no violations
|
||||
"guardrail_intervened": Content blocked due to policy violations
|
||||
"guardrail_failed_to_respond": Technical error or API failure
|
||||
"""
|
||||
if response.status_code == 200:
|
||||
if self._check_bedrock_response_for_exception(response):
|
||||
return "failure"
|
||||
return "guardrail_failed_to_respond"
|
||||
|
||||
# Check if the guardrail would block content
|
||||
try:
|
||||
_json_response = response.json()
|
||||
bedrock_guardrail_response = BedrockGuardrailResponse(**_json_response)
|
||||
if self._should_raise_guardrail_blocked_exception(bedrock_guardrail_response):
|
||||
return "guardrail_intervened"
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return "success"
|
||||
return "failure"
|
||||
return "guardrail_failed_to_respond"
|
||||
|
||||
def _get_http_exception_for_blocked_guardrail(
|
||||
self, response: BedrockGuardrailResponse
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Union, Type
|
||||
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Type, Union
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -12,11 +14,11 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.javelin import (
|
||||
JavelinGuardInput,
|
||||
JavelinGuardRequest,
|
||||
JavelinGuardResponse,
|
||||
JavelinGuardInput,
|
||||
)
|
||||
from fastapi import HTTPException
|
||||
from litellm.types.utils import GuardrailStatus
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
|
@ -95,7 +97,7 @@ class JavelinGuardrail(CustomGuardrail):
|
|||
if self.application:
|
||||
headers["x-javelin-application"] = self.application
|
||||
|
||||
status: Literal["success", "failure", "blocked"] = "failure"
|
||||
status: GuardrailStatus = "guardrail_failed_to_respond"
|
||||
javelin_response: Optional[JavelinGuardResponse] = None
|
||||
exception_str = ""
|
||||
|
||||
|
|
@ -122,7 +124,7 @@ class JavelinGuardrail(CustomGuardrail):
|
|||
status = "success"
|
||||
return javelin_response
|
||||
except Exception as e:
|
||||
status = "failure"
|
||||
status = "guardrail_failed_to_respond"
|
||||
exception_str = str(e)
|
||||
return {"assessments": []}
|
||||
finally:
|
||||
|
|
@ -178,12 +180,12 @@ class JavelinGuardrail(CustomGuardrail):
|
|||
"""
|
||||
Pre-call hook for the Javelin guardrail.
|
||||
"""
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_last_user_message,
|
||||
)
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("Javelin Guardrail: pre_call_hook")
|
||||
verbose_proxy_logger.debug("Javelin Guardrail: Request data: %s", data)
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import (
|
|||
LakeraAIRequest,
|
||||
LakeraAIResponse,
|
||||
)
|
||||
from litellm.types.utils import GuardrailStatus
|
||||
|
||||
|
||||
class LakeraAIGuardrail(CustomGuardrail):
|
||||
|
|
@ -70,7 +71,7 @@ class LakeraAIGuardrail(CustomGuardrail):
|
|||
"""
|
||||
Call the Lakera AI v2 guard API.
|
||||
"""
|
||||
status: Literal["success", "failure"] = "success"
|
||||
status: GuardrailStatus = "success"
|
||||
exception_str: str = ""
|
||||
start_time: datetime = datetime.now()
|
||||
lakera_response: Optional[LakeraAIResponse] = None
|
||||
|
|
@ -99,7 +100,7 @@ class LakeraAIGuardrail(CustomGuardrail):
|
|||
lakera_response = LakeraAIResponse(**response.json())
|
||||
return lakera_response, masked_entity_count
|
||||
except Exception as e:
|
||||
status = "failure"
|
||||
status = "guardrail_failed_to_respond"
|
||||
exception_str = str(e)
|
||||
raise e
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
GuardrailStatus,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
)
|
||||
|
|
@ -329,14 +330,14 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
guardrail_response = metadata.get("_model_armor_response", {})
|
||||
|
||||
# Determine status – default to "success" but prefer the explicit value if present.
|
||||
guardrail_status: Literal["success", "failure", "blocked"] = metadata.get(
|
||||
guardrail_status: GuardrailStatus = metadata.get(
|
||||
"_model_armor_status", "success"
|
||||
) # type: ignore
|
||||
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_json_response=guardrail_response,
|
||||
request_data=request_data,
|
||||
guardrail_status=guardrail_status, # type: ignore
|
||||
guardrail_status=guardrail_status,
|
||||
duration=duration,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
|
|
|
|||
|
|
@ -8,7 +8,8 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import os
|
||||
from typing import Any, Dict, Final, Literal, Optional, Union, Type, TYPE_CHECKING
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, Final, Literal, Optional, Type, Union
|
||||
from urllib.parse import urljoin
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -23,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import EmbeddingResponse, ImageResponse
|
||||
from litellm.types.utils import EmbeddingResponse, GuardrailStatus, ImageResponse
|
||||
|
||||
# Constants
|
||||
USER_ROLE: Final[Literal["user"]] = "user"
|
||||
|
|
@ -204,6 +205,7 @@ class NomaGuardrail(CustomGuardrail):
|
|||
user_auth: UserAPIKeyAuth,
|
||||
) -> Optional[str]:
|
||||
"""Shared logic for processing user message checks"""
|
||||
start_time = datetime.now()
|
||||
extra_data = self.get_guardrail_dynamic_request_body_params(request_data)
|
||||
|
||||
user_message = await self._extract_user_message(request_data)
|
||||
|
|
@ -218,6 +220,23 @@ class NomaGuardrail(CustomGuardrail):
|
|||
user_auth=user_auth,
|
||||
extra_data=extra_data,
|
||||
)
|
||||
|
||||
end_time = datetime.now()
|
||||
duration = (end_time - start_time).total_seconds()
|
||||
|
||||
# Determine guardrail status based on response
|
||||
guardrail_status = self._determine_guardrail_status(response_json)
|
||||
|
||||
# Always log guardrail information for consistency
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_provider="noma",
|
||||
guardrail_json_response=response_json,
|
||||
request_data=request_data,
|
||||
guardrail_status=guardrail_status,
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=end_time.timestamp(),
|
||||
duration=duration,
|
||||
)
|
||||
|
||||
if self.monitor_mode:
|
||||
await self._handle_verdict_background(
|
||||
|
|
@ -248,6 +267,8 @@ class NomaGuardrail(CustomGuardrail):
|
|||
user_auth: UserAPIKeyAuth,
|
||||
) -> Optional[str]:
|
||||
"""Shared logic for processing LLM response checks"""
|
||||
|
||||
start_time = datetime.now()
|
||||
extra_data = self.get_guardrail_dynamic_request_body_params(request_data)
|
||||
|
||||
if not isinstance(response, litellm.ModelResponse):
|
||||
|
|
@ -271,6 +292,23 @@ class NomaGuardrail(CustomGuardrail):
|
|||
user_auth=user_auth,
|
||||
extra_data=extra_data,
|
||||
)
|
||||
|
||||
end_time = datetime.now()
|
||||
duration = (end_time - start_time).total_seconds()
|
||||
|
||||
# Determine guardrail status based on response
|
||||
guardrail_status = self._determine_guardrail_status(response_json)
|
||||
|
||||
# Always log guardrail information for consistency
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_provider="noma",
|
||||
guardrail_json_response=response_json,
|
||||
request_data=request_data,
|
||||
guardrail_status=guardrail_status,
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=end_time.timestamp(),
|
||||
duration=duration,
|
||||
)
|
||||
|
||||
if self.monitor_mode:
|
||||
await self._handle_verdict_background(
|
||||
|
|
@ -294,6 +332,41 @@ class NomaGuardrail(CustomGuardrail):
|
|||
await self._check_verdict(ASSISTANT_ROLE, content, response_json)
|
||||
return content
|
||||
|
||||
def _determine_guardrail_status(self, response_json: dict) -> GuardrailStatus:
|
||||
"""
|
||||
Determine the guardrail status based on NOMA API response.
|
||||
|
||||
Args:
|
||||
response_json: Response from NOMA API
|
||||
|
||||
Returns:
|
||||
"success": Content allowed through with no violations
|
||||
"guardrail_intervened": Content blocked due to policy violations
|
||||
"guardrail_failed_to_respond": Technical error or API failure
|
||||
"""
|
||||
try:
|
||||
# Check if we got a valid response structure
|
||||
if not isinstance(response_json, dict):
|
||||
return "guardrail_failed_to_respond"
|
||||
|
||||
# Get the verdict from the response
|
||||
verdict = response_json.get("verdict", True)
|
||||
|
||||
# If verdict is True, content is allowed
|
||||
if verdict is True:
|
||||
return "success"
|
||||
|
||||
# If verdict is False, content is blocked/flagged
|
||||
if verdict is False:
|
||||
return "guardrail_intervened"
|
||||
|
||||
# If verdict is missing or invalid, treat as failure
|
||||
return "guardrail_failed_to_respond"
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error determining NOMA guardrail status: {str(e)}")
|
||||
return "guardrail_failed_to_respond"
|
||||
|
||||
def _should_only_sensitive_data_failed(self, classification_obj: dict) -> bool:
|
||||
"""
|
||||
Check if only sensitive data detectors (PII, PCI, secrets) have result=true in the classification.
|
||||
|
|
@ -539,8 +612,22 @@ class NomaGuardrail(CustomGuardrail):
|
|||
try:
|
||||
return await self._check_user_message(data, user_api_key_dict)
|
||||
except NomaBlockedMessage:
|
||||
# Blocked requests were already logged in _process_user_message_check with "blocked" status
|
||||
raise
|
||||
except Exception as e:
|
||||
# Log technical failures
|
||||
from datetime import datetime
|
||||
start_time = datetime.now()
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_provider="noma",
|
||||
guardrail_json_response=str(e),
|
||||
request_data=data,
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=start_time.timestamp(),
|
||||
duration=0.0,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.error(f"Noma pre-call hook failed: {str(e)}")
|
||||
|
||||
if self.block_failures:
|
||||
|
|
@ -580,8 +667,22 @@ class NomaGuardrail(CustomGuardrail):
|
|||
try:
|
||||
return await self._check_user_message(data, user_api_key_dict)
|
||||
except NomaBlockedMessage:
|
||||
# Blocked requests were already logged in _process_user_message_check with "blocked" status
|
||||
raise
|
||||
except Exception as e:
|
||||
# Log technical failures
|
||||
from datetime import datetime
|
||||
start_time = datetime.now()
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_provider="noma",
|
||||
guardrail_json_response=str(e),
|
||||
request_data=data,
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=start_time.timestamp(),
|
||||
duration=0.0,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.error(f"Noma moderation hook failed: {str(e)}")
|
||||
|
||||
if self.block_failures:
|
||||
|
|
@ -615,8 +716,22 @@ class NomaGuardrail(CustomGuardrail):
|
|||
try:
|
||||
return await self._check_llm_response(data, response, user_api_key_dict)
|
||||
except NomaBlockedMessage:
|
||||
# Blocked requests were already logged in _process_llm_response_check with "blocked" status
|
||||
raise
|
||||
except Exception as e:
|
||||
# Log technical failures
|
||||
from datetime import datetime
|
||||
start_time = datetime.now()
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
guardrail_provider="noma",
|
||||
guardrail_json_response=str(e),
|
||||
request_data=data,
|
||||
guardrail_status="guardrail_failed_to_respond",
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=start_time.timestamp(),
|
||||
duration=0.0,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.error(f"Noma post-call hook failed: {str(e)}")
|
||||
if self.block_failures:
|
||||
raise
|
||||
|
|
|
|||
|
|
@ -10,14 +10,12 @@
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
from litellm._uuid import uuid
|
||||
from datetime import datetime
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
|
|
@ -29,6 +27,7 @@ import aiohttp
|
|||
import litellm # noqa: E401
|
||||
from litellm import get_secret
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.exceptions import BlockedPiiEntityError
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -45,6 +44,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.presidio import (
|
|||
PresidioAnalyzeResponseItem,
|
||||
)
|
||||
from litellm.types.utils import CallTypes as LitellmCallTypes
|
||||
from litellm.types.utils import GuardrailStatus
|
||||
from litellm.utils import (
|
||||
EmbeddingResponse,
|
||||
ImageResponse,
|
||||
|
|
@ -324,7 +324,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
"""
|
||||
start_time = datetime.now()
|
||||
analyze_results: Optional[Union[List[PresidioAnalyzeResponseItem], Dict]] = None
|
||||
status: Literal["success", "failure"] = "success"
|
||||
status: GuardrailStatus = "success"
|
||||
masked_entity_count: Dict[str, int] = {}
|
||||
exception_str: str = ""
|
||||
try:
|
||||
|
|
@ -356,7 +356,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
)
|
||||
return redacted_text["text"]
|
||||
except Exception as e:
|
||||
status = "failure"
|
||||
status = "guardrail_failed_to_respond"
|
||||
exception_str = str(e)
|
||||
raise e
|
||||
finally:
|
||||
|
|
|
|||
|
|
@ -26,19 +26,28 @@ model_list:
|
|||
- model_name: vertex_ai/*
|
||||
litellm_params:
|
||||
model: vertex_ai/*
|
||||
- model_name: "grok-4"
|
||||
model_info:
|
||||
mode: completion
|
||||
litellm_params:
|
||||
model: oci/xai.grok-4
|
||||
oci_key: ocid1.api_key.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk
|
||||
oci_region: us-phoenix-1
|
||||
oci_user: ocid1.user.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk
|
||||
oci_fingerprint: aa:bb:cc:dd:ee:ff:11:22:33:44:55:66:77:88:99:00
|
||||
oci_tenancy: ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk
|
||||
oci_key_file: /path/to/oci_api_key.pem
|
||||
oci_compartment_id: ocid1.compartment.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk
|
||||
drop_params: True
|
||||
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: lakera
|
||||
- guardrail_name: "bedrock-pre-guard"
|
||||
litellm_params:
|
||||
guardrail: lakera_v2
|
||||
mode: pre_call
|
||||
api_key: os.environ/LAKERA_API_KEY
|
||||
default_on: false
|
||||
project_id: project-9770817088
|
||||
breakdown: true
|
||||
payload: true
|
||||
dev_info: true
|
||||
guardrail: bedrock # supported values: "aporia", "bedrock", "lakera"
|
||||
mode: "during_call"
|
||||
guardrailIdentifier: ff6ujrregl1q
|
||||
guardrailVersion: "DRAFT"
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["datadog"]
|
||||
|
|
|
|||
|
|
@ -2031,6 +2031,13 @@ class GuardrailMode(TypedDict, total=False):
|
|||
default: Optional[str]
|
||||
|
||||
|
||||
GuardrailStatus = Literal[
|
||||
"success",
|
||||
"guardrail_intervened",
|
||||
"guardrail_failed_to_respond",
|
||||
"not_run"
|
||||
]
|
||||
|
||||
class StandardLoggingGuardrailInformation(TypedDict, total=False):
|
||||
guardrail_name: Optional[str]
|
||||
guardrail_provider: Optional[str]
|
||||
|
|
@ -2039,7 +2046,7 @@ class StandardLoggingGuardrailInformation(TypedDict, total=False):
|
|||
]
|
||||
guardrail_request: Optional[dict]
|
||||
guardrail_response: Optional[Union[dict, str, List[dict]]]
|
||||
guardrail_status: Literal["success", "failure", "blocked"]
|
||||
guardrail_status: GuardrailStatus
|
||||
start_time: Optional[float]
|
||||
end_time: Optional[float]
|
||||
duration: Optional[float]
|
||||
|
|
@ -2082,6 +2089,20 @@ class CostBreakdown(TypedDict):
|
|||
tool_usage_cost: float # Cost of usage of built-in tools
|
||||
|
||||
|
||||
class StandardLoggingPayloadStatusFields(TypedDict, total=False):
|
||||
"""Status fields for easy filtering and analytics"""
|
||||
llm_api_status: StandardLoggingPayloadStatus
|
||||
"""Status of the LLM API call - 'success' if completed, 'failure' if errored"""
|
||||
guardrail_status: GuardrailStatus
|
||||
"""
|
||||
Status of guardrail execution:
|
||||
- 'success': Guardrail ran and allowed content through
|
||||
- 'guardrail_intervened': Guardrail blocked or modified content
|
||||
- 'guardrail_failed_to_respond': Guardrail had technical failure
|
||||
- 'not_run': No guardrail was run
|
||||
"""
|
||||
|
||||
|
||||
class StandardLoggingPayload(TypedDict):
|
||||
id: str
|
||||
trace_id: str # Trace multiple LLM calls belonging to same overall request (e.g. fallbacks/retries)
|
||||
|
|
@ -2093,6 +2114,7 @@ class StandardLoggingPayload(TypedDict):
|
|||
StandardLoggingModelCostFailureDebugInformation
|
||||
]
|
||||
status: StandardLoggingPayloadStatus
|
||||
status_fields: StandardLoggingPayloadStatusFields
|
||||
custom_llm_provider: Optional[str]
|
||||
total_tokens: int
|
||||
prompt_tokens: int
|
||||
|
|
|
|||
79
tests/guardrails_tests/conftest.py
Normal file
79
tests/guardrails_tests/conftest.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
# conftest.py
|
||||
|
||||
import importlib
|
||||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
import litellm
|
||||
import asyncio
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def event_loop():
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
loop = asyncio.new_event_loop()
|
||||
yield loop
|
||||
loop.close()
|
||||
|
||||
@pytest.fixture(scope="function", autouse=True)
|
||||
def setup_and_teardown():
|
||||
"""
|
||||
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
|
||||
"""
|
||||
curr_dir = os.getcwd() # Get the current working directory
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the project directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
import asyncio
|
||||
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
# flush all logs
|
||||
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
|
||||
|
||||
|
||||
importlib.reload(litellm)
|
||||
|
||||
try:
|
||||
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
|
||||
import litellm.proxy.proxy_server
|
||||
|
||||
importlib.reload(litellm.proxy.proxy_server)
|
||||
except Exception as e:
|
||||
print(f"Error reloading litellm.proxy.proxy_server: {e}")
|
||||
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_event_loop_policy().new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
print(litellm)
|
||||
# from litellm import Router, completion, aembedding, acompletion, embedding
|
||||
yield
|
||||
|
||||
# Teardown code (executes after the yield point)
|
||||
loop.close() # Close the loop created earlier
|
||||
asyncio.set_event_loop(None) # Remove the reference to the loop
|
||||
|
||||
|
||||
|
||||
def pytest_collection_modifyitems(config, items):
|
||||
# Separate tests in 'test_amazing_proxy_custom_logger.py' and other tests
|
||||
custom_logger_tests = [
|
||||
item for item in items if "custom_logger" in item.parent.name
|
||||
]
|
||||
other_tests = [item for item in items if "custom_logger" not in item.parent.name]
|
||||
|
||||
# Sort tests based on their names
|
||||
custom_logger_tests.sort(key=lambda x: x.name)
|
||||
other_tests.sort(key=lambda x: x.name)
|
||||
|
||||
# Reorder the items list
|
||||
items[:] = custom_logger_tests + other_tests
|
||||
|
|
@ -15,8 +15,9 @@ from litellm.types.guardrails import GuardrailEventHooks
|
|||
from typing import Optional
|
||||
|
||||
|
||||
class TestCustomLogger(CustomLogger):
|
||||
class CustomLoggerForTesting(CustomLogger):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.standard_logging_payload: Optional[StandardLoggingPayload] = None
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
@ -28,7 +29,7 @@ async def test_standard_logging_payload_includes_guardrail_information():
|
|||
"""
|
||||
Test that the standard logging payload includes the guardrail information when a guardrail is applied
|
||||
"""
|
||||
test_custom_logger = TestCustomLogger()
|
||||
test_custom_logger = CustomLoggerForTesting()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
presidio_guard = _OPTIONAL_PresidioPIIMasking(
|
||||
guardrail_name="presidio_guard",
|
||||
|
|
@ -177,4 +178,469 @@ async def test_langfuse_trace_includes_guardrail_information():
|
|||
assert output_item["entity_type"] == "PHONE_NUMBER"
|
||||
assert "score" in output_item
|
||||
assert "start" in output_item
|
||||
assert "end" in output_item
|
||||
assert "end" in output_item
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrail_status_blocked():
|
||||
"""
|
||||
Test that Bedrock guardrail sets correct status fields when blocking content.
|
||||
|
||||
This test verifies that when Bedrock guardrail blocks content:
|
||||
1. The guardrail_information contains guardrail_status="blocked"
|
||||
2. The status_fields.guardrail_status is set to "guardrail_intervened"
|
||||
3. The status_fields.llm_api_status remains "success" (mock LLM call succeeds)
|
||||
"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
litellm._turn_on_debug()
|
||||
|
||||
# Setup custom logger to capture standard logging payload
|
||||
test_custom_logger = CustomLoggerForTesting()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
|
||||
# Create Bedrock guardrail with mock AWS credentials
|
||||
bedrock_guard = BedrockGuardrail(
|
||||
guardrail_name="bedrock_guard",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
guardrailIdentifier="test-id",
|
||||
guardrailVersion="1",
|
||||
aws_access_key_id="test-key",
|
||||
aws_secret_access_key="test-secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
|
||||
# Mock Bedrock API response indicating content was blocked
|
||||
# action="GUARDRAIL_INTERVENED" means the guardrail blocked the request
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"outputs": [{"text": "Blocked"}],
|
||||
"assessments": [{
|
||||
"topicPolicy": {
|
||||
"topics": [{"name": "harmful", "action": "BLOCKED"}]
|
||||
}
|
||||
}]
|
||||
}
|
||||
bedrock_guard.async_handler.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "harmful content"}],
|
||||
"mock_response": "Hello",
|
||||
"metadata": {}
|
||||
}
|
||||
|
||||
# Mock should_run_guardrail to ensure guardrail logic executes
|
||||
with patch.object(bedrock_guard, 'should_run_guardrail', return_value=True):
|
||||
# Call guardrail pre_call hook - this will raise an exception when content is blocked
|
||||
try:
|
||||
await bedrock_guard.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=None,
|
||||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
except Exception:
|
||||
# Expected exception when guardrail blocks content
|
||||
pass
|
||||
|
||||
# Call litellm.acompletion to trigger logging callbacks
|
||||
# This populates the standard_logging_payload in our custom logger
|
||||
response = await litellm.acompletion(**request_data)
|
||||
await asyncio.sleep(1)
|
||||
|
||||
# Verify the standard logging payload was captured
|
||||
assert test_custom_logger.standard_logging_payload is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"] is not None
|
||||
|
||||
# Verify guardrail information fields
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_status"] == "guardrail_intervened"
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_provider"] == "bedrock"
|
||||
|
||||
# Verify the new typed status fields
|
||||
# guardrail_status should be "guardrail_intervened" when content is blocked
|
||||
# llm_api_status should be "success" since the mock LLM call itself succeeded
|
||||
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
|
||||
assert status_fields.get("llm_api_status") == "success"
|
||||
assert status_fields.get("guardrail_status") == "guardrail_intervened"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrail_status_success():
|
||||
"""
|
||||
Test that Bedrock guardrail sets correct status fields when allowing content.
|
||||
|
||||
This test verifies that when Bedrock guardrail allows content through:
|
||||
1. The guardrail_information contains guardrail_status="success"
|
||||
2. The status_fields.guardrail_status is set to "success"
|
||||
3. The status_fields.llm_api_status is "success"
|
||||
"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
# Reset callbacks completely to avoid event loop conflicts
|
||||
litellm.callbacks = []
|
||||
await asyncio.sleep(0.1) # Let previous callbacks finish
|
||||
|
||||
# Setup custom logger to capture standard logging payload
|
||||
test_custom_logger = CustomLoggerForTesting()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
|
||||
# Create Bedrock guardrail
|
||||
bedrock_guard = BedrockGuardrail(
|
||||
guardrail_name="bedrock_guard",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
guardrailIdentifier="test-id",
|
||||
guardrailVersion="1",
|
||||
aws_access_key_id="test-key",
|
||||
aws_secret_access_key="test-secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
|
||||
# Mock success response
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"action": "NONE",
|
||||
"outputs": [{"text": "Safe content"}],
|
||||
"assessments": []
|
||||
}
|
||||
bedrock_guard.async_handler.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "safe content"}],
|
||||
"mock_response": "Hello",
|
||||
"metadata": {}
|
||||
}
|
||||
|
||||
# Mock should_run_guardrail to return True
|
||||
with patch.object(bedrock_guard, 'should_run_guardrail', return_value=True):
|
||||
await bedrock_guard.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=None,
|
||||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
|
||||
# Call litellm.acompletion to trigger logging
|
||||
response = await litellm.acompletion(**request_data)
|
||||
await asyncio.sleep(1)
|
||||
|
||||
# Check standard logging payload status fields
|
||||
assert test_custom_logger.standard_logging_payload is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"] is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_status"] == "success"
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_provider"] == "bedrock"
|
||||
|
||||
# Check status fields
|
||||
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
|
||||
assert status_fields.get("llm_api_status") == "success"
|
||||
assert status_fields.get("guardrail_status") == "success"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_guardrail_status_failure():
|
||||
"""
|
||||
Test that Bedrock guardrail sets correct status fields when the API endpoint fails.
|
||||
|
||||
This test verifies that when Bedrock guardrail API is down/fails:
|
||||
1. The guardrail_information contains guardrail_status="failure"
|
||||
2. The status_fields.guardrail_status is set to "guardrail_failed_to_respond"
|
||||
3. The exception is still raised (maintaining existing behavior)
|
||||
"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import httpx
|
||||
|
||||
# Reset callbacks completely to avoid event loop conflicts
|
||||
litellm.callbacks = []
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Setup custom logger to capture standard logging payload
|
||||
test_custom_logger = CustomLoggerForTesting()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
|
||||
# Create Bedrock guardrail
|
||||
bedrock_guard = BedrockGuardrail(
|
||||
guardrail_name="bedrock_guard",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
guardrailIdentifier="test-id",
|
||||
guardrailVersion="1",
|
||||
aws_access_key_id="test-key",
|
||||
aws_secret_access_key="test-secret",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
|
||||
# Mock network failure (endpoint down)
|
||||
bedrock_guard.async_handler.post = AsyncMock(
|
||||
side_effect=httpx.ConnectError("Connection failed")
|
||||
)
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "test content"}],
|
||||
"mock_response": "Hello",
|
||||
"metadata": {}
|
||||
}
|
||||
|
||||
# Mock should_run_guardrail to return True
|
||||
with patch.object(bedrock_guard, 'should_run_guardrail', return_value=True):
|
||||
# Call guardrail (will raise exception on network failure)
|
||||
try:
|
||||
await bedrock_guard.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=None,
|
||||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
except Exception:
|
||||
# Expected exception when endpoint is down
|
||||
pass
|
||||
|
||||
# Call litellm.acompletion to trigger logging
|
||||
response = await litellm.acompletion(**request_data)
|
||||
await asyncio.sleep(1)
|
||||
|
||||
# Check standard logging payload status fields
|
||||
assert test_custom_logger.standard_logging_payload is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"] is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_provider"] == "bedrock"
|
||||
|
||||
# Check status fields
|
||||
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
|
||||
assert status_fields.get("llm_api_status") == "success"
|
||||
assert status_fields.get("guardrail_status") == "guardrail_failed_to_respond"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_noma_guardrail_status_blocked():
|
||||
"""
|
||||
Test that Noma guardrail sets correct status fields when blocking content.
|
||||
|
||||
This test verifies that when Noma guardrail blocks content (verdict=False):
|
||||
1. The guardrail_information contains guardrail_status="blocked"
|
||||
2. The status_fields.guardrail_status is set to "guardrail_intervened"
|
||||
3. The status_fields.llm_api_status remains "success"
|
||||
"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.noma.noma import NomaGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
# Reset callbacks completely to avoid event loop conflicts
|
||||
litellm.callbacks = []
|
||||
await asyncio.sleep(0.1) # Let previous callbacks finish
|
||||
|
||||
# Setup custom logger to capture standard logging payload
|
||||
test_custom_logger = CustomLoggerForTesting()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
|
||||
# Create Noma guardrail
|
||||
noma_guard = NomaGuardrail(
|
||||
guardrail_name="noma_guard",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
api_key="test-key",
|
||||
monitor_mode=False,
|
||||
)
|
||||
|
||||
# Mock blocked response
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"verdict": False,
|
||||
"originalResponse": {
|
||||
"prompt": {
|
||||
"topicDetector": {"harmful": {"result": True}}
|
||||
}
|
||||
}
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
noma_guard.async_handler.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "harmful content"}],
|
||||
"mock_response": "Hello",
|
||||
"metadata": {}
|
||||
}
|
||||
|
||||
# Mock should_run_guardrail to return True
|
||||
with patch.object(noma_guard, 'should_run_guardrail', return_value=True):
|
||||
# Call guardrail (will raise exception on block)
|
||||
try:
|
||||
await noma_guard.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=None,
|
||||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Call litellm.acompletion to trigger logging
|
||||
response = await litellm.acompletion(**request_data)
|
||||
await asyncio.sleep(1)
|
||||
|
||||
# Check standard logging payload status fields
|
||||
assert test_custom_logger.standard_logging_payload is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"] is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_status"] == "guardrail_intervened"
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_provider"] == "noma"
|
||||
|
||||
# Check status fields
|
||||
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
|
||||
assert status_fields.get("llm_api_status") == "success"
|
||||
assert status_fields.get("guardrail_status") == "guardrail_intervened"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_noma_guardrail_status_success():
|
||||
"""
|
||||
Test that Noma guardrail sets correct status fields when allowing content.
|
||||
|
||||
This test verifies that when Noma guardrail allows content (verdict=True):
|
||||
1. The guardrail_information contains guardrail_status="success"
|
||||
2. The status_fields.guardrail_status is set to "success"
|
||||
3. The status_fields.llm_api_status is "success"
|
||||
"""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.noma.noma import NomaGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
# Reset callbacks completely to avoid event loop conflicts
|
||||
litellm.callbacks = []
|
||||
await asyncio.sleep(0.1) # Let previous callbacks finish
|
||||
|
||||
# Setup custom logger to capture standard logging payload
|
||||
test_custom_logger = CustomLoggerForTesting()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
|
||||
# Create Noma guardrail
|
||||
noma_guard = NomaGuardrail(
|
||||
guardrail_name="noma_guard",
|
||||
event_hook=GuardrailEventHooks.pre_call,
|
||||
api_key="test-key",
|
||||
monitor_mode=False,
|
||||
)
|
||||
|
||||
# Mock success response
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {
|
||||
"verdict": True,
|
||||
"originalResponse": {"prompt": {}}
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
noma_guard.async_handler.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "safe content"}],
|
||||
"mock_response": "Hello",
|
||||
"metadata": {}
|
||||
}
|
||||
|
||||
# Mock should_run_guardrail to return True
|
||||
with patch.object(noma_guard, 'should_run_guardrail', return_value=True):
|
||||
await noma_guard.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=None,
|
||||
data=request_data,
|
||||
call_type="completion"
|
||||
)
|
||||
|
||||
# Call litellm.acompletion to trigger logging
|
||||
response = await litellm.acompletion(**request_data)
|
||||
await asyncio.sleep(1)
|
||||
|
||||
# Check standard logging payload status fields
|
||||
assert test_custom_logger.standard_logging_payload is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"] is not None
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_status"] == "success"
|
||||
assert test_custom_logger.standard_logging_payload["guardrail_information"]["guardrail_provider"] == "noma"
|
||||
|
||||
# Check status fields
|
||||
status_fields = test_custom_logger.standard_logging_payload.get("status_fields", {})
|
||||
assert status_fields.get("llm_api_status") == "success"
|
||||
assert status_fields.get("guardrail_status") == "success"
|
||||
|
||||
|
||||
def test_guardrail_status_fields_computation():
|
||||
"""
|
||||
Test that status fields are computed correctly from guardrail information.
|
||||
|
||||
This unit test verifies the _get_status_fields function correctly maps:
|
||||
- guardrail_status="blocked" -> status_fields.guardrail_status="guardrail_intervened" (legacy)
|
||||
- guardrail_status="guardrail_intervened" -> status_fields.guardrail_status="guardrail_intervened"
|
||||
- guardrail_status="success" -> status_fields.guardrail_status="success"
|
||||
- guardrail_status="failure" -> status_fields.guardrail_status="guardrail_failed_to_respond" (legacy)
|
||||
- guardrail_status="guardrail_failed_to_respond" -> status_fields.guardrail_status="guardrail_failed_to_respond"
|
||||
- no guardrail -> status_fields.guardrail_status="not_run"
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import _get_status_fields
|
||||
|
||||
# Test guardrail_intervened status (content was blocked by guardrail)
|
||||
intervened_info = {"guardrail_status": "guardrail_intervened"}
|
||||
status_fields_intervened = _get_status_fields(
|
||||
status="success",
|
||||
guardrail_information=intervened_info,
|
||||
error_str=None
|
||||
)
|
||||
assert status_fields_intervened["llm_api_status"] == "success"
|
||||
assert status_fields_intervened["guardrail_status"] == "guardrail_intervened"
|
||||
|
||||
# Test legacy blocked status (for backward compatibility)
|
||||
blocked_info = {"guardrail_status": "blocked"}
|
||||
status_fields_blocked = _get_status_fields(
|
||||
status="success",
|
||||
guardrail_information=blocked_info,
|
||||
error_str=None
|
||||
)
|
||||
assert status_fields_blocked["llm_api_status"] == "success"
|
||||
assert status_fields_blocked["guardrail_status"] == "guardrail_intervened"
|
||||
|
||||
# Test success status
|
||||
success_info = {"guardrail_status": "success"}
|
||||
status_fields_success = _get_status_fields(
|
||||
status="success",
|
||||
guardrail_information=success_info,
|
||||
error_str=None
|
||||
)
|
||||
assert status_fields_success["llm_api_status"] == "success"
|
||||
assert status_fields_success["guardrail_status"] == "success"
|
||||
|
||||
# Test guardrail_failed_to_respond status
|
||||
failed_info = {"guardrail_status": "guardrail_failed_to_respond"}
|
||||
status_fields_failed = _get_status_fields(
|
||||
status="failure",
|
||||
guardrail_information=failed_info,
|
||||
error_str=None
|
||||
)
|
||||
assert status_fields_failed["llm_api_status"] == "failure"
|
||||
assert status_fields_failed["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
|
||||
# Test legacy failure status (for backward compatibility)
|
||||
failure_info = {"guardrail_status": "failure"}
|
||||
status_fields_failure = _get_status_fields(
|
||||
status="failure",
|
||||
guardrail_information=failure_info,
|
||||
error_str=None
|
||||
)
|
||||
assert status_fields_failure["llm_api_status"] == "failure"
|
||||
assert status_fields_failure["guardrail_status"] == "guardrail_failed_to_respond"
|
||||
|
||||
# Test no guardrail run
|
||||
no_guardrail = None
|
||||
status_fields_no_guardrail = _get_status_fields(
|
||||
status="success",
|
||||
guardrail_information=no_guardrail,
|
||||
error_str=None
|
||||
)
|
||||
assert status_fields_no_guardrail["llm_api_status"] == "success"
|
||||
assert status_fields_no_guardrail["guardrail_status"] == "not_run"
|
||||
Loading…
Add table
Reference in a new issue