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:
Jason Roberts 2025-10-18 15:57:51 -05:00 • committed by GitHub
parent 645f84c02e
commit c471bf1f16
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 1196 additions and 147 deletions

View file

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

View file

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

View file

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

View file

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

View file

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