mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(guardrails): Add content masking and streaming support to PANW Prisma AIRS guardrail (#15666)
* feat(guardrails): Add content masking and streaming support to PANW Prisma AIRS - Add mask_request_content and mask_response_content parameters - Implement content masking for prompts and responses - Add streaming support with real-time masking - Add comprehensive test coverage (28 tests) - Update documentation with masking examples and security notes * fix(guardrails): Fix PANW Prisma AIRS env var fallback and text completion support
This commit is contained in:
parent
645f84c02e
commit
c471bf1f16
5 changed files with 1196 additions and 147 deletions
|
|
@ -11,10 +11,13 @@ LiteLLM supports PANW Prisma AIRS (AI Runtime Security) guardrails via the [Pris
|
|||
- ✅ **Real-time prompt injection detection**
|
||||
- ✅ **Malicious content filtering**
|
||||
- ✅ **Data loss prevention (DLP)**
|
||||
- ✅ **Sensitive content masking** - Automatically mask PII, credit cards, SSNs instead of blocking
|
||||
- ✅ **Comprehensive threat detection** for AI models and datasets
|
||||
- ✅ **Model-agnostic protection** across public and private models
|
||||
- ✅ **Synchronous scanning** with immediate response
|
||||
- ✅ **Configurable security profiles**
|
||||
- ✅ **Streaming support** - Real-time masking for streaming responses
|
||||
- ✅ **Fail-closed security** - Blocks requests if PANW API is unavailable (maximum security)
|
||||
|
||||
## Quick Start
|
||||
|
||||
|
|
@ -42,9 +45,9 @@ guardrails:
|
|||
litellm_params:
|
||||
guardrail: panw_prisma_airs
|
||||
mode: "pre_call" # Run before LLM call
|
||||
api_key: os.environ/AIRS_API_KEY # Your PANW API key
|
||||
profile_name: os.environ/AIRS_API_PROFILE_NAME # Security profile from Strata Cloud Manager
|
||||
api_base: "https://service.api.aisecurity.paloaltonetworks.com/v1/scan/sync/request" # Optional
|
||||
api_key: os.environ/PANW_PRISMA_AIRS_API_KEY # Your Prisma AIRS API key
|
||||
profile_name: os.environ/PANW_PRISMA_AIRS_PROFILE_NAME # Security profile from Strata Cloud Manager
|
||||
api_base: "https://service.api.aisecurity.paloaltonetworks.com"
|
||||
```
|
||||
|
||||
#### Supported values for `mode`
|
||||
|
|
@ -56,8 +59,8 @@ guardrails:
|
|||
### 3. Start LiteLLM Gateway
|
||||
|
||||
```bash title="Set environment variables"
|
||||
export AIRS_API_KEY="your-panw-api-key"
|
||||
export AIRS_API_PROFILE_NAME="your-security-profile"
|
||||
export PANW_PRISMA_AIRS_API_KEY="your-panw-api-key"
|
||||
export PANW_PRISMA_AIRS_PROFILE_NAME="your-security-profile"
|
||||
export OPENAI_API_KEY="sk-proj-..."
|
||||
```
|
||||
|
||||
|
|
@ -197,16 +200,16 @@ Expected successful response:
|
|||
|-----------|----------|-------------|---------|
|
||||
| `api_key` | Yes | Your PANW Prisma AIRS API key from Strata Cloud Manager | - |
|
||||
| `profile_name` | Yes | Security profile name configured in Strata Cloud Manager | - |
|
||||
| `api_base` | No | Custom API endpoint | `https://service.api.aisecurity.paloaltonetworks.com/v1/scan/sync/request` |
|
||||
| `api_base` | No | Custom API base URL (without /v1/scan/sync/request path) | `https://service.api.aisecurity.paloaltonetworks.com` |
|
||||
| `mode` | No | When to run the guardrail | `pre_call` |
|
||||
|
||||
## Environment Variables
|
||||
|
||||
```bash
|
||||
export AIRS_API_KEY="your-panw-api-key"
|
||||
export AIRS_API_PROFILE_NAME="your-security-profile"
|
||||
# Optional custom endpoint
|
||||
export PANW_API_ENDPOINT="https://custom-endpoint.com/v1/scan/sync/request"
|
||||
export PANW_PRISMA_AIRS_API_KEY="your-panw-api-key"
|
||||
export PANW_PRISMA_AIRS_PROFILE_NAME="your-security-profile"
|
||||
# Optional custom base URL (without /v1/scan/sync/request path)
|
||||
export PANW_PRISMA_AIRS_API_BASE="https://custom-endpoint.com"
|
||||
```
|
||||
|
||||
## Advanced Configuration
|
||||
|
|
@ -221,17 +224,125 @@ guardrails:
|
|||
litellm_params:
|
||||
guardrail: panw_prisma_airs
|
||||
mode: "pre_call"
|
||||
api_key: os.environ/AIRS_API_KEY
|
||||
api_key: os.environ/PANW_PRISMA_AIRS_API_KEY
|
||||
profile_name: "strict-policy" # High security profile
|
||||
|
||||
- guardrail_name: "panw-permissive-security"
|
||||
litellm_params:
|
||||
guardrail: panw_prisma_airs
|
||||
mode: "post_call"
|
||||
api_key: os.environ/AIRS_API_KEY
|
||||
api_key: os.environ/PANW_PRISMA_AIRS_API_KEY
|
||||
profile_name: "permissive-policy" # Lower security profile
|
||||
```
|
||||
|
||||
### Content Masking
|
||||
|
||||
PANW Prisma AIRS can automatically mask sensitive content (PII, credit cards, SSNs, etc.) instead of blocking requests. This allows your application to continue functioning while protecting sensitive data.
|
||||
|
||||
#### How It Works
|
||||
|
||||
1. **Detection**: PANW scans content and identifies sensitive data
|
||||
2. **Masking**: Sensitive data is replaced with placeholders (e.g., `XXXXXXXXXX` or `{PHONE}`)
|
||||
3. **Pass-through**: Masked content is sent to the LLM or returned to the user
|
||||
|
||||
#### Configuration Options
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: "panw-with-masking"
|
||||
litellm_params:
|
||||
guardrail: panw_prisma_airs
|
||||
mode: "post_call" # Scan both input and output
|
||||
api_key: os.environ/PANW_PRISMA_AIRS_API_KEY
|
||||
profile_name: "default"
|
||||
mask_request_content: true # Mask sensitive data in prompts
|
||||
mask_response_content: true # Mask sensitive data in responses
|
||||
```
|
||||
|
||||
**Masking Parameters:**
|
||||
|
||||
- `mask_request_content: true` - When PANW detects sensitive data in prompts, mask it instead of blocking
|
||||
- `mask_response_content: true` - When PANW detects sensitive data in responses, mask it instead of blocking
|
||||
- `mask_on_block: true` - Backwards compatible flag that enables both request and response masking
|
||||
|
||||
:::warning Important: Masking is Controlled by PANW Security Profile
|
||||
The **actual masking behavior** (what content gets masked and how) is controlled by your **PANW Prisma AIRS security profile** configured in Strata Cloud Manager. The LiteLLM config settings (`mask_request_content`, `mask_response_content`) only control whether to:
|
||||
- **Apply the masked content** returned by PANW and allow the request to continue, OR
|
||||
- **Block the request** entirely when sensitive data is detected
|
||||
|
||||
LiteLLM does not alter or configure your PANW security profile. To change what content gets masked, update your profile settings in Strata Cloud Manager.
|
||||
:::
|
||||
|
||||
:::info Security Posture
|
||||
The guardrail is **fail-closed** by default - if the PANW API is unavailable, requests are blocked to ensure no unscanned content reaches your LLM. This provides maximum security.
|
||||
:::
|
||||
|
||||
#### Example: Masking Credit Card Numbers
|
||||
|
||||
<Tabs>
|
||||
<TabItem label="Without Masking" value="no-mask">
|
||||
|
||||
**Request:**
|
||||
```json
|
||||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "My credit card is 4929-3813-3266-4295"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
**Response:** ❌ **Blocked with 400 error**
|
||||
|
||||
</TabItem>
|
||||
<TabItem label="With Masking" value="with-mask">
|
||||
|
||||
**Request:**
|
||||
```json
|
||||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "My credit card is 4929-3813-3266-4295"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
**Masked prompt sent to LLM:**
|
||||
```json
|
||||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "My credit card is XXXXXXXXXXXXXXXXXX"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
**Response:** ✅ **Allowed with masked content**
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
#### Masking Capabilities
|
||||
|
||||
The guardrail masks sensitive content in:
|
||||
|
||||
- ✅ **Chat messages** - User prompts and assistant responses
|
||||
- ✅ **Streaming responses** - Real-time masking of streamed content
|
||||
- ✅ **Multi-choice responses** - All choices in the response
|
||||
- ✅ **Tool/function calls** - Arguments passed to tools and functions
|
||||
- ✅ **Content lists** - Mixed content types (text, images, etc.)
|
||||
|
||||
#### Complete Example
|
||||
|
||||
```yaml
|
||||
guardrails:
|
||||
- guardrail_name: "panw-production-security"
|
||||
litellm_params:
|
||||
guardrail: panw_prisma_airs
|
||||
mode: "post_call" # Scan input and output
|
||||
api_key: os.environ/PANW_PRISMA_AIRS_API_KEY
|
||||
profile_name: "production-profile"
|
||||
mask_request_content: true # Mask sensitive prompts
|
||||
mask_response_content: true # Mask sensitive responses
|
||||
```
|
||||
|
||||
## Use Cases
|
||||
|
||||
From [official Prisma AIRS documentation](https://docs.paloaltonetworks.com/ai-runtime-security/activation-and-onboarding/ai-runtime-security-api-intercept-overview):
|
||||
|
|
@ -245,7 +356,7 @@ From [official Prisma AIRS documentation](https://docs.paloaltonetworks.com/ai-r
|
|||
## Next Steps
|
||||
|
||||
- Configure your security policies in [Strata Cloud Manager](https://apps.paloaltonetworks.com/)
|
||||
- Review the [Prisma AIRS API documentation](https://pan.dev/prisma-airs/api/airuntimesecurity/scan-sync-request/) for advanced features
|
||||
- Review the [Prisma AIRS API documentation](https://pan.dev/airs/) for advanced features
|
||||
- Set up monitoring and alerting for threat detections in your PANW dashboard
|
||||
- Consider implementing both pre_call and post_call guardrails for comprehensive protection
|
||||
- Monitor detection events and tune your security profiles based on your application needs
|
||||
|
|
@ -13,8 +13,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
|
||||
guardrail_name = guardrail.get("guardrail_name")
|
||||
profile_name = cast(Optional[str], getattr(litellm_params, "profile_name", None))
|
||||
if not litellm_params.api_key:
|
||||
raise ValueError("PANW Prisma AIRS: api_key is required")
|
||||
|
||||
# Note: api_key can be None here - handler will fallback to PANW_PRISMA_AIRS_API_KEY env var
|
||||
if not profile_name:
|
||||
raise ValueError("PANW Prisma AIRS: profile_name is required")
|
||||
if not guardrail_name:
|
||||
|
|
|
|||
|
|
@ -1,12 +1,13 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
PANW Prisma AIRS Built-in Guardrail for LiteLLM
|
||||
Palo Alto Networks Prisma AI Runtime Security (AIRS) Guardrail Integration for LiteLLM
|
||||
|
||||
Provides real-time threat detection, DLP, URL filtering, content masking, and policy enforcement for AI applications.
|
||||
"""
|
||||
|
||||
import os
|
||||
from litellm._uuid import uuid
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Type, cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Type
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -30,32 +31,48 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
"""
|
||||
LiteLLM Built-in Guardrail for Palo Alto Networks Prisma AI Runtime Security (AIRS).
|
||||
|
||||
This guardrail scans prompts and responses using the PANW Prisma AIRS API to detect
|
||||
malicious content, injection attempts, and policy violations.
|
||||
Scans prompts and responses using PANW Prisma AIRS API to detect malicious content,
|
||||
injection attempts, and policy violations. Supports content masking and fail-closed error handling.
|
||||
|
||||
Configuration:
|
||||
guardrail_name: Name of the guardrail instance
|
||||
api_key: PANW Prisma AIRS API key
|
||||
api_base: PANW Prisma AIRS API endpoint
|
||||
profile_name: PANW Prisma AIRS security profile name
|
||||
default_on: Whether to enable by default
|
||||
api_base: PANW Prisma AIRS API endpoint (default: https://service.api.aisecurity.paloaltonetworks.com)
|
||||
profile_name: PANW security profile name
|
||||
mask_request_content: Apply masking to prompts (default: False)
|
||||
mask_response_content: Apply masking to responses (default: False)
|
||||
mask_on_block: Backwards compatible flag that enables both request and response masking
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
guardrail_name: str,
|
||||
api_key: str,
|
||||
api_base: str,
|
||||
profile_name: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
default_on: bool = True,
|
||||
mask_on_block: bool = False,
|
||||
mask_request_content: bool = False,
|
||||
mask_response_content: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
"""Initialize PANW Prisma AIRS guardrail handler."""
|
||||
|
||||
# Initialize parent CustomGuardrail
|
||||
super().__init__(guardrail_name=guardrail_name, default_on=default_on, **kwargs)
|
||||
# Masking configuration - mask_on_block enables both for backwards compatibility
|
||||
self.mask_on_block = mask_on_block
|
||||
_mask_request_content = mask_request_content or mask_on_block
|
||||
_mask_response_content = mask_response_content or mask_on_block
|
||||
|
||||
# Store configuration
|
||||
# Initialize parent CustomGuardrail with masking flags
|
||||
super().__init__(
|
||||
guardrail_name=guardrail_name,
|
||||
default_on=default_on,
|
||||
mask_request_content=_mask_request_content,
|
||||
mask_response_content=_mask_response_content,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
# Store configuration with env var fallbacks
|
||||
self.api_key = api_key or os.getenv("PANW_PRISMA_AIRS_API_KEY")
|
||||
self.api_base = (
|
||||
api_base
|
||||
|
|
@ -63,9 +80,17 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
or "https://service.api.aisecurity.paloaltonetworks.com"
|
||||
)
|
||||
self.profile_name = profile_name
|
||||
|
||||
# Validate required configuration
|
||||
if not self.api_key:
|
||||
raise ValueError(
|
||||
"PANW Prisma AIRS: api_key is required. "
|
||||
"Set it via config or PANW_PRISMA_AIRS_API_KEY environment variable."
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Initialized PANW Prisma AIRS Guardrail: {guardrail_name}"
|
||||
verbose_proxy_logger.info(
|
||||
f"Initialized PANW Prisma AIRS Guardrail: {guardrail_name} "
|
||||
f"(mask_request={self.mask_request_content}, mask_response={self.mask_response_content})"
|
||||
)
|
||||
|
||||
def _extract_text_from_messages(self, messages: List[Dict[str, Any]]) -> str:
|
||||
|
|
@ -104,20 +129,38 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
return " ".join(text_parts) if text_parts else ""
|
||||
|
||||
def _extract_response_text(self, response: ModelResponse) -> str:
|
||||
"""Extract text from LLM response."""
|
||||
"""
|
||||
Extract all text content from LLM response.
|
||||
Handles multiple choices, tool calls, and function calls.
|
||||
Returns concatenated text for scanning.
|
||||
"""
|
||||
try:
|
||||
from litellm.types.utils import Choices
|
||||
|
||||
if (
|
||||
hasattr(response, "choices")
|
||||
and response.choices
|
||||
and len(response.choices) > 0
|
||||
and hasattr(response.choices[0], "message")
|
||||
):
|
||||
return cast(Choices, response.choices[0]).message.content or ""
|
||||
except (AttributeError, IndexError):
|
||||
|
||||
text_parts = []
|
||||
|
||||
if hasattr(response, "choices") and response.choices:
|
||||
for choice in response.choices:
|
||||
if isinstance(choice, Choices):
|
||||
# Extract message content
|
||||
if choice.message.content:
|
||||
text_parts.append(str(choice.message.content))
|
||||
|
||||
# Extract tool call arguments
|
||||
if hasattr(choice.message, "tool_calls") and choice.message.tool_calls:
|
||||
for tool_call in choice.message.tool_calls:
|
||||
if hasattr(tool_call, "function") and hasattr(tool_call.function, "arguments"):
|
||||
text_parts.append(str(tool_call.function.arguments))
|
||||
|
||||
# Extract function call arguments (legacy)
|
||||
if hasattr(choice.message, "function_call") and choice.message.function_call:
|
||||
if hasattr(choice.message.function_call, "arguments"):
|
||||
text_parts.append(str(choice.message.function_call.arguments))
|
||||
|
||||
return " ".join(text_parts) if text_parts else ""
|
||||
except (AttributeError, IndexError) as e:
|
||||
verbose_proxy_logger.error(
|
||||
"PANW Prisma AIRS: Error extracting response text"
|
||||
f"PANW Prisma AIRS: Error extracting response text: {str(e)}"
|
||||
)
|
||||
return ""
|
||||
|
||||
|
|
@ -191,6 +234,88 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
verbose_proxy_logger.error(f"PANW Prisma AIRS: API call failed: {str(e)}")
|
||||
return {"action": "block", "category": "api_error"}
|
||||
|
||||
def _get_masked_text(self, scan_result: Dict[str, Any], is_response: bool = False) -> Optional[str]:
|
||||
"""Extract masked text from PANW scan result."""
|
||||
masked_key = "response_masked_data" if is_response else "prompt_masked_data"
|
||||
masked_data = scan_result.get(masked_key)
|
||||
if masked_data and isinstance(masked_data, dict):
|
||||
return masked_data.get("data")
|
||||
return None
|
||||
|
||||
def _apply_masking_to_messages(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
masked_text: str
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Apply masked text to the last user message."""
|
||||
if not messages:
|
||||
return messages
|
||||
|
||||
for i, message in enumerate(reversed(messages)):
|
||||
if message.get("role") == "user":
|
||||
new_message = message.copy()
|
||||
content = message.get("content")
|
||||
|
||||
if isinstance(content, str):
|
||||
new_message["content"] = masked_text
|
||||
elif isinstance(content, list):
|
||||
new_content = []
|
||||
for part in content:
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
new_content.append({"type": "text", "text": masked_text})
|
||||
else:
|
||||
new_content.append(part)
|
||||
new_message["content"] = new_content
|
||||
|
||||
idx = len(messages) - i - 1
|
||||
return messages[:idx] + [new_message] + messages[idx+1:]
|
||||
|
||||
return messages
|
||||
|
||||
def _apply_masking_to_response(
|
||||
self,
|
||||
response: ModelResponse,
|
||||
masked_text: str
|
||||
) -> None:
|
||||
"""
|
||||
Apply masked text to all content in response in-place.
|
||||
Handles message content, tool calls, and function calls across all choices.
|
||||
Preserves list-based content structure (e.g., multimodal messages).
|
||||
"""
|
||||
from litellm.types.utils import Choices
|
||||
|
||||
if not hasattr(response, "choices") or not response.choices:
|
||||
return
|
||||
|
||||
for choice in response.choices:
|
||||
if isinstance(choice, Choices):
|
||||
# Mask message content - handle both string and list formats
|
||||
content = choice.message.content
|
||||
if content:
|
||||
if isinstance(content, str):
|
||||
choice.message.content = masked_text
|
||||
elif isinstance(content, list):
|
||||
# Preserve list structure, only replace text parts
|
||||
new_content = []
|
||||
for part in content: # type: ignore
|
||||
if isinstance(part, dict) and part.get("type") == "text":
|
||||
new_content.append({"type": "text", "text": masked_text})
|
||||
else:
|
||||
# Preserve non-text parts (images, etc.)
|
||||
new_content.append(part)
|
||||
choice.message.content = new_content # type: ignore
|
||||
|
||||
# Mask tool call arguments
|
||||
if hasattr(choice.message, "tool_calls") and choice.message.tool_calls:
|
||||
for tool_call in choice.message.tool_calls:
|
||||
if hasattr(tool_call, "function") and hasattr(tool_call.function, "arguments"):
|
||||
tool_call.function.arguments = masked_text
|
||||
|
||||
# Mask function call arguments (legacy)
|
||||
if hasattr(choice.message, "function_call") and choice.message.function_call:
|
||||
if hasattr(choice.message.function_call, "arguments"):
|
||||
choice.message.function_call.arguments = masked_text
|
||||
|
||||
def _build_error_detail(
|
||||
self, scan_result: Dict[str, Any], is_response: bool = False
|
||||
) -> Dict[str, Any]:
|
||||
|
|
@ -253,96 +378,350 @@ class PanwPrismaAirsHandler(CustomGuardrail):
|
|||
|
||||
Raises HTTPException if content should be blocked.
|
||||
"""
|
||||
verbose_proxy_logger.debug("PANW Prisma AIRS: Running pre-call prompt scan")
|
||||
|
||||
# Extract prompt text from messages
|
||||
messages = data.get("messages", [])
|
||||
prompt_text = self._extract_text_from_messages(messages)
|
||||
|
||||
if not prompt_text:
|
||||
verbose_proxy_logger.warning(
|
||||
"PANW Prisma AIRS: No user prompt found in request"
|
||||
)
|
||||
return None
|
||||
|
||||
# Prepare metadata
|
||||
metadata = {
|
||||
"user": data.get("user", "litellm_user"),
|
||||
"model": data.get("model", "unknown"),
|
||||
}
|
||||
|
||||
# Scan prompt with PANW Prisma AIRS
|
||||
scan_result = await self._call_panw_api(
|
||||
content=prompt_text, is_response=False, metadata=metadata
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
action = scan_result.get("action", "block")
|
||||
category = scan_result.get("category", "unknown")
|
||||
verbose_proxy_logger.info("PANW Prisma AIRS: Running pre-call prompt scan")
|
||||
|
||||
if action == "allow":
|
||||
verbose_proxy_logger.debug(
|
||||
f"PANW Prisma AIRS: Response allowed (Category: {category})"
|
||||
# Check if guardrail should run for this request
|
||||
event_type: GuardrailEventHooks = GuardrailEventHooks.pre_call
|
||||
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
|
||||
return data
|
||||
|
||||
try:
|
||||
# Extract prompt text from messages (chat completion) or prompt (text completion)
|
||||
messages = data.get("messages", [])
|
||||
prompt_text = self._extract_text_from_messages(messages)
|
||||
|
||||
# Fallback to prompt field for text completion requests
|
||||
if not prompt_text:
|
||||
prompt_value = data.get("prompt")
|
||||
if isinstance(prompt_value, str):
|
||||
prompt_text = prompt_value
|
||||
elif isinstance(prompt_value, list):
|
||||
# Handle list of prompts (batch text completion)
|
||||
prompt_text = " ".join(str(p) for p in prompt_value if p)
|
||||
else:
|
||||
prompt_text = ""
|
||||
|
||||
if not prompt_text:
|
||||
verbose_proxy_logger.warning(
|
||||
"PANW Prisma AIRS: No user prompt found in request (checked 'messages' and 'prompt' fields)"
|
||||
)
|
||||
return None
|
||||
|
||||
# Prepare metadata
|
||||
metadata = {
|
||||
"user": data.get("user") or "litellm_user",
|
||||
"model": data.get("model") or "unknown",
|
||||
}
|
||||
|
||||
# Scan prompt with PANW Prisma AIRS
|
||||
scan_result = await self._call_panw_api(
|
||||
content=prompt_text, is_response=False, metadata=metadata
|
||||
)
|
||||
|
||||
else:
|
||||
error_detail = self._build_error_detail(scan_result, is_response=True)
|
||||
action = scan_result.get("action", "block")
|
||||
category = scan_result.get("category", "unknown")
|
||||
masked_text = self._get_masked_text(scan_result, is_response=False)
|
||||
|
||||
# If action is "allow", apply masking if available and allow through
|
||||
if action == "allow":
|
||||
if masked_text:
|
||||
if messages:
|
||||
data["messages"] = self._apply_masking_to_messages(messages, masked_text)
|
||||
elif "prompt" in data:
|
||||
data["prompt"] = masked_text
|
||||
verbose_proxy_logger.info(
|
||||
f"PANW Prisma AIRS: Prompt allowed with masking (Category: {category})"
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.info(
|
||||
f"PANW Prisma AIRS: Prompt allowed (Category: {category})"
|
||||
)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return None
|
||||
|
||||
# Action is "block" - check if we should mask instead of blocking
|
||||
if masked_text and self.mask_request_content:
|
||||
if messages:
|
||||
data["messages"] = self._apply_masking_to_messages(messages, masked_text)
|
||||
elif "prompt" in data:
|
||||
data["prompt"] = masked_text
|
||||
verbose_proxy_logger.warning(
|
||||
"PANW Prisma AIRS: Prompt blocked but masked instead (mask_request_content=True)"
|
||||
)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return None
|
||||
|
||||
# Block the request
|
||||
error_detail = self._build_error_detail(scan_result, is_response=False)
|
||||
verbose_proxy_logger.warning(
|
||||
f"PANW Prisma AIRS: {error_detail['error']['message']}"
|
||||
)
|
||||
raise HTTPException(status_code=400, detail=error_detail)
|
||||
|
||||
return None
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"PANW Prisma AIRS scan failed: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": {
|
||||
"message": "Security scan failed - request blocked for safety",
|
||||
"type": "guardrail_scan_error",
|
||||
"code": "panw_prisma_airs_scan_failed",
|
||||
"guardrail": self.guardrail_name,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: Dict[str, Any],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: ModelResponse,
|
||||
) -> ModelResponse:
|
||||
response: Any,
|
||||
) -> Any:
|
||||
"""
|
||||
Post-call hook to scan LLM responses before returning to user.
|
||||
|
||||
Raises HTTPException if response should be blocked.
|
||||
"""
|
||||
verbose_proxy_logger.debug("PANW Prisma AIRS: Running post-call response scan")
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
# Extract response text
|
||||
response_text = self._extract_response_text(response)
|
||||
|
||||
if not response_text:
|
||||
verbose_proxy_logger.warning(
|
||||
"PANW Prisma AIRS: No response content found to scan"
|
||||
)
|
||||
# Only process ModelResponse objects
|
||||
if not isinstance(response, ModelResponse):
|
||||
return response
|
||||
|
||||
# Prepare metadata
|
||||
metadata = {
|
||||
"user": data.get("user", "litellm_user"),
|
||||
"model": data.get("model", "unknown"),
|
||||
}
|
||||
verbose_proxy_logger.info("PANW Prisma AIRS: Running post-call response scan")
|
||||
|
||||
# Scan response with PANW Prisma AIRS
|
||||
scan_result = await self._call_panw_api(
|
||||
content=response_text, is_response=True, metadata=metadata
|
||||
)
|
||||
# Check if guardrail should run for this request
|
||||
event_type: GuardrailEventHooks = GuardrailEventHooks.post_call
|
||||
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
|
||||
return response
|
||||
|
||||
action = scan_result.get("action", "block")
|
||||
category = scan_result.get("category", "unknown")
|
||||
try:
|
||||
# Extract response text
|
||||
response_text = self._extract_response_text(response)
|
||||
|
||||
if action == "allow":
|
||||
verbose_proxy_logger.debug(
|
||||
f"PANW Prisma AIRS: Response allowed (Category: {category})"
|
||||
if not response_text:
|
||||
verbose_proxy_logger.warning(
|
||||
"PANW Prisma AIRS: No response content found to scan"
|
||||
)
|
||||
return response
|
||||
|
||||
# Prepare metadata
|
||||
metadata = {
|
||||
"user": data.get("user") or "litellm_user",
|
||||
"model": data.get("model") or "unknown",
|
||||
}
|
||||
|
||||
# Scan response with PANW Prisma AIRS
|
||||
scan_result = await self._call_panw_api(
|
||||
content=response_text, is_response=True, metadata=metadata
|
||||
)
|
||||
|
||||
else:
|
||||
action = scan_result.get("action", "block")
|
||||
category = scan_result.get("category", "unknown")
|
||||
masked_text = self._get_masked_text(scan_result, is_response=True)
|
||||
|
||||
# If action is "allow", apply masking if available and allow through
|
||||
if action == "allow":
|
||||
if masked_text:
|
||||
self._apply_masking_to_response(response, masked_text)
|
||||
verbose_proxy_logger.info(
|
||||
f"PANW Prisma AIRS: Response allowed with masking (Category: {category})"
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.info(
|
||||
f"PANW Prisma AIRS: Response allowed (Category: {category})"
|
||||
)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return response
|
||||
|
||||
# Action is "block" - check if we should mask instead of blocking
|
||||
if masked_text and self.mask_response_content:
|
||||
self._apply_masking_to_response(response, masked_text)
|
||||
verbose_proxy_logger.warning(
|
||||
"PANW Prisma AIRS: Response blocked but masked instead (mask_response_content=True)"
|
||||
)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return response
|
||||
|
||||
# Block the response
|
||||
error_detail = self._build_error_detail(scan_result, is_response=True)
|
||||
verbose_proxy_logger.warning(
|
||||
f"PANW Prisma AIRS: {error_detail['error']['message']}"
|
||||
)
|
||||
raise HTTPException(status_code=400, detail=error_detail)
|
||||
|
||||
return response
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"PANW Prisma AIRS scan failed: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": {
|
||||
"message": "Security scan failed - response blocked for safety",
|
||||
"type": "guardrail_scan_error",
|
||||
"code": "panw_prisma_airs_scan_failed",
|
||||
"guardrail": self.guardrail_name,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
async def _scan_and_process_streaming_response(
|
||||
self,
|
||||
assembled_model_response: ModelResponse,
|
||||
request_data: dict,
|
||||
) -> Tuple[bool, ModelResponse]:
|
||||
"""
|
||||
Scan assembled streaming response and apply masking if needed.
|
||||
Returns (content_was_modified, response).
|
||||
"""
|
||||
content_was_modified = False
|
||||
response_text = self._extract_response_text(assembled_model_response)
|
||||
|
||||
if not response_text or not response_text.strip():
|
||||
verbose_proxy_logger.info("PANW Prisma AIRS: No content to scan in streaming response")
|
||||
return content_was_modified, assembled_model_response
|
||||
|
||||
# Prepare metadata and scan
|
||||
metadata = {
|
||||
"user": request_data.get("user") or "litellm_user",
|
||||
"model": request_data.get("model") or "unknown",
|
||||
}
|
||||
|
||||
scan_result = await self._call_panw_api(
|
||||
content=response_text, is_response=True, metadata=metadata
|
||||
)
|
||||
|
||||
action = scan_result.get("action", "block")
|
||||
category = scan_result.get("category", "unknown")
|
||||
masked_text = self._get_masked_text(scan_result, is_response=True)
|
||||
|
||||
# Handle scan results
|
||||
if action == "allow":
|
||||
if masked_text:
|
||||
self._apply_masking_to_response(assembled_model_response, masked_text)
|
||||
content_was_modified = True
|
||||
verbose_proxy_logger.info(
|
||||
f"PANW Prisma AIRS: Streaming response allowed with masking (Category: {category})"
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.info(
|
||||
f"PANW Prisma AIRS: Streaming response allowed (Category: {category})"
|
||||
)
|
||||
elif masked_text and self.mask_response_content:
|
||||
self._apply_masking_to_response(assembled_model_response, masked_text)
|
||||
content_was_modified = True
|
||||
verbose_proxy_logger.warning(
|
||||
"PANW Prisma AIRS: Streaming response blocked but masked instead (mask_response_content=True)"
|
||||
)
|
||||
else:
|
||||
error_detail = self._build_error_detail(scan_result, is_response=True)
|
||||
verbose_proxy_logger.warning(
|
||||
f"PANW Prisma AIRS: {error_detail['error']['message']}"
|
||||
)
|
||||
raise HTTPException(status_code=400, detail=error_detail)
|
||||
|
||||
return content_was_modified, assembled_model_response
|
||||
|
||||
@log_guardrail_information
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: Any,
|
||||
request_data: dict,
|
||||
):
|
||||
"""
|
||||
Process streaming response chunks and scan the assembled response.
|
||||
"""
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
|
||||
# Check if guardrail should run for this request
|
||||
from litellm.types.guardrails import GuardrailEventHooks as EventHooks
|
||||
|
||||
if not self.should_run_guardrail(
|
||||
data=request_data, event_type=EventHooks.post_call
|
||||
):
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
verbose_proxy_logger.info("PANW Prisma AIRS: Running post-call streaming scan")
|
||||
|
||||
all_chunks = []
|
||||
content_was_modified = False
|
||||
|
||||
try:
|
||||
# Collect all chunks
|
||||
async for chunk in response:
|
||||
all_chunks.append(chunk)
|
||||
|
||||
# Assemble complete response from chunks
|
||||
assembled_model_response = stream_chunk_builder(chunks=all_chunks)
|
||||
|
||||
if isinstance(assembled_model_response, ModelResponse):
|
||||
# Scan and process the assembled response
|
||||
content_was_modified, assembled_model_response = await self._scan_and_process_streaming_response(
|
||||
assembled_model_response, request_data
|
||||
)
|
||||
|
||||
# Add guardrail to applied guardrails header for observability
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=request_data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
|
||||
# Only use MockResponseIterator if content was modified
|
||||
# Otherwise, yield original chunks to preserve streaming behavior
|
||||
if content_was_modified:
|
||||
mock_response = MockResponseIterator(model_response=assembled_model_response)
|
||||
async for chunk in mock_response:
|
||||
yield chunk
|
||||
else:
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
else:
|
||||
# If not a ModelResponse, just yield original chunks
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"PANW Prisma AIRS streaming error: {str(e)}")
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": {
|
||||
"message": "Security scan failed - streaming response blocked for safety",
|
||||
"type": "guardrail_scan_error",
|
||||
"code": "panw_prisma_airs_scan_failed",
|
||||
"guardrail": self.guardrail_name,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
|
||||
|
|
|
|||
|
|
@ -16,8 +16,22 @@ class PanwPrismaAirsGuardrailConfigModel(GuardrailConfigModel):
|
|||
)
|
||||
|
||||
profile_name: str = Field(
|
||||
default="default",
|
||||
description="PANW Prisma AIRS security profile name. Required.",
|
||||
description="PANW Prisma AIRS security profile name configured in Strata Cloud Manager. Required.",
|
||||
)
|
||||
|
||||
mask_on_block: bool = Field(
|
||||
default=False,
|
||||
description="Backwards compatible flag that enables both request and response masking. When True, enables both mask_request_content and mask_response_content.",
|
||||
)
|
||||
|
||||
mask_request_content: bool = Field(
|
||||
default=False,
|
||||
description="Apply masking to prompts that would be blocked. When True, masked content is sent to the LLM instead of blocking the request.",
|
||||
)
|
||||
|
||||
mask_response_content: bool = Field(
|
||||
default=False,
|
||||
description="Apply masking to responses that would be blocked. When True, masked content is returned to the user instead of blocking the response.",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -56,23 +56,21 @@ class TestPanwAirsInitialization:
|
|||
guardrail_config = {"guardrail_name": "test_guardrail"}
|
||||
|
||||
with patch("litellm.logging_callback_manager.add_litellm_callback"):
|
||||
handler = initialize_guardrail(litellm_params, guardrail_config)
|
||||
handler = initialize_guardrail(litellm_params, guardrail_config)
|
||||
|
||||
assert isinstance(handler, PanwPrismaAirsHandler)
|
||||
assert handler.guardrail_name == "test_guardrail"
|
||||
|
||||
def test_missing_api_key_raises_error(self):
|
||||
"""Test that missing API key raises ValueError."""
|
||||
litellm_params = SimpleNamespace(
|
||||
profile_name="test_profile",
|
||||
api_base=None,
|
||||
default_on=True,
|
||||
api_key=None, # Missing API key
|
||||
)
|
||||
guardrail_config = {"guardrail_name": "test_guardrail"}
|
||||
|
||||
# Test direct handler initialization without api_key or env var
|
||||
with pytest.raises(ValueError, match="api_key is required"):
|
||||
initialize_guardrail(litellm_params, guardrail_config)
|
||||
PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
profile_name="test_profile",
|
||||
api_key=None, # No API key provided
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
def test_missing_profile_name_raises_error(self):
|
||||
"""Test that missing profile name raises ValueError."""
|
||||
|
|
@ -80,12 +78,12 @@ class TestPanwAirsInitialization:
|
|||
api_key="test_key",
|
||||
api_base=None,
|
||||
default_on=True,
|
||||
profile_name=None, # Missing profile name
|
||||
profile_name=None,
|
||||
)
|
||||
guardrail_config = {"guardrail_name": "test_guardrail"}
|
||||
|
||||
with pytest.raises(ValueError, match="profile_name is required"):
|
||||
initialize_guardrail(litellm_params, guardrail_config)
|
||||
initialize_guardrail(litellm_params, guardrail_config)
|
||||
|
||||
|
||||
class TestPanwAirsPromptScanning:
|
||||
|
|
@ -93,22 +91,20 @@ class TestPanwAirsPromptScanning:
|
|||
|
||||
@pytest.fixture
|
||||
def handler(self):
|
||||
"""Create test handler."""
|
||||
return PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
@pytest.fixture
|
||||
def user_api_key_dict(self):
|
||||
"""Mock user API key dict."""
|
||||
return UserAPIKeyAuth(api_key="test_key")
|
||||
|
||||
@pytest.fixture
|
||||
def safe_prompt_data(self):
|
||||
"""Safe prompt data."""
|
||||
return {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "What is the capital of France?"}],
|
||||
|
|
@ -117,7 +113,6 @@ class TestPanwAirsPromptScanning:
|
|||
|
||||
@pytest.fixture
|
||||
def malicious_prompt_data(self):
|
||||
"""Malicious prompt data."""
|
||||
return {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [
|
||||
|
|
@ -134,7 +129,6 @@ class TestPanwAirsPromptScanning:
|
|||
self, handler, user_api_key_dict, safe_prompt_data
|
||||
):
|
||||
"""Test that safe prompts are allowed."""
|
||||
# Mock PANW API response - allow
|
||||
mock_response = {"action": "allow", "category": "benign"}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", return_value=mock_response):
|
||||
|
|
@ -145,7 +139,6 @@ class TestPanwAirsPromptScanning:
|
|||
call_type="completion",
|
||||
)
|
||||
|
||||
# Should return None (not blocked)
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -153,7 +146,6 @@ class TestPanwAirsPromptScanning:
|
|||
self, handler, user_api_key_dict, malicious_prompt_data
|
||||
):
|
||||
"""Test that malicious prompts are blocked."""
|
||||
# Mock PANW API response - block
|
||||
mock_response = {"action": "block", "category": "malicious"}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", return_value=mock_response):
|
||||
|
|
@ -165,7 +157,6 @@ class TestPanwAirsPromptScanning:
|
|||
call_type="completion",
|
||||
)
|
||||
|
||||
# Verify exception details
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "PANW Prisma AI Security policy" in str(exc_info.value.detail)
|
||||
assert "malicious" in str(exc_info.value.detail)
|
||||
|
|
@ -182,17 +173,14 @@ class TestPanwAirsPromptScanning:
|
|||
call_type="completion",
|
||||
)
|
||||
|
||||
# Should return None (not blocked, no content to scan)
|
||||
assert result is None
|
||||
|
||||
def test_extract_text_from_messages(self, handler):
|
||||
"""Test text extraction from various message formats."""
|
||||
# Test simple string content
|
||||
messages = [{"role": "user", "content": "Hello world"}]
|
||||
text = handler._extract_text_from_messages(messages)
|
||||
assert text == "Hello world"
|
||||
|
||||
# Test complex content format
|
||||
messages = [
|
||||
{
|
||||
"role": "user",
|
||||
|
|
@ -205,7 +193,6 @@ class TestPanwAirsPromptScanning:
|
|||
text = handler._extract_text_from_messages(messages)
|
||||
assert text == "Analyze this image"
|
||||
|
||||
# Test multiple messages (should get last user message)
|
||||
messages = [
|
||||
{"role": "user", "content": "First message"},
|
||||
{"role": "assistant", "content": "Assistant response"},
|
||||
|
|
@ -220,27 +207,24 @@ class TestPanwAirsResponseScanning:
|
|||
|
||||
@pytest.fixture
|
||||
def handler(self):
|
||||
"""Create test handler."""
|
||||
return PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
@pytest.fixture
|
||||
def user_api_key_dict(self):
|
||||
"""Mock user API key dict."""
|
||||
return UserAPIKeyAuth(api_key="test_key")
|
||||
|
||||
@pytest.fixture
|
||||
def request_data(self):
|
||||
"""Request data."""
|
||||
return {"model": "gpt-3.5-turbo", "user": "test_user"}
|
||||
|
||||
@pytest.fixture
|
||||
def safe_response(self):
|
||||
"""Safe LLM response."""
|
||||
return ModelResponse(
|
||||
id="test_id",
|
||||
choices=[
|
||||
|
|
@ -256,7 +240,6 @@ class TestPanwAirsResponseScanning:
|
|||
|
||||
@pytest.fixture
|
||||
def harmful_response(self):
|
||||
"""Harmful LLM response."""
|
||||
return ModelResponse(
|
||||
id="test_id",
|
||||
choices=[
|
||||
|
|
@ -276,7 +259,6 @@ class TestPanwAirsResponseScanning:
|
|||
self, handler, user_api_key_dict, request_data, safe_response
|
||||
):
|
||||
"""Test that safe responses are allowed."""
|
||||
# Mock PANW API response - allow
|
||||
mock_response = {"action": "allow", "category": "benign"}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", return_value=mock_response):
|
||||
|
|
@ -286,7 +268,6 @@ class TestPanwAirsResponseScanning:
|
|||
response=safe_response,
|
||||
)
|
||||
|
||||
# Should return original response
|
||||
assert result == safe_response
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -294,7 +275,6 @@ class TestPanwAirsResponseScanning:
|
|||
self, handler, user_api_key_dict, request_data, harmful_response
|
||||
):
|
||||
"""Test that harmful responses are blocked."""
|
||||
# Mock PANW API response - block
|
||||
mock_response = {"action": "block", "category": "harmful"}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", return_value=mock_response):
|
||||
|
|
@ -305,7 +285,6 @@ class TestPanwAirsResponseScanning:
|
|||
response=harmful_response,
|
||||
)
|
||||
|
||||
# Verify exception details
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Response blocked by PANW Prisma AI Security policy" in str(
|
||||
exc_info.value.detail
|
||||
|
|
@ -318,12 +297,12 @@ class TestPanwAirsAPIIntegration:
|
|||
|
||||
@pytest.fixture
|
||||
def handler(self):
|
||||
"""Create test handler."""
|
||||
return PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -352,7 +331,6 @@ class TestPanwAirsAPIIntegration:
|
|||
@pytest.mark.asyncio
|
||||
async def test_api_error_handling(self, handler):
|
||||
"""Test API error handling (fail closed)."""
|
||||
# Mock the HTTP client to raise an exception
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs.panw_prisma_airs.get_async_httpx_client"
|
||||
) as mock_client:
|
||||
|
|
@ -362,18 +340,14 @@ class TestPanwAirsAPIIntegration:
|
|||
|
||||
result = await handler._call_panw_api("test content")
|
||||
|
||||
# Should fail closed (block) when API is unavailable
|
||||
assert result["action"] == "block"
|
||||
assert result["category"] == "api_error"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_api_response_handling(self, handler):
|
||||
"""Test handling of invalid API responses."""
|
||||
# Mock HTTP client to return invalid response (missing "action" field)
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"invalid": "response"
|
||||
} # Missing "action" field
|
||||
mock_response.json.return_value = {"invalid": "response"}
|
||||
mock_response.raise_for_status.return_value = None
|
||||
|
||||
with patch(
|
||||
|
|
@ -385,7 +359,6 @@ class TestPanwAirsAPIIntegration:
|
|||
|
||||
result = await handler._call_panw_api("test content")
|
||||
|
||||
# Should fail closed (block) when API response is invalid
|
||||
assert result["action"] == "block"
|
||||
assert result["category"] == "api_error"
|
||||
|
||||
|
|
@ -396,7 +369,6 @@ class TestPanwAirsAPIIntegration:
|
|||
content="", is_response=False, metadata={"user": "test", "model": "gpt-3.5"}
|
||||
)
|
||||
|
||||
# Should allow empty content without API call
|
||||
assert result["action"] == "allow"
|
||||
assert result["category"] == "empty"
|
||||
|
||||
|
|
@ -413,13 +385,13 @@ class TestPanwAirsConfiguration:
|
|||
mode="pre_call",
|
||||
api_key="test_key",
|
||||
profile_name="test_profile",
|
||||
api_base=None, # No api_base provided
|
||||
api_base=None,
|
||||
default_on=True,
|
||||
)
|
||||
guardrail_config = {"guardrail_name": "test"}
|
||||
|
||||
with patch("litellm.logging_callback_manager.add_litellm_callback"):
|
||||
handler = initialize_guardrail(litellm_params, guardrail_config)
|
||||
handler = initialize_guardrail(litellm_params, guardrail_config)
|
||||
|
||||
assert handler.api_base == "https://service.api.aisecurity.paloaltonetworks.com"
|
||||
|
||||
|
|
@ -439,7 +411,7 @@ class TestPanwAirsConfiguration:
|
|||
guardrail_config = {"guardrail_name": "test"}
|
||||
|
||||
with patch("litellm.logging_callback_manager.add_litellm_callback"):
|
||||
handler = initialize_guardrail(litellm_params, guardrail_config)
|
||||
handler = initialize_guardrail(litellm_params, guardrail_config)
|
||||
|
||||
assert handler.api_base == custom_base
|
||||
|
||||
|
|
@ -455,16 +427,589 @@ class TestPanwAirsConfiguration:
|
|||
api_base=None,
|
||||
default_on=True,
|
||||
)
|
||||
guardrail_config = {
|
||||
"guardrail_name": "test_guardrail",
|
||||
} # No guardrail_name
|
||||
guardrail_config = {"guardrail_name": "test_guardrail"}
|
||||
|
||||
with patch("litellm.logging_callback_manager.add_litellm_callback"):
|
||||
handler = initialize_guardrail(litellm_params, guardrail_config)
|
||||
handler = initialize_guardrail(litellm_params, guardrail_config)
|
||||
|
||||
assert handler.guardrail_name == "test_guardrail"
|
||||
|
||||
|
||||
class TestPanwAirsMaskingFunctionality:
|
||||
"""Test content masking features."""
|
||||
|
||||
def test_mask_on_block_backwards_compatibility(self):
|
||||
"""Test that mask_on_block enables both request and response masking."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
mask_on_block=True, # Should enable both masking flags
|
||||
)
|
||||
|
||||
# Verify both masking flags are enabled
|
||||
assert handler.mask_on_block is True
|
||||
assert handler.mask_request_content is True
|
||||
assert handler.mask_response_content is True
|
||||
|
||||
def test_mask_on_block_overrides_individual_flags(self):
|
||||
"""Test that mask_on_block=True overrides individual masking flags."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
mask_on_block=True,
|
||||
mask_request_content=False, # Should be overridden
|
||||
mask_response_content=False, # Should be overridden
|
||||
)
|
||||
|
||||
# mask_on_block should take precedence
|
||||
assert handler.mask_on_block is True
|
||||
assert handler.mask_request_content is True
|
||||
assert handler.mask_response_content is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_masking_on_block(self):
|
||||
"""Test that prompts are masked instead of blocked when mask_request_content=True."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
mask_request_content=True,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "Sensitive content"}],
|
||||
}
|
||||
|
||||
mock_response = {
|
||||
"action": "block",
|
||||
"category": "sensitive",
|
||||
"prompt_masked_data": {"data": "XXXXXXXXX content"},
|
||||
}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", return_value=mock_response):
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert data["messages"][0]["content"] == "XXXXXXXXX content"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_masking_with_content_list(self):
|
||||
"""Test that content lists are properly masked when mask_request_content=True."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
mask_request_content=True,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "My SSN is 123-45-6789"},
|
||||
{"type": "image", "url": "data:image/jpeg;base64,abc123"}
|
||||
]
|
||||
}],
|
||||
}
|
||||
|
||||
mock_response = {
|
||||
"action": "block",
|
||||
"category": "sensitive_data",
|
||||
"prompt_masked_data": {"data": "My SSN is XXXXXXXXXX"},
|
||||
}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", return_value=mock_response):
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
# Verify masking was applied to text content
|
||||
assert result is None
|
||||
assert isinstance(data["messages"][0]["content"], list)
|
||||
assert data["messages"][0]["content"][0]["type"] == "text"
|
||||
assert data["messages"][0]["content"][0]["text"] == "My SSN is XXXXXXXXXX"
|
||||
# Image should remain unchanged
|
||||
assert data["messages"][0]["content"][1]["type"] == "image"
|
||||
assert data["messages"][0]["content"][1]["url"] == "data:image/jpeg;base64,abc123"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_masking_on_block(self):
|
||||
"""Test that responses are masked instead of blocked when mask_response_content=True."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
mask_response_content=True,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
data = {"model": "gpt-3.5-turbo"}
|
||||
response = ModelResponse(
|
||||
id="test_id",
|
||||
choices=[
|
||||
Choices(
|
||||
index=0,
|
||||
message=Message(role="assistant", content="Sensitive response"),
|
||||
)
|
||||
],
|
||||
model="gpt-3.5-turbo",
|
||||
)
|
||||
|
||||
mock_response = {
|
||||
"action": "block",
|
||||
"category": "sensitive",
|
||||
"response_masked_data": {"data": "XXXXXXXXX response"},
|
||||
}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", return_value=mock_response):
|
||||
result = await handler.async_post_call_success_hook(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=response,
|
||||
)
|
||||
|
||||
assert result.choices[0].message.content == "XXXXXXXXX response"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fail_closed_on_api_error(self):
|
||||
"""Test fail-closed behavior on API errors (guardrail blocks on scan failures)."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth()
|
||||
data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [{"role": "user", "content": "Test content"}],
|
||||
}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", side_effect=Exception("API Error")):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=None,
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "Security scan failed" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
class TestPanwAirsAdvancedFeatures:
|
||||
"""Test advanced features: multi-choice, tool calls, streaming observability."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multi_choice_response_extraction(self):
|
||||
"""Test extraction of text from responses with multiple choices."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
# Create multi-choice response
|
||||
response = ModelResponse(
|
||||
id="test_id",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="First choice content", role="assistant"),
|
||||
),
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=1,
|
||||
message=Message(content="Second choice content", role="assistant"),
|
||||
),
|
||||
],
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion",
|
||||
)
|
||||
|
||||
extracted_text = handler._extract_response_text(response)
|
||||
assert "First choice content" in extracted_text
|
||||
assert "Second choice content" in extracted_text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_extraction(self):
|
||||
"""Test extraction of text from responses with tool calls."""
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Function
|
||||
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
# Create a proper ModelResponse with tool calls
|
||||
response = ModelResponse(
|
||||
id="test_id",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="tool_calls",
|
||||
index=0,
|
||||
message=Message(
|
||||
content=None,
|
||||
role="assistant",
|
||||
tool_calls=[
|
||||
ChatCompletionMessageToolCall(
|
||||
id="call_123",
|
||||
type="function",
|
||||
function=Function(
|
||||
name="get_weather",
|
||||
arguments='{"location": "San Francisco", "ssn": "123-45-6789"}'
|
||||
)
|
||||
)
|
||||
]
|
||||
),
|
||||
),
|
||||
],
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion",
|
||||
)
|
||||
|
||||
extracted_text = handler._extract_response_text(response)
|
||||
assert "123-45-6789" in extracted_text
|
||||
assert "San Francisco" in extracted_text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_call_masking(self):
|
||||
"""Test masking of tool call arguments when blocked."""
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Function
|
||||
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
mask_response_content=True,
|
||||
)
|
||||
|
||||
# Create a proper ModelResponse with tool calls
|
||||
response = ModelResponse(
|
||||
id="test_id",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="tool_calls",
|
||||
index=0,
|
||||
message=Message(
|
||||
content=None,
|
||||
role="assistant",
|
||||
tool_calls=[
|
||||
ChatCompletionMessageToolCall(
|
||||
id="call_123",
|
||||
type="function",
|
||||
function=Function(
|
||||
name="get_weather",
|
||||
arguments='{"location": "San Francisco", "ssn": "123-45-6789"}'
|
||||
)
|
||||
)
|
||||
]
|
||||
),
|
||||
),
|
||||
],
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion",
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test_key")
|
||||
data = {"messages": [{"role": "user", "content": "test"}], "model": "gpt-4"}
|
||||
|
||||
# Mock PANW API to return block with masking
|
||||
mock_scan_result = {
|
||||
"action": "block",
|
||||
"category": "sensitive_data",
|
||||
"response_masked_data": {
|
||||
"data": '{"location": "San Francisco", "ssn": "XXXXXXXXXX"}'
|
||||
}
|
||||
}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.return_value = mock_scan_result
|
||||
|
||||
result = await handler.async_post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict, response=response, data=data
|
||||
)
|
||||
|
||||
# Verify arguments were masked
|
||||
assert result.choices[0].message.tool_calls[0].function.arguments == '{"location": "San Francisco", "ssn": "XXXXXXXXXX"}'
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multi_choice_masking(self):
|
||||
"""Test masking applied to all choices in multi-choice response."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
mask_response_content=True,
|
||||
)
|
||||
|
||||
# Create multi-choice response
|
||||
response = ModelResponse(
|
||||
id="test_id",
|
||||
choices=[
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(content="SSN is 123-45-6789", role="assistant"),
|
||||
),
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=1,
|
||||
message=Message(content="Another SSN: 987-65-4321", role="assistant"),
|
||||
),
|
||||
],
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion",
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test_key")
|
||||
data = {"messages": [{"role": "user", "content": "test"}], "model": "gpt-4"}
|
||||
|
||||
mock_scan_result = {
|
||||
"action": "block",
|
||||
"category": "sensitive_data",
|
||||
"response_masked_data": {"data": "SSN is XXXXXXXXXX"}
|
||||
}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.return_value = mock_scan_result
|
||||
|
||||
result = await handler.async_post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict, response=response, data=data
|
||||
)
|
||||
|
||||
# Verify all choices were masked
|
||||
assert result.choices[0].message.content == "SSN is XXXXXXXXXX"
|
||||
assert result.choices[1].message.content == "SSN is XXXXXXXXXX"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_hook_adds_guardrail_header(self):
|
||||
"""Test that streaming hook adds guardrail to applied guardrails header."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
api_base="https://test.panw.com/api",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="test_key")
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": "test"}],
|
||||
"model": "gpt-4"
|
||||
}
|
||||
|
||||
# Create mock streaming chunks
|
||||
from litellm.types.utils import StreamingChoices, Delta
|
||||
|
||||
mock_chunks = [
|
||||
ModelResponse(
|
||||
id="test_id",
|
||||
choices=[StreamingChoices(delta=Delta(content="Hello", role="assistant"), finish_reason=None, index=0)],
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
),
|
||||
ModelResponse(
|
||||
id="test_id",
|
||||
choices=[StreamingChoices(delta=Delta(content=" world", role="assistant"), finish_reason="stop", index=0)],
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
),
|
||||
]
|
||||
|
||||
async def mock_response_iter():
|
||||
for chunk in mock_chunks:
|
||||
yield chunk
|
||||
|
||||
mock_scan_result = {"action": "allow", "category": "safe"}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
|
||||
with patch("litellm.proxy.common_utils.callback_utils.add_guardrail_to_applied_guardrails_header") as mock_header:
|
||||
mock_api.return_value = mock_scan_result
|
||||
|
||||
chunks_received = []
|
||||
async for chunk in handler.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_response_iter(),
|
||||
request_data=request_data,
|
||||
):
|
||||
chunks_received.append(chunk)
|
||||
|
||||
# Verify header function was called
|
||||
assert mock_header.called
|
||||
mock_header.assert_called_once_with(
|
||||
request_data=request_data,
|
||||
guardrail_name="test_panw_airs"
|
||||
)
|
||||
|
||||
|
||||
class TestTextCompletionSupport:
|
||||
"""Test support for text completion (non-chat) requests."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_completion_prompt_extraction(self):
|
||||
"""Test that guardrail can extract and scan text completion prompts."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test_key", user_id="test_user", team_id="test_team"
|
||||
)
|
||||
|
||||
# Text completion request (no messages, just prompt)
|
||||
data = {
|
||||
"prompt": "Complete this sentence: AI security is",
|
||||
"model": "gpt-3.5-turbo-instruct",
|
||||
"max_tokens": 50
|
||||
}
|
||||
|
||||
mock_scan_result = {"action": "allow", "category": "safe"}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.return_value = mock_scan_result
|
||||
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type="text_completion",
|
||||
)
|
||||
|
||||
# Verify API was called with the prompt text
|
||||
mock_api.assert_called_once()
|
||||
call_args = mock_api.call_args
|
||||
assert call_args.kwargs["content"] == "Complete this sentence: AI security is"
|
||||
assert call_args.kwargs["is_response"] is False
|
||||
|
||||
# Verify request was allowed through
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_completion_with_masking(self):
|
||||
"""Test that masking works with text completion prompts."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
mask_request_content=True,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test_key", user_id="test_user", team_id="test_team"
|
||||
)
|
||||
|
||||
data = {
|
||||
"prompt": "Send money to account 123-456-7890",
|
||||
"model": "gpt-3.5-turbo-instruct",
|
||||
}
|
||||
|
||||
# Simulate PANW blocking but providing masked content
|
||||
mock_scan_result = {
|
||||
"action": "block",
|
||||
"category": "dlp",
|
||||
"prompt_masked_data": {"data": "Send money to account XXXXXXXXXX"}
|
||||
}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.return_value = mock_scan_result
|
||||
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type="text_completion",
|
||||
)
|
||||
|
||||
# Verify the prompt was masked
|
||||
assert result is None
|
||||
assert data["prompt"] == "Send money to account XXXXXXXXXX"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_text_completion_with_list_prompts(self):
|
||||
"""Test that guardrail handles batch text completion (list of prompts)."""
|
||||
handler = PanwPrismaAirsHandler(
|
||||
guardrail_name="test_panw_airs",
|
||||
api_key="test_api_key",
|
||||
profile_name="test_profile",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test_key", user_id="test_user", team_id="test_team"
|
||||
)
|
||||
|
||||
# Batch completion request
|
||||
data = {
|
||||
"prompt": ["Tell me a joke", "What is AI?"],
|
||||
"model": "gpt-3.5-turbo-instruct",
|
||||
}
|
||||
|
||||
mock_scan_result = {"action": "allow", "category": "safe"}
|
||||
|
||||
with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api:
|
||||
mock_api.return_value = mock_scan_result
|
||||
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=MagicMock(),
|
||||
data=data,
|
||||
call_type="text_completion",
|
||||
)
|
||||
|
||||
# Verify API was called with joined prompts
|
||||
mock_api.assert_called_once()
|
||||
call_args = mock_api.call_args
|
||||
assert "Tell me a joke" in call_args.kwargs["content"]
|
||||
assert "What is AI?" in call_args.kwargs["content"]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run tests
|
||||
pytest.main([__file__, "-v"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue