From c471bf1f16c2ba484c614e906800d6a54b16865a Mon Sep 17 00:00:00 2001 From: Jason Roberts <51415896+jroberts2600@users.noreply.github.com> Date: Sat, 18 Oct 2025 15:57:51 -0500 Subject: [PATCH] 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 --- .../docs/proxy/guardrails/panw_prisma_airs.md | 137 +++- .../panw_prisma_airs/__init__.py | 4 +- .../panw_prisma_airs/panw_prisma_airs.py | 539 ++++++++++++--- .../guardrail_hooks/panw_prisma_airs.py | 18 +- .../guardrail_hooks/test_panw_prisma_airs.py | 645 ++++++++++++++++-- 5 files changed, 1196 insertions(+), 147 deletions(-) diff --git a/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md b/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md index 20cbc60a3e9..97f3e7efe54 100644 --- a/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md +++ b/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md @@ -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 + + + + +**Request:** +```json +{ + "messages": [ + {"role": "user", "content": "My credit card is 4929-3813-3266-4295"} + ] +} +``` + +**Response:** ❌ **Blocked with 400 error** + + + + +**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** + + + + +#### 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 \ No newline at end of file diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/__init__.py index f7c05fb8c45..e69077401e9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/__init__.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index 8ca29506771..c6d2543f9f4 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -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"]]: diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/panw_prisma_airs.py b/litellm/types/proxy/guardrails/guardrail_hooks/panw_prisma_airs.py index 2d728f7076c..c23c1542674 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/panw_prisma_airs.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/panw_prisma_airs.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py index 12d84f9530e..4d351c80da1 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py @@ -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"])