From 5733f711fd174ad34ef0e516440e3265fa4d2548 Mon Sep 17 00:00:00 2001 From: Jason Roberts <51415896+jroberts2600@users.noreply.github.com> Date: Thu, 6 Nov 2025 21:34:37 -0600 Subject: [PATCH] feat(guardrails): panw prisma airs guardrail deduplication and enhanced session tracking (#16273) * feat(guardrails): Add deduplication and session tracking - Implement deduplication logic to prevent duplicate scans (via call_id; add _check_and_mark_scanned) caused by LiteLLM callback system - Add session tracking using litellm_trace_id as AI Session ID for Prisma AIRS SCM logging - Extract helper methods _extract_prompt_from_request maintainability - Use httpxSpecialProvider import (LoggingCallback -> GuardrailCallback) - Add comprehensive tests for deduplication and session tracking (7 new tests) - Update documentation with multi-turn conversation tracking examples * docs: update PANW Prisma AIRS multi-turn conversation example to use industry-standard terminology - Clearer example for conversation tracking - Updated terminology from 'AI Session ID' to 'Prisma AIRS AI Session ID' for clarity * fix: remove unused asyncio import * fix: correct mypy type ignore comment --- .../docs/proxy/guardrails/panw_prisma_airs.md | 73 ++- .../panw_prisma_airs/panw_prisma_airs.py | 105 +++- .../guardrail_hooks/test_panw_prisma_airs.py | 451 ++++++++++++++++-- 3 files changed, 560 insertions(+), 69 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 e1d6ddf5928..edf2a05d24c 100644 --- a/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md +++ b/docs/my-website/docs/proxy/guardrails/panw_prisma_airs.md @@ -4,12 +4,12 @@ import TabItem from '@theme/TabItem'; # PANW Prisma AIRS -LiteLLM supports PANW Prisma AIRS (AI Runtime Security) guardrails via the [Prisma AIRS Scan API](https://pan.dev/prisma-airs/api/airuntimesecurity/scan-sync-request/). This integration provides **Security-as-Code** for AI applications using Palo Alto Networks' AI security platform. +LiteLLM supports PANW Prisma AIRS (AI Runtime Security) guardrails via the [Prisma AIRS Scan API](https://pan.dev/prisma-airs/api/airuntimesecurity/airuntimesecurityapi//). This integration provides **Security-as-Code** for AI applications using Palo Alto Networks' AI security platform. ## Features - ✅ **Real-time prompt injection detection** -- ✅ **Malicious content filtering** +- ✅ **Malicious URL detection** - ✅ **Data loss prevention (DLP)** - ✅ **Sensitive content masking** - Automatically mask PII, credit cards, SSNs instead of blocking - ✅ **Comprehensive threat detection** for AI models and datasets @@ -17,6 +17,7 @@ LiteLLM supports PANW Prisma AIRS (AI Runtime Security) guardrails via the [Pris - ✅ **Synchronous scanning** with immediate response - ✅ **Configurable security profiles** - ✅ **Streaming support** - Real-time masking for streaming responses +- ✅ **Multi-turn conversation tracking** - Automatic session grouping in Prisma AIRS SCM logs - ✅ **Fail-closed security** - Blocks requests if PANW API is unavailable (maximum security) ## Quick Start @@ -237,6 +238,74 @@ You can override guardrail settings on a per-request basis using the `metadata` - **Note:** If your API key is not linked to a profile, you must provide `profile_name` or `profile_id` ::: +## Multi-Turn Conversation Tracking + +PANW Prisma AIRS automatically tracks multi-turn conversations using LiteLLM's `litellm_trace_id`. This enables you to: + +- **Group related requests** - All requests in a conversation share the same AI Session ID in Prisma AIRS SCM logs +- **Track conversation context** - See the full history of prompts and responses for a user session +- **Analyze attack patterns** - Identify sophisticated multi-turn attacks across conversation history + +### How It Works + +LiteLLM automatically generates a unique `litellm_trace_id` for each conversation session. The PANW guardrail uses this as the PANW transaction ID (which maps to "AI Session ID" in Strata Cloud Manager): + +``` +Conversation Session: litellm_trace_id = "abc-123-def-456" + +Turn 1 (User): "What's the capital of France?" + → Scan ID: scan_001 | Prisma AIRS AI Session ID: abc-123-def-456 + +Turn 2 (Assistant): "Paris is the capital of France." + → Scan ID: scan_002 | Prisma AIRS AI Session ID: abc-123-def-456 + +Turn 3 (User): "What's the population?" + → Scan ID: scan_003 | Prisma AIRS AI Session ID: abc-123-def-456 + +Turn 4 (Assistant): "Paris has approximately 2.1 million residents." + → Scan ID: scan_004 | Prisma AIRS AI Session ID: abc-123-def-456 +``` + +All scans appear under the same AI Session ID in Prisma AIRS logs, making it easy to: +- Review complete conversation history (all 4 turns grouped together) +- Identify patterns across multiple turns +- Correlate security events within a session +- Track the flow of user prompts and AI responses + +### Session Tracking + +LiteLLM automatically generates a unique `litellm_trace_id` for each request, which the PANW guardrail uses as the AI Session ID in Strata Cloud Manager. All prompt and response scans for a request are automatically grouped under the same session. + +#### Custom Session IDs (Per-App Tracking) + +You can provide your own `litellm_trace_id` to track sessions on a per-app or per-conversation basis: + +```bash +curl -X POST http://localhost:4000/v1/chat/completions \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer sk-1234" \ + -d '{ + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "capital of France"}], + "litellm_trace_id": "my-app-session-123", # Custom AI Session ID + "metadata": { + "profile_name": "dev-allow-all-profile", # Override security profile + "user_ip": "192.168.1.1", # Track user IP + "app_name": "eng" # Custom app identifier + }, + "guardrails": ["panw-prisma-airs-pre-guard", "panw-prisma-airs-post-guard"] + }' +``` + +**Result in PANW SCM:** +- AI Session ID: `my-app-session-123` +- All prompt and response scans will be grouped under this custom session ID +- Perfect for tracking multi-turn conversations or per-application sessions + +:::tip Viewing Sessions in Prisma AIRS SCM Logs +In Strata Cloud Manager, navigate to **AI Runtime > Sessions** to view all AI Session IDs and their associated scans. Click on a session to see the complete conversation history with security analysis. +::: + ## Environment Variables ```bash 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 814fb6deb2c..36fdfecaab8 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 @@ -205,11 +205,10 @@ class PanwPrismaAirsHandler(CustomGuardrail): if not content.strip(): return {"action": "allow", "category": "empty"} - # Use litellm_call_id if provided, fallback to generating UUID - litellm_call_id = call_id or str(uuid.uuid4()) - - # Build transaction ID with prompt/response prefix + full call_id for correlation - transaction_id = f"litellm-{'resp' if is_response else 'req'}-{litellm_call_id}" + # Use litellm_trace_id as Prisma AIRS AI Session ID for session grouping + transaction_id = metadata.get("litellm_trace_id") if metadata else None + if not transaction_id: + transaction_id = call_id or str(uuid.uuid4()) # Build Prisma AIRS API metadata # Handle app_name: LiteLLM by default, or LiteLLM-{user_app_name} if user provides one @@ -264,7 +263,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): payload["ai_profile"] = ai_profile if is_response: - payload["metadata"]["is_response"] = True # type: ignore[index] + payload["metadata"]["is_response"] = True # type: ignore[call-overload, index] headers = { "Content-Type": "application/json", @@ -275,7 +274,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): try: # Use LiteLLM's async HTTP client async_client = get_async_httpx_client( - llm_provider=httpxSpecialProvider.LoggingCallback + llm_provider=httpxSpecialProvider.GuardrailCallback ) response = await async_client.post( @@ -496,8 +495,69 @@ class PanwPrismaAirsHandler(CustomGuardrail): if "app_name" in user_metadata: metadata["app_name"] = user_metadata["app_name"] + # Include litellm_trace_id for session tracking + if data.get("litellm_trace_id"): + metadata["litellm_trace_id"] = data["litellm_trace_id"] + return metadata + def _check_and_mark_scanned(self, data: dict, scan_type: str) -> bool: + """ + Check if request has already been scanned and mark it as scanned. + + Args: + data: Request data dictionary + scan_type: Type of scan ('pre', 'post', 'streaming') + + Returns: + True if already scanned (should skip), False if needs scanning + """ + call_id = data.get("litellm_call_id") + if not call_id: + call_id = str(uuid.uuid4()) + data["litellm_call_id"] = call_id + + scan_key = f"_panw_{scan_type}_scanned_{call_id}" + litellm_metadata = data.setdefault("litellm_metadata", {}) + + if litellm_metadata.get(scan_key): + verbose_proxy_logger.debug( + f"PANW Prisma AIRS: Skipping duplicate {scan_type}-call scan" + ) + return True # Already scanned + + litellm_metadata[scan_key] = True + return False # Needs scanning + + def _extract_prompt_from_request(self, data: dict) -> str: + """ + Extract prompt text from request data. + + Handles both chat completion (messages) and text completion (prompt) formats. + + Args: + data: Request data dictionary + + Returns: + Extracted prompt text, or empty string if not found + """ + # Extract from messages (chat 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 = "" + + return prompt_text + @log_guardrail_information async def async_pre_call_hook( self, @@ -534,21 +594,14 @@ class PanwPrismaAirsHandler(CustomGuardrail): 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) + # Prevent duplicate scans by checking if already processed + if self._check_and_mark_scanned(data, "pre"): + return data - # 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 = "" + try: + # Extract prompt text from request + prompt_text = self._extract_prompt_from_request(data) + messages = data.get("messages", []) # Keep for masking operations if not prompt_text: verbose_proxy_logger.warning( @@ -659,6 +712,10 @@ class PanwPrismaAirsHandler(CustomGuardrail): if self.should_run_guardrail(data=data, event_type=event_type) is not True: return response + # Prevent duplicate scans by checking if already processed + if self._check_and_mark_scanned(data, "post"): + return response + try: # Extract response text response_text = self._extract_response_text(response) @@ -819,6 +876,12 @@ class PanwPrismaAirsHandler(CustomGuardrail): yield chunk return + # Prevent duplicate scans by checking if already processed + if self._check_and_mark_scanned(request_data, "streaming"): + async for chunk in response: + yield chunk + return + verbose_proxy_logger.info("PANW Prisma AIRS: Running post-call streaming scan") all_chunks = [] 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 101a1d8275a..77a7daf0de4 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 @@ -55,7 +55,7 @@ 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" @@ -70,7 +70,7 @@ class TestPanwAirsInitialization: profile_name="test_profile", api_key=None, # No API key provided default_on=True, - ) + ) def test_api_key_with_linked_profile(self): """Test initialization with API key that has a linked profile (no explicit profile_name needed).""" @@ -82,7 +82,9 @@ class TestPanwAirsInitialization: default_on=True, ) assert handler.api_key == "test_api_key_with_linked_profile" - assert handler.profile_name is None # Should be None, PANW API will use linked profile + assert ( + handler.profile_name is None + ) # Should be None, PANW API will use linked profile class TestPanwAirsPromptScanning: @@ -390,7 +392,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 == "https://service.api.aisecurity.paloaltonetworks.com" @@ -410,7 +412,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 @@ -429,7 +431,7 @@ class TestPanwAirsConfiguration: 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" @@ -521,13 +523,15 @@ class TestPanwAirsMaskingFunctionality: 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"} - ] - }], + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "My SSN is 123-45-6789"}, + {"type": "image", "url": "data:image/jpeg;base64,abc123"}, + ], + } + ], } mock_response = { @@ -551,7 +555,9 @@ class TestPanwAirsMaskingFunctionality: 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" + assert ( + data["messages"][0]["content"][1]["url"] == "data:image/jpeg;base64,abc123" + ) @pytest.mark.asyncio async def test_response_masking_on_block(self): @@ -610,7 +616,9 @@ class TestPanwAirsMaskingFunctionality: "messages": [{"role": "user", "content": "Test content"}], } - with patch.object(handler, "_call_panw_api", side_effect=Exception("API Error")): + 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, @@ -665,7 +673,7 @@ class TestPanwAirsAdvancedFeatures: 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", @@ -690,10 +698,10 @@ class TestPanwAirsAdvancedFeatures: type="function", function=Function( name="get_weather", - arguments='{"location": "San Francisco", "ssn": "123-45-6789"}' - ) + arguments='{"location": "San Francisco", "ssn": "123-45-6789"}', + ), ) - ] + ], ), ), ], @@ -710,7 +718,7 @@ class TestPanwAirsAdvancedFeatures: 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", @@ -736,10 +744,10 @@ class TestPanwAirsAdvancedFeatures: type="function", function=Function( name="get_weather", - arguments='{"location": "San Francisco", "ssn": "123-45-6789"}' - ) + arguments='{"location": "San Francisco", "ssn": "123-45-6789"}', + ), ) - ] + ], ), ), ], @@ -757,18 +765,23 @@ class TestPanwAirsAdvancedFeatures: "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: + 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"}' + 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): @@ -794,7 +807,9 @@ class TestPanwAirsAdvancedFeatures: Choices( finish_reason="stop", index=1, - message=Message(content="Another SSN: 987-65-4321", role="assistant"), + message=Message( + content="Another SSN: 987-65-4321", role="assistant" + ), ), ], created=1234567890, @@ -808,12 +823,14 @@ class TestPanwAirsAdvancedFeatures: mock_scan_result = { "action": "block", "category": "sensitive_data", - "response_masked_data": {"data": "SSN is XXXXXXXXXX"} + "response_masked_data": {"data": "SSN is XXXXXXXXXX"}, } - with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: + 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 ) @@ -836,23 +853,35 @@ class TestPanwAirsAdvancedFeatures: user_api_key_dict = UserAPIKeyAuth(api_key="test_key") request_data = { "messages": [{"role": "user", "content": "test"}], - "model": "gpt-4" + "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)], + 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)], + choices=[ + StreamingChoices( + delta=Delta(content=" world", role="assistant"), + finish_reason="stop", + index=0, + ) + ], created=1234567890, model="gpt-4", object="chat.completion.chunk", @@ -865,10 +894,14 @@ class TestPanwAirsAdvancedFeatures: 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: + 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, @@ -880,8 +913,7 @@ class TestPanwAirsAdvancedFeatures: # Verify header function was called assert mock_header.called mock_header.assert_called_once_with( - request_data=request_data, - guardrail_name="test_panw_airs" + request_data=request_data, guardrail_name="test_panw_airs" ) @@ -906,12 +938,14 @@ class TestTextCompletionSupport: data = { "prompt": "Complete this sentence: AI security is", "model": "gpt-3.5-turbo-instruct", - "max_tokens": 50 + "max_tokens": 50, } mock_scan_result = {"action": "allow", "category": "safe"} - with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: + 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( @@ -924,7 +958,9 @@ class TestTextCompletionSupport: # 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["content"] == "Complete this sentence: AI security is" + ) assert call_args.kwargs["is_response"] is False # Verify request was allowed through @@ -954,10 +990,12 @@ class TestTextCompletionSupport: mock_scan_result = { "action": "block", "category": "dlp", - "prompt_masked_data": {"data": "Send money to account XXXXXXXXXX"} + "prompt_masked_data": {"data": "Send money to account XXXXXXXXXX"}, } - with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: + 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( @@ -993,7 +1031,9 @@ class TestTextCompletionSupport: mock_scan_result = {"action": "allow", "category": "safe"} - with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: + 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( @@ -1010,5 +1050,324 @@ class TestTextCompletionSupport: assert "What is AI?" in call_args.kwargs["content"] +class TestPanwAirsDeduplication: + """Test deduplication of callback invocations.""" + + @pytest.mark.asyncio + async def test_duplicate_pre_call_scan_prevented(self): + """Test that duplicate pre-call scans are prevented.""" + 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") + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Test"}], + "litellm_call_id": "test-call-123", + } + + mock_response = {"action": "allow", "category": "benign"} + + with patch.object( + handler, "_call_panw_api", return_value=mock_response + ) as mock_api: + # First call - should scan + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=None, + data=data, + call_type="completion", + ) + assert mock_api.call_count == 1 + + # Second call with same call_id - should skip + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=None, + data=data, + call_type="completion", + ) + # Still 1 - no additional scan + assert mock_api.call_count == 1 + + @pytest.mark.asyncio + async def test_duplicate_post_call_scan_prevented(self): + """Test that duplicate post-call scans are prevented.""" + 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") + data = { + "model": "gpt-3.5-turbo", + "litellm_call_id": "test-call-456", + } + response = ModelResponse( + id="test_id", + choices=[ + Choices( + index=0, + message=Message(role="assistant", content="Test response"), + ) + ], + model="gpt-3.5-turbo", + ) + + mock_response = {"action": "allow", "category": "benign"} + + with patch.object( + handler, "_call_panw_api", return_value=mock_response + ) as mock_api: + # First call + await handler.async_post_call_success_hook( + data=data, + user_api_key_dict=user_api_key_dict, + response=response, + ) + assert mock_api.call_count == 1 + + # Second call - should skip + await handler.async_post_call_success_hook( + data=data, + user_api_key_dict=user_api_key_dict, + response=response, + ) + assert mock_api.call_count == 1 + + @pytest.mark.asyncio + async def test_duplicate_streaming_scan_prevented(self): + """Test that duplicate streaming scans are prevented.""" + 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") + request_data = { + "model": "gpt-3.5-turbo", + "litellm_call_id": "test-call-789", + "messages": [{"role": "user", "content": "test"}], + } + + # 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", + ), + ] + + 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: + mock_api.return_value = mock_scan_result + + # First call - should scan + 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) + + assert mock_api.call_count == 1 + + # Second call with same call_id - should skip + 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, + ): + pass + + # Still 1 - no additional scan + assert mock_api.call_count == 1 + + +class TestPanwAirsSessionTracking: + """Test session tracking with litellm_trace_id.""" + + @pytest.mark.asyncio + async def test_litellm_trace_id_used_as_transaction_id(self): + """Test that litellm_trace_id is used as PANW transaction ID.""" + handler = PanwPrismaAirsHandler( + guardrail_name="test_panw_airs", + api_key="test_api_key", + profile_name="test_profile", + default_on=True, + ) + + trace_id = "abc-123-def-456" + metadata = { + "user": "test_user", + "model": "gpt-4", + "litellm_trace_id": trace_id, + } + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs.panw_prisma_airs.get_async_httpx_client" + ) as mock_client: + mock_async_client = AsyncMock() + mock_response = MagicMock() + mock_response.json.return_value = {"action": "allow", "category": "benign"} + mock_response.raise_for_status.return_value = None + mock_async_client.post = AsyncMock(return_value=mock_response) + mock_client.return_value = mock_async_client + + await handler._call_panw_api( + content="Test content", + is_response=False, + metadata=metadata, + ) + + # Verify tr_id in API payload matches trace_id + call_args = mock_async_client.post.call_args + payload = call_args.kwargs["json"] + assert payload["tr_id"] == trace_id + + @pytest.mark.asyncio + async def test_fallback_to_call_id_when_trace_id_missing(self): + """Test fallback to call_id when litellm_trace_id is missing.""" + handler = PanwPrismaAirsHandler( + guardrail_name="test_panw_airs", + api_key="test_api_key", + profile_name="test_profile", + default_on=True, + ) + + call_id = "fallback-call-789" + metadata = { + "user": "test_user", + "model": "gpt-4", + # No litellm_trace_id + } + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs.panw_prisma_airs.get_async_httpx_client" + ) as mock_client: + mock_async_client = AsyncMock() + mock_response = MagicMock() + mock_response.json.return_value = {"action": "allow", "category": "benign"} + mock_response.raise_for_status.return_value = None + mock_async_client.post = AsyncMock(return_value=mock_response) + mock_client.return_value = mock_async_client + + await handler._call_panw_api( + content="Test content", + is_response=False, + metadata=metadata, + call_id=call_id, + ) + + # Verify tr_id falls back to call_id + call_args = mock_async_client.post.call_args + payload = call_args.kwargs["json"] + assert payload["tr_id"] == call_id + + @pytest.mark.asyncio + async def test_trace_id_extraction_from_request_data(self): + """Test that litellm_trace_id is extracted from request data.""" + handler = PanwPrismaAirsHandler( + guardrail_name="test_panw_airs", + api_key="test_api_key", + profile_name="test_profile", + default_on=True, + ) + + trace_id = "session-xyz-789" + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Test"}], + "litellm_trace_id": trace_id, + } + + # Extract metadata + metadata = handler._prepare_metadata_from_request(data) + + # Verify trace_id is included in metadata + assert "litellm_trace_id" in metadata + assert metadata["litellm_trace_id"] == trace_id + + @pytest.mark.asyncio + async def test_same_trace_id_for_prompt_and_response(self): + """Test that prompt and response scans use the same trace_id.""" + handler = PanwPrismaAirsHandler( + guardrail_name="test_panw_airs", + api_key="test_api_key", + profile_name="test_profile", + default_on=True, + ) + + trace_id = "conversation-session-123" + + with patch( + "litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs.panw_prisma_airs.get_async_httpx_client" + ) as mock_client: + mock_async_client = AsyncMock() + mock_response = MagicMock() + mock_response.json.return_value = {"action": "allow", "category": "benign"} + mock_response.raise_for_status.return_value = None + mock_async_client.post = AsyncMock(return_value=mock_response) + mock_client.return_value = mock_async_client + + # Prompt scan + await handler._call_panw_api( + content="User prompt", + is_response=False, + metadata={ + "litellm_trace_id": trace_id, + "user": "test", + "model": "gpt-4", + }, + ) + prompt_payload = mock_async_client.post.call_args.kwargs["json"] + prompt_tr_id = prompt_payload["tr_id"] + + # Response scan + await handler._call_panw_api( + content="Assistant response", + is_response=True, + metadata={ + "litellm_trace_id": trace_id, + "user": "test", + "model": "gpt-4", + }, + ) + response_payload = mock_async_client.post.call_args.kwargs["json"] + response_tr_id = response_payload["tr_id"] + + # Both should use the same trace_id + assert prompt_tr_id == trace_id + assert response_tr_id == trace_id + assert prompt_tr_id == response_tr_id + + if __name__ == "__main__": pytest.main([__file__, "-v"])