Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py: 8%

828 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1#!/usr/bin/env python3 

2""" 

3Palo Alto Networks Prisma AI Runtime Security (AIRS) Guardrail Integration for LiteLLM 

4 

5Provides real-time threat detection, DLP, URL filtering, content masking, and policy enforcement for AI applications. 

6""" 

7 

8import json 

9import os 

10import re 

11from collections.abc import AsyncIterable, Mapping, Sequence 

12from datetime import datetime 

13from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias 

14from urllib.parse import urlparse 

15 

16import httpx 

17from fastapi import HTTPException 

18from pydantic import BaseModel, ConfigDict, ValidationError, field_validator 

19 

20from litellm._logging import verbose_proxy_logger 

21from litellm._uuid import uuid 

22from litellm.caching import DualCache 

23from litellm.integrations.custom_guardrail import ( 

24 CustomGuardrail, 

25 log_guardrail_information, 

26) 

27from litellm.llms.base_llm.guardrail_translation.utils import ( 

28 effective_scan_only_tool_results_for_guardrail, 

29) 

30from litellm.llms.custom_httpx.http_handler import ( 

31 AsyncHTTPHandler, 

32 get_async_httpx_client, 

33 httpxSpecialProvider, 

34) 

35from litellm.proxy._types import UserAPIKeyAuth 

36from litellm.proxy.common_utils.callback_utils import ( 

37 add_guardrail_scan_id, 

38 add_guardrail_to_applied_guardrails_header, 

39) 

40from litellm.types.guardrails import GuardrailEventHooks 

41from litellm.types.utils import ( 

42 CallTypes, 

43 CallTypesLiteral, 

44 ChatCompletionDeltaCustomToolCall, 

45 ChatCompletionDeltaToolCall, 

46 ChatCompletionMessageCustomToolCall, 

47 ChatCompletionMessageToolCall, 

48 ChatCompletionToolCallChunk, 

49 Choices, 

50 GenericGuardrailAPIInputs, 

51 ModelResponse, 

52 ModelResponseStream, 

53) 

54 

55ToolCallLike: TypeAlias = ( 

56 ChatCompletionMessageToolCall 

57 | ChatCompletionDeltaToolCall 

58 | ChatCompletionMessageCustomToolCall 

59 | ChatCompletionDeltaCustomToolCall 

60 | ChatCompletionToolCallChunk 

61) 

62 

63 

64class _ToolCallFunctionSlice(BaseModel): 

65 model_config = ConfigDict(from_attributes=True, extra="ignore") 

66 

67 name: str | None = None 

68 arguments: str | None = None 

69 

70 @field_validator("name", "arguments", mode="before") 

71 @classmethod 

72 def _coerce_to_scannable_text(cls, value: object) -> str | None: 

73 """Accept any shape a client can put here and render it scannable. 

74 

75 The OpenAI request path forwards client-supplied ``tool_calls`` verbatim, so a 

76 client can post a dict for ``arguments`` or a non-string for ``name``. Rejecting 

77 either would fail validation for the whole slice, which reads as an unscannable 

78 tool call and skips it silently -- the one outcome a scanner must never have. 

79 A caller could otherwise suppress the scan on a tool call just by sending 

80 ``"name": 123``. 

81 """ 

82 if value is None or isinstance(value, str): 

83 return value 

84 return json.dumps(value) if isinstance(value, (dict, list)) else str(value) 

85 

86 

87class _ToolCallSlice(BaseModel): 

88 model_config = ConfigDict(from_attributes=True, extra="ignore") 

89 

90 function: _ToolCallFunctionSlice | None = None 

91 

92 

93if TYPE_CHECKING: 93 ↛ 94line 93 didn't jump to line 94 because the condition on line 93 was never true

94 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

95 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel 

96 

97 

98class PanwPrismaAirsHandler(CustomGuardrail): 

99 """ 

100 LiteLLM Built-in Guardrail for Palo Alto Networks Prisma AI Runtime Security (AIRS). 

101 

102 Scans prompts and responses using PANW Prisma AIRS API to detect malicious content, 

103 injection attempts, and policy violations. Supports content masking and fail-closed error handling. 

104 

105 Configuration: 

106 guardrail_name: Name of the guardrail instance 

107 api_key: PANW Prisma AIRS API key 

108 api_base: PANW Prisma AIRS API endpoint (default: https://service.api.aisecurity.paloaltonetworks.com) 

109 profile_name: PANW security profile name (optional if API key has linked profile) 

110 app_name: Application name for tracking in Prisma AIRS analytics (default: "LiteLLM") 

111 mask_request_content: Apply masking to prompts (default: False) 

112 mask_response_content: Apply masking to responses (default: False) 

113 mask_on_block: Backwards compatible flag that enables both request and response masking 

114 """ 

115 

116 _PROVIDER_NAME = "panw_prisma_airs" 

117 

118 #: AIRS fields withheld from the client-visible error detail. 

119 #: ``response_masked_data`` is the model's own generation. The block branch that builds 

120 #: this detail is only reached when ``mask_response_content`` is False, so echoing it 

121 #: back would hand the caller exactly the text the operator declined to deliver. 

122 #: ``prompt_masked_data`` is deliberately NOT withheld: it is the caller's own input, 

123 #: and it is one of the fields the ticket asks for. 

124 _CLIENT_HIDDEN_SCAN_FIELDS: Final = frozenset({"response_masked_data"}) 

125 

126 def __init__( 

127 self, 

128 guardrail_name: str, 

129 profile_name: str | None = None, 

130 api_key: str | None = None, 

131 api_base: str | None = None, 

132 default_on: bool = True, 

133 mask_on_block: bool = False, 

134 mask_request_content: bool = False, 

135 mask_response_content: bool = False, 

136 app_name: str | None = None, 

137 fallback_on_error: Literal["block", "allow"] = "block", 

138 timeout: float = 10.0, 

139 violation_message_template: str | None = None, 

140 http_client: AsyncHTTPHandler | None = None, 

141 **kwargs, 

142 ): 

143 """Initialize PANW Prisma AIRS guardrail handler.""" 

144 

145 # Masking configuration - mask_on_block enables both for backwards compatibility 

146 self.mask_on_block = mask_on_block 

147 _mask_request_content: Final = mask_request_content or mask_on_block 

148 _mask_response_content: Final = mask_response_content or mask_on_block 

149 

150 # Initialize parent CustomGuardrail with masking flags 

151 super().__init__( 

152 guardrail_name=guardrail_name, 

153 default_on=default_on, 

154 supported_event_hooks=list(self.get_supported_event_hooks()), 

155 mask_request_content=_mask_request_content, 

156 mask_response_content=_mask_response_content, 

157 violation_message_template=violation_message_template, 

158 **kwargs, 

159 ) 

160 

161 # Store configuration with env var fallbacks 

162 self.api_key = api_key or os.getenv("PANW_PRISMA_AIRS_API_KEY") 

163 self.api_base = ( 

164 api_base or os.getenv("PANW_PRISMA_AIRS_API_BASE") or "https://service.api.aisecurity.paloaltonetworks.com" 

165 ) 

166 self.profile_name = profile_name 

167 

168 # Handle app_name: Default to "LiteLLM", or prefix user's app_name with "LiteLLM-" 

169 if app_name: 

170 self.app_name = f"LiteLLM-{app_name}" 

171 else: 

172 self.app_name = "LiteLLM" 

173 

174 # Validate required configuration 

175 if not self.api_key: 

176 raise ValueError( 

177 "PANW Prisma AIRS: api_key is required. " 

178 "Set it via config or PANW_PRISMA_AIRS_API_KEY environment variable." 

179 ) 

180 

181 # Warn if no profile is configured (user must have API key with linked profile) 

182 if not self.profile_name: 

183 verbose_proxy_logger.warning( 

184 "PANW Prisma AIRS Guardrail '%s': No profile_name configured. Ensure your API key has a linked profile in Strata Cloud Manager, or provide 'profile_name'/'profile_id' via config or per-request metadata. Requests will fail if the API key is not linked to a profile.", 

185 guardrail_name, 

186 ) 

187 

188 self.http_client = http_client 

189 self.fallback_on_error = fallback_on_error 

190 # Coerce defensively. The dashboard UI persists this field as a JSON 

191 # string, and Pydantic extras (the path that splats model_dump into 

192 # this handler) preserve whatever type the user supplied. A string 

193 # value would otherwise reach httpx, which raises TypeError on its 

194 # internal '<=' comparison and surfaces as a misleading api_error. 

195 self.timeout = float(timeout) if timeout is not None else 10.0 

196 

197 # Tri-state: None = not set (default-on for Anthropic), True = explicit on, False = explicit off 

198 self.experimental_use_latest_role_message_only: bool | None = kwargs.get( 

199 "experimental_use_latest_role_message_only" 

200 ) 

201 

202 if self.fallback_on_error == "allow": 

203 verbose_proxy_logger.warning( 

204 "PANW Prisma AIRS Guardrail '%s': fallback_on_error='allow' - requests will proceed without scanning when API is unavailable.", 

205 guardrail_name, 

206 ) 

207 

208 verbose_proxy_logger.info( 

209 "Initialized PANW Prisma AIRS Guardrail: %s (profile=%s, mask_request=%s, mask_response=%s, fallback_on_error=%s, timeout=%s)", 

210 guardrail_name, 

211 self.profile_name or "API-key-linked", 

212 self.mask_request_content, 

213 self.mask_response_content, 

214 self.fallback_on_error, 

215 self.timeout, 

216 ) 

217 

218 # MCP event → base-call compatibility map. 

219 # Allows guardrails configured with mode: pre_call / during_call to 

220 # automatically run on MCP tool invocations (pre_mcp_call / during_mcp_call). 

221 _MCP_COMPAT_MAP = { 

222 GuardrailEventHooks.pre_mcp_call: GuardrailEventHooks.pre_call, 

223 GuardrailEventHooks.during_mcp_call: GuardrailEventHooks.during_call, 

224 } 

225 

226 def should_run_guardrail(self, data: Mapping[str, object], event_type: GuardrailEventHooks) -> bool: 

227 if super().should_run_guardrail(data, event_type): 

228 return True 

229 compat: Final = self._MCP_COMPAT_MAP.get(event_type) 

230 if compat is not None: 

231 if super().should_run_guardrail(data, compat): 

232 return True 

233 return False 

234 

235 def _extract_text_from_messages(self, messages: Sequence[Mapping[str, object]]) -> str: 

236 """Extract text content from messages array.""" 

237 if not isinstance(messages, list) or not messages: 

238 return "" 

239 

240 # Find the last user message 

241 for message in reversed(messages): 

242 if message.get("role") not in ("user", "developer"): 

243 continue 

244 

245 content = message.get("content") 

246 if not content: 

247 continue 

248 

249 if isinstance(content, str): 

250 return content 

251 

252 if isinstance(content, list): 

253 return self._extract_text_from_content_list(content) 

254 

255 return "" 

256 

257 def _extract_text_from_content_list(self, content_list: list[dict[str, Any]]) -> str: 

258 """Extract text from content list format.""" 

259 text_parts: Final = [ 

260 part.get("text", "") 

261 for part in content_list 

262 if isinstance(part, dict) and part.get("type") == "text" and part.get("text") 

263 ] 

264 return " ".join(text_parts) if text_parts else "" 

265 

266 def _extract_response_text(self, response: ModelResponse) -> str: 

267 """ 

268 Extract all text content from LLM response. 

269 Handles multiple choices, tool calls, and function calls. 

270 Returns concatenated text for scanning. 

271 """ 

272 try: 

273 text_parts: Final = [] 

274 

275 if hasattr(response, "choices") and response.choices: 

276 for choice in response.choices: 

277 if isinstance(choice, Choices): 

278 # Extract message content 

279 if choice.message.content: 

280 text_parts.append(str(choice.message.content)) 

281 

282 # Extract tool call arguments 

283 if hasattr(choice.message, "tool_calls") and choice.message.tool_calls: 

284 for tool_call in choice.message.tool_calls: 

285 if hasattr(tool_call, "function") and hasattr(tool_call.function, "arguments"): 

286 text_parts.append(str(tool_call.function.arguments)) 

287 

288 # Extract function call arguments (legacy) 

289 if hasattr(choice.message, "function_call") and choice.message.function_call: 

290 if hasattr(choice.message.function_call, "arguments"): 

291 text_parts.append(str(choice.message.function_call.arguments)) 

292 

293 return " ".join(text_parts) if text_parts else "" 

294 except (AttributeError, IndexError) as e: 

295 verbose_proxy_logger.error("PANW Prisma AIRS: Error extracting response text: %s", e) 

296 return "" 

297 

298 async def _call_panw_api( 

299 self, 

300 content: str = "", 

301 is_response: bool = False, 

302 metadata: Mapping[str, object] | None = None, 

303 call_id: object = None, 

304 tool_event: dict[str, Any] | None = None, 

305 ) -> dict[str, object]: 

306 """Call PANW Prisma AIRS API to scan content or a tool_event.""" 

307 

308 if tool_event is None and not content.strip(): 

309 return {"action": "allow", "category": "empty"} 

310 

311 # tr_id is optional in the AIRS API. Allow call_id=None only for 

312 # MCP tool_events (ecosystem == "mcp"). All other paths (content 

313 # scans, non-MCP tool_events) remain fail-closed. 

314 if not call_id: 

315 _is_mcp_tool_event: Final = ( 

316 tool_event is not None 

317 and isinstance(tool_event.get("metadata"), dict) 

318 and tool_event["metadata"].get("ecosystem") == "mcp" 

319 ) 

320 if not _is_mcp_tool_event: 

321 return { 

322 "action": "block", 

323 "category": "missing_call_id", 

324 "_always_block": True, 

325 } 

326 

327 # Build Prisma AIRS API metadata 

328 # Handle app_name: LiteLLM by default, or LiteLLM-{user_app_name} if user provides one 

329 user_app_name: Final = metadata.get("app_name") if metadata else None 

330 if user_app_name: 

331 app_name_value = f"LiteLLM-{user_app_name}" 

332 else: 

333 app_name_value = self.app_name # Defaults to "LiteLLM" 

334 

335 panw_metadata: Final[dict[str, object]] = { 

336 "app_user": ( 

337 (metadata.get("app_user") or metadata.get("user") or "litellm_user") if metadata else "litellm_user" 

338 ), 

339 "ai_model": metadata.get("model", "unknown") if metadata else "unknown", 

340 "app_name": app_name_value, 

341 "source": "litellm_builtin_guardrail", 

342 } 

343 

344 # Include user_ip if available (from LiteLLM metadata or request) 

345 if metadata and metadata.get("user_ip"): 

346 panw_metadata["user_ip"] = metadata["user_ip"] 

347 elif metadata and metadata.get("requester_ip_address"): 

348 panw_metadata["user_ip"] = metadata["requester_ip_address"] 

349 

350 # Forward litellm_trace_id in AIRS metadata for session correlation 

351 if metadata and metadata.get("litellm_trace_id"): 

352 panw_metadata["litellm_trace_id"] = metadata["litellm_trace_id"] 

353 

354 # Build contents: tool_event takes priority, else prompt/response text 

355 contents: Sequence[Mapping[str, object]] 

356 if tool_event is not None: 

357 contents = [{"tool_event": tool_event}] 

358 else: 

359 contents = [{"response" if is_response else "prompt": content}] 

360 

361 payload: Final[dict[str, object]] = { 

362 "metadata": panw_metadata, 

363 "contents": contents, 

364 } 

365 # Use per-request litellm_call_id as AIRS tr_id; keep litellm_trace_id in metadata. 

366 if call_id: 

367 payload["tr_id"] = call_id 

368 

369 # Build ai_profile object per PANW API schema 

370 # Priority: per-request profile_id > per-request profile_name > config profile_name 

371 # Note: If both are provided, PANW API uses profile_id (profile_id takes precedence) 

372 profile_name = None 

373 profile_id = None 

374 

375 if metadata: 

376 profile_id = metadata.get("profile_id") 

377 profile_name = metadata.get("profile_name", self.profile_name) 

378 else: 

379 profile_name = self.profile_name 

380 

381 # Add ai_profile to payload if profile is specified 

382 # If neither profile_name nor profile_id is provided, PANW API will use the 

383 # profile linked to the API key (if configured in Strata Cloud Manager) 

384 if profile_name or profile_id: 

385 ai_profile: Final[dict[str, object]] = {} 

386 if profile_id: 

387 ai_profile["profile_id"] = profile_id 

388 if profile_name: 

389 ai_profile["profile_name"] = profile_name 

390 payload["ai_profile"] = ai_profile 

391 

392 if is_response and tool_event is None: 

393 panw_metadata["is_response"] = True 

394 

395 headers: Final = { 

396 "Content-Type": "application/json", 

397 "Accept": "application/json", 

398 "x-pan-token": self.api_key or "", # api_key validated in __init__, never None 

399 } 

400 

401 try: 

402 # Use LiteLLM's async HTTP client 

403 async_client: Final = self.http_client or get_async_httpx_client( 

404 llm_provider=httpxSpecialProvider.GuardrailCallback 

405 ) 

406 

407 # Bypass wrapper to access follow_redirects parameter 

408 response: Final = await async_client.client.post( 

409 f"{self.api_base}/v1/scan/sync/request", 

410 headers=headers, 

411 json=payload, 

412 timeout=self.timeout, 

413 follow_redirects=False, # Prevent redirect attacks 

414 ) 

415 response.raise_for_status() 

416 

417 result: Final[dict[str, object]] = response.json() 

418 

419 # Validate response format 

420 if "action" not in result: 

421 verbose_proxy_logger.error("PANW Prisma AIRS: Invalid API response format: %s", result) 

422 return {"action": "block", "category": "api_error"} 

423 

424 # Check for profile-related errors from PANW API 

425 if result.get("action") == "block" and "error" in result: 

426 error_msg: Final = str(result.get("error", "")).lower() 

427 if "profile" in error_msg and ( 

428 "not found" in error_msg or "required" in error_msg or "invalid" in error_msg 

429 ): 

430 verbose_proxy_logger.error( 

431 "PANW Prisma AIRS: Profile configuration error. Ensure your API key has a linked profile in Strata Cloud Manager, or provide 'profile_name' or 'profile_id' in config/metadata. PANW API response: %s", 

432 result, 

433 ) 

434 

435 verbose_proxy_logger.debug( 

436 "PANW Prisma AIRS: Scan result - Action: %s, Category: %s", 

437 result.get("action"), 

438 result.get("category", "unknown"), 

439 ) 

440 return result 

441 

442 except httpx.HTTPStatusError as e: 

443 status: Final = e.response.status_code 

444 error_body = "" 

445 try: 

446 error_body = e.response.text 

447 except Exception: 

448 pass 

449 

450 # Enhanced 400 diagnostics for tool_event schema debugging 

451 if status == 400: 

452 diag_parts: Final = ["PANW Prisma AIRS: HTTP 400 from AIRS API."] 

453 if tool_event is not None: 

454 diag_parts.append(f"tool_event.metadata={tool_event.get('metadata')}") 

455 has_input: Final = "input" in tool_event 

456 input_len: Final = len(tool_event["input"]) if has_input else 0 

457 diag_parts.append(f"input present={has_input}, len={input_len}") 

458 diag_parts.append(f"response body: {error_body[:500]}") 

459 verbose_proxy_logger.error(" | ".join(diag_parts)) 

460 

461 is_profile_error: Final = any( 

462 phrase in error_body.lower() 

463 for phrase in [ 

464 "profile not found", 

465 "profile required", 

466 "invalid profile", 

467 ] 

468 ) 

469 

470 if status in (401, 403) or is_profile_error: 

471 verbose_proxy_logger.error( 

472 "PANW Prisma AIRS: Authentication/config error (HTTP %s). Check API key and profile configuration.", 

473 status, 

474 ) 

475 return { 

476 "action": "block", 

477 "category": "config_error", 

478 "_always_block": True, 

479 } 

480 elif status == 429 or status >= 500: 

481 # Transient: rate-limit and server errors — safe to fail-open 

482 verbose_proxy_logger.error("PANW Prisma AIRS: API error (HTTP %s): %s", status, error_body[:500]) 

483 return { 

484 "action": "block", 

485 "category": f"http_{status}_error", 

486 "_is_transient": True, 

487 } 

488 else: 

489 # Permanent 4xx client errors (400, 404, etc.) — must not bypass scanning 

490 if status != 400: # 400 already logged with diagnostics above 

491 verbose_proxy_logger.error("PANW Prisma AIRS: API error (HTTP %s): %s", status, error_body[:500]) 

492 return { 

493 "action": "block", 

494 "category": f"http_{status}_error", 

495 "_always_block": True, 

496 } 

497 

498 except httpx.TimeoutException as e: 

499 verbose_proxy_logger.error("PANW Prisma AIRS: Timeout error: %s", e) 

500 return { 

501 "action": "block", 

502 "category": "timeout_error", 

503 "_is_transient": True, 

504 } 

505 

506 except httpx.RequestError as e: 

507 verbose_proxy_logger.error("PANW Prisma AIRS: Network/request error: %s", e) 

508 return { 

509 "action": "block", 

510 "category": "network_error", 

511 "_is_transient": True, 

512 } 

513 

514 except Exception as e: 

515 verbose_proxy_logger.error("PANW Prisma AIRS: Unexpected error: %s", e) 

516 return {"action": "block", "category": "api_error", "_is_transient": True} 

517 

518 @staticmethod 

519 def _get_mcp_server_name(request_data: dict, mcp_tool_name: str) -> str: 

520 """Resolve MCP server name from request data or MCP registry.""" 

521 if request_data.get("mcp_server_name"): 

522 return request_data["mcp_server_name"] 

523 if request_data.get("server_name"): 

524 return request_data["server_name"] 

525 try: 

526 from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( 

527 global_mcp_server_manager, 

528 ) 

529 

530 server_id: Final = request_data.get("server_id") 

531 if server_id: 

532 server: Final = global_mcp_server_manager.get_mcp_server_by_id(server_id) 

533 if server: 

534 return ( 

535 getattr(server, "alias", None) 

536 or getattr(server, "server_name", None) 

537 or getattr(server, "name", None) 

538 or getattr(server, "server_id", None) 

539 or "unknown" 

540 ) 

541 return global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.get(mcp_tool_name, "unknown") 

542 except ImportError: 

543 return "unknown" 

544 except Exception: 

545 verbose_proxy_logger.debug( 

546 "PANW Prisma AIRS: unexpected error resolving MCP server name", 

547 exc_info=True, 

548 ) 

549 return "unknown" 

550 

551 def _get_masked_text(self, scan_result: Mapping[str, object], is_response: bool = False) -> str | None: 

552 """Extract masked text from PANW scan result.""" 

553 masked_key: Final = "response_masked_data" if is_response else "prompt_masked_data" 

554 masked_data: Final = scan_result.get(masked_key) 

555 if masked_data and isinstance(masked_data, dict): 

556 return masked_data.get("data") 

557 return None 

558 

559 @staticmethod 

560 def _mask_content_list(content_list: list, masked_text: str) -> list: 

561 """Replace text parts in a content list, preserving non-text parts (images, etc.).""" 

562 new_content: Final = [] 

563 for part in content_list: 

564 if isinstance(part, dict) and part.get("type") == "text": 

565 new_content.append({"type": "text", "text": masked_text}) 

566 else: 

567 new_content.append(part) 

568 return new_content 

569 

570 @staticmethod 

571 def _apply_mcp_masking( 

572 request_data: dict, 

573 original_args: object, 

574 masked_text: str, 

575 *, 

576 is_blocked: bool = True, 

577 ) -> None: 

578 """Write masked arguments back to MCP request_data fields. 

579 

580 - ``arguments`` is the authoritative field that ``call_mcp_tool`` 

581 reads, so it must be updated first. 

582 - ``mcp_arguments`` is mirrored for consistency / test observability. 

583 - If the original args were structured (dict/list), attempt 

584 ``json.loads`` to preserve the type; block if the masked text 

585 is not valid JSON (to avoid corrupting structured args). 

586 - If neither ``arguments`` nor ``mcp_arguments`` is present in 

587 request_data, block — do not silently invent a new field. 

588 """ 

589 has_arguments: Final = "arguments" in request_data 

590 has_mcp_arguments: Final = "mcp_arguments" in request_data 

591 if not has_arguments and not has_mcp_arguments: 

592 raise HTTPException( 

593 status_code=400, 

594 detail={ 

595 "error": { 

596 "message": "MCP request blocked: no rewritable argument field present", 

597 "type": "guardrail_violation", 

598 "code": "panw_prisma_airs_blocked", 

599 } 

600 }, 

601 ) 

602 

603 # If the original args were structured, preserve the type. 

604 if isinstance(original_args, (dict, list)): 

605 try: 

606 parsed: Final[object] = json.loads(masked_text) 

607 except (json.JSONDecodeError, TypeError): 

608 raise HTTPException( 

609 status_code=400, 

610 detail={ 

611 "error": { 

612 "message": "MCP request blocked: masked data is not valid JSON for structured arguments", 

613 "type": "guardrail_violation", 

614 "code": "panw_prisma_airs_blocked", 

615 } 

616 }, 

617 ) 

618 masked_value: object = parsed 

619 else: 

620 masked_value = masked_text 

621 

622 if has_arguments: 

623 request_data["arguments"] = masked_value 

624 if has_mcp_arguments: 

625 request_data["mcp_arguments"] = masked_value 

626 

627 if is_blocked: 

628 verbose_proxy_logger.warning( 

629 "PANW Prisma AIRS: MCP request blocked but masked instead (mask_request_content=True)" 

630 ) 

631 else: 

632 verbose_proxy_logger.info("PANW Prisma AIRS: MCP request allowed with PII masking applied") 

633 

634 def _apply_masking_to_messages( 

635 self, messages: list[dict[str, object]], masked_text: str 

636 ) -> Sequence[Mapping[str, object]]: 

637 """Apply masked text to the last user message.""" 

638 if not messages: 

639 return messages 

640 

641 for i, message in enumerate(reversed(messages)): 

642 if message.get("role") == "user": 

643 new_message = message.copy() 

644 content = message.get("content") 

645 

646 if isinstance(content, str): 

647 new_message["content"] = masked_text 

648 elif isinstance(content, list): 

649 new_message["content"] = self._mask_content_list(content, masked_text) 

650 

651 idx = len(messages) - i - 1 

652 return messages[:idx] + [new_message] + messages[idx + 1 :] 

653 

654 return messages 

655 

656 def _apply_masking_to_response(self, response: ModelResponse, masked_text: str) -> None: 

657 """ 

658 Apply masked text to all content in response in-place. 

659 Handles message content, tool calls, and function calls across all choices. 

660 Preserves list-based content structure (e.g., multimodal messages). 

661 """ 

662 if not hasattr(response, "choices") or not response.choices: 

663 return 

664 

665 for choice in response.choices: 

666 if isinstance(choice, Choices): 

667 # Mask message content - handle both string and list formats 

668 content = choice.message.content 

669 if content: 

670 if isinstance(content, str): 

671 choice.message.content = masked_text 

672 elif isinstance(content, list): 

673 choice.message.content = self._mask_content_list(content, masked_text) 

674 

675 # Mask tool call arguments 

676 if hasattr(choice.message, "tool_calls") and choice.message.tool_calls: 

677 for tool_call in choice.message.tool_calls: 

678 if hasattr(tool_call, "function") and hasattr(tool_call.function, "arguments"): 

679 tool_call.function.arguments = masked_text 

680 

681 # Mask function call arguments (legacy) 

682 if hasattr(choice.message, "function_call") and choice.message.function_call: 

683 if hasattr(choice.message.function_call, "arguments"): 

684 choice.message.function_call.arguments = masked_text 

685 

686 def _build_error_detail( 

687 self, 

688 scan_result: Mapping[str, object], 

689 is_response: bool = False, 

690 ) -> Mapping[str, Mapping[str, object]]: 

691 """Build enhanced error detail with scan information.""" 

692 action_type: Final = "Response" if is_response else "Prompt" 

693 code_suffix: Final = "_response_blocked" if is_response else "_blocked" 

694 

695 category: Final = scan_result.get("category", "unknown") 

696 default_msg: Final = f"{action_type} blocked by PANW Prisma AI Security policy (Category: {category})" 

697 

698 # Use custom violation message template if configured 

699 error_msg: Final = self.render_violation_message( 

700 default=default_msg, 

701 context={ 

702 "guardrail_name": self.guardrail_name, 

703 "category": category, 

704 "action_type": action_type, 

705 "default_message": default_msg, 

706 }, 

707 ) 

708 

709 return { 

710 "error": { 

711 **{ 

712 key: value 

713 for key, value in scan_result.items() 

714 if not key.startswith("_") and key not in self._CLIENT_HIDDEN_SCAN_FIELDS 

715 }, 

716 "message": error_msg, 

717 "type": "guardrail_violation", 

718 "code": f"panw_prisma_airs{code_suffix}", 

719 "guardrail": self.guardrail_name, 

720 "category": category, 

721 } 

722 } 

723 

724 def _record_scan_id( 

725 self, request_data: dict[str, object], scan_result: Mapping[str, object], stage: GuardrailEventHooks 

726 ) -> None: 

727 """Surface the AIRS scan id on the response, so allowed calls are auditable too.""" 

728 scan_id: Final = scan_result.get("scan_id") 

729 add_guardrail_scan_id( 

730 request_data=request_data, 

731 scan_id=str(scan_id) if scan_id else None, 

732 guardrail_name=self.guardrail_name, 

733 provider=self._PROVIDER_NAME, 

734 stage=stage, 

735 ) 

736 

737 def _handle_api_error_with_logging( 

738 self, 

739 scan_result: dict[str, object], 

740 data: dict[str, object], 

741 start_time: datetime, 

742 event_type: GuardrailEventHooks, 

743 is_response: bool = False, 

744 ) -> None: 

745 """Handle API errors with fail-open/fail-closed logic.""" 

746 end_time: Final = datetime.now() 

747 duration: Final = (end_time - start_time).total_seconds() 

748 category: Final = scan_result.get("category", "api_error") 

749 

750 self.add_standard_logging_guardrail_information_to_request_data( 

751 guardrail_provider=self._PROVIDER_NAME, 

752 guardrail_json_response=scan_result, 

753 request_data=data, 

754 guardrail_status="guardrail_failed_to_respond", 

755 start_time=start_time.timestamp(), 

756 end_time=end_time.timestamp(), 

757 duration=duration, 

758 event_type=event_type, 

759 ) 

760 

761 if scan_result.get("_always_block"): 

762 is_config: Final = category == "config_error" 

763 raise HTTPException( 

764 status_code=500, 

765 detail={ 

766 "error": { 

767 "message": ( 

768 "Security scan failed - configuration error" 

769 if is_config 

770 else "Security scan failed - request blocked for safety" 

771 ), 

772 "type": ("guardrail_config_error" if is_config else "guardrail_scan_error"), 

773 "code": ("panw_prisma_airs_config_error" if is_config else "panw_prisma_airs_scan_failed"), 

774 "guardrail": self.guardrail_name, 

775 "category": category, 

776 } 

777 }, 

778 ) 

779 

780 if scan_result.get("_is_transient") and self.fallback_on_error == "allow": 

781 verbose_proxy_logger.warning( 

782 "PANW Prisma AIRS: Allowing %s without scanning (fallback_on_error='allow', error: %s)", 

783 "response" if is_response else "request", 

784 category, 

785 ) 

786 add_guardrail_to_applied_guardrails_header( 

787 request_data=data, guardrail_name=f"{self.guardrail_name}:unscanned" 

788 ) 

789 return 

790 

791 raise HTTPException( 

792 status_code=500, 

793 detail={ 

794 "error": { 

795 "message": "Security scan failed - request blocked for safety", 

796 "type": "guardrail_scan_error", 

797 "code": "panw_prisma_airs_scan_failed", 

798 "guardrail": self.guardrail_name, 

799 "category": category, 

800 } 

801 }, 

802 ) 

803 

804 def _prepare_metadata_from_request(self, data: dict[str, Any]) -> dict[str, object]: 

805 """ 

806 Extract and prepare metadata from request data for PANW API call. 

807 

808 Supported metadata fields (from request.metadata): 

809 - profile_name: AI security profile name (PANW API field) 

810 - profile_id: AI security profile ID (PANW API field, takes precedence) 

811 - user_ip: User IP address for tracking 

812 - app_name: Application identifier (will be prefixed with "LiteLLM-") 

813 

814 Note: If neither profile_name nor profile_id is provided, PANW API will use 

815 the profile linked to the API key (configured in Strata Cloud Manager). 

816 If both are provided, PANW API uses profile_id (profile_id takes precedence). 

817 """ 

818 user_metadata: Final = data.get("metadata", {}) or {} 

819 requester_meta: Final = user_metadata.get("requester_metadata", {}) or {} 

820 metadata: Final[dict[str, object]] = { 

821 "user": data.get("user") or "litellm_user", 

822 "model": data.get("model") or "unknown", 

823 } 

824 

825 # Pass through PANW API fields (check requester_metadata fallback for /v1/messages routes) 

826 for key in ("profile_name", "profile_id", "user_ip", "app_name", "app_user"): 

827 val = user_metadata.get(key) or requester_meta.get(key) 

828 if val: 

829 metadata[key] = val 

830 

831 # Include litellm_trace_id for session tracking. 

832 # Sources (checked in priority order): 

833 # 1. data["litellm_trace_id"] — top-level body field 

834 # 2. metadata["litellm_trace_id"] — user passes in request metadata 

835 # 3. metadata["trace_id"] — x-litellm-trace-id header 

836 # (litellm_pre_call_utils stores it as "trace_id", not "litellm_trace_id") 

837 # 4. requester_metadata["litellm_trace_id"] — deep copy for /v1/messages routes 

838 trace_id: Final = ( 

839 data.get("litellm_trace_id") 

840 or user_metadata.get("litellm_trace_id") 

841 or user_metadata.get("trace_id") 

842 or requester_meta.get("litellm_trace_id") 

843 ) 

844 if trace_id: 

845 metadata["litellm_trace_id"] = trace_id 

846 

847 return metadata 

848 

849 @staticmethod 

850 def _extract_text_from_sse_bytes(chunks: Sequence[bytes]) -> str: 

851 """Extract text from Anthropic SSE byte chunks (content_block_delta → text_delta).""" 

852 texts: Final[list[str]] = [] 

853 raw: Final = b"".join(chunks).decode("utf-8", errors="replace") 

854 for line in raw.split("\n"): 

855 line = line.strip() 

856 if not line.startswith("data: "): 

857 continue 

858 try: 

859 data = json.loads(line[6:]) 

860 except (json.JSONDecodeError, ValueError): 

861 continue 

862 if not isinstance(data, dict): 

863 continue 

864 if data.get("type") == "content_block_delta": 

865 delta = data.get("delta") or {} 

866 if delta.get("type") == "text_delta": 

867 texts.append(delta.get("text", "")) 

868 return "".join(texts) 

869 

870 @staticmethod 

871 def _extract_text_from_streaming_events(chunks: Sequence[object]) -> str: 

872 """Extract text from /v1/responses streaming events (object or dict).""" 

873 

874 def _attr(c, key): 

875 val = getattr(c, key, None) 

876 if val is None and isinstance(c, dict): 

877 val = c.get(key) 

878 return val 

879 

880 parts: Final[list[str]] = [] 

881 for chunk in chunks: 

882 if _attr(chunk, "type") == "response.output_text.delta": 

883 delta = _attr(chunk, "delta") 

884 if isinstance(delta, str): 

885 parts.append(delta) 

886 # Defense-in-depth: handle dict chat.completion.chunk format 

887 elif isinstance(chunk, dict) and chunk.get("object") == "chat.completion.chunk": 

888 for choice in chunk.get("choices") or []: 

889 if isinstance(choice, dict): 

890 delta = choice.get("delta") or {} 

891 content = delta.get("content") 

892 if isinstance(content, str): 

893 parts.append(content) 

894 # Fallback: response.output_text.done carries full text if no deltas captured 

895 if not parts: 

896 for chunk in chunks: 

897 if _attr(chunk, "type") == "response.output_text.done": 

898 text = _attr(chunk, "text") 

899 if isinstance(text, str): 

900 parts.append(text) 

901 return "".join(parts) 

902 

903 async def _scan_raw_streaming_text(self, text: str, request_data: dict, start_time: datetime) -> None: 

904 """Scan text from non-ModelResponse streaming chunks. Raises HTTPException(400) on block. 

905 

906 Note: response masking is not supported on raw streaming paths 

907 (/v1/messages, /v1/responses) because the response is raw SSE 

908 bytes/events that cannot be reliably reconstructed. If 

909 mask_response_content is configured, a warning is logged and the 

910 response is blocked instead. Request-side masking 

911 (mask_request_content) is unaffected — it runs in async_pre_call_hook 

912 before streaming begins. 

913 """ 

914 if not text or not text.strip(): 

915 return 

916 

917 metadata: Final = self._prepare_metadata_from_request(request_data) 

918 scan_result: Final = await self._call_panw_api( 

919 content=text, 

920 is_response=True, 

921 metadata=metadata, 

922 call_id=request_data.get("litellm_call_id"), 

923 ) 

924 if scan_result.get("_is_transient") or scan_result.get("_always_block"): 

925 self._handle_api_error_with_logging( 

926 scan_result, 

927 request_data, 

928 start_time, 

929 is_response=True, 

930 event_type=GuardrailEventHooks.post_call, 

931 ) 

932 return # _always_block raises inside; transient errors fail-open here 

933 action: Final = scan_result.get("action", "block") 

934 if action != "allow": 

935 masked_text: Final = self._get_masked_text(scan_result, is_response=True) 

936 if masked_text and self.mask_response_content: 

937 verbose_proxy_logger.warning( 

938 "PANW Prisma AIRS: mask_response_content is configured but " 

939 "cannot be applied to raw streaming responses (/v1/messages " 

940 "or /v1/responses). Blocking response instead." 

941 ) 

942 raise HTTPException( 

943 status_code=400, 

944 detail=self._build_error_detail(scan_result, is_response=True), 

945 ) 

946 # Success logging + observability header 

947 end_time: Final = datetime.now() 

948 self.add_standard_logging_guardrail_information_to_request_data( 

949 guardrail_provider=self._PROVIDER_NAME, 

950 guardrail_json_response=scan_result, 

951 request_data=request_data, 

952 guardrail_status="success", 

953 start_time=start_time.timestamp(), 

954 end_time=end_time.timestamp(), 

955 duration=(end_time - start_time).total_seconds(), 

956 event_type=GuardrailEventHooks.post_call, 

957 ) 

958 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) 

959 self._record_scan_id(request_data, scan_result, GuardrailEventHooks.post_call) 

960 

961 def _check_and_mark_scanned(self, data: dict, scan_type: str) -> bool: 

962 """ 

963 Check if request has already been scanned and mark it as scanned. 

964 

965 Args: 

966 data: Request data dictionary 

967 scan_type: Type of scan ('pre', 'post', 'streaming') 

968 

969 Returns: 

970 True if already scanned (should skip), False if needs scanning 

971 """ 

972 call_id = data.get("litellm_call_id") 

973 if not call_id: 

974 call_id = str(uuid.uuid4()) 

975 data["litellm_call_id"] = call_id 

976 verbose_proxy_logger.warning( 

977 "PANW Prisma AIRS: litellm_call_id missing from request data, synthesized %s for %s scan deduplication", 

978 call_id, 

979 scan_type, 

980 ) 

981 

982 scan_key: Final = f"_panw_{scan_type}_scanned_{call_id}" 

983 litellm_metadata: Final = data.setdefault("litellm_metadata", {}) 

984 

985 if litellm_metadata.get(scan_key): 

986 verbose_proxy_logger.debug("PANW Prisma AIRS: Skipping duplicate %s-call scan", scan_type) 

987 return True # Already scanned 

988 

989 litellm_metadata[scan_key] = True 

990 return False # Needs scanning 

991 

992 def _extract_prompt_from_request(self, data: dict) -> str: 

993 """ 

994 Extract prompt text from request data. 

995 

996 Handles both chat completion (messages) and text completion (prompt) formats. 

997 

998 Args: 

999 data: Request data dictionary 

1000 

1001 Returns: 

1002 Extracted prompt text, or empty string if not found 

1003 """ 

1004 # Extract from messages (chat completion) 

1005 messages: Final = data.get("messages", []) 

1006 prompt_text = self._extract_text_from_messages(messages) 

1007 

1008 # Fallback to prompt field for text completion requests 

1009 if not prompt_text: 

1010 prompt_value: Final = data.get("prompt") 

1011 if isinstance(prompt_value, str): 

1012 prompt_text = prompt_value 

1013 elif isinstance(prompt_value, list): 

1014 # Handle list of prompts (batch text completion) 

1015 prompt_text = " ".join(str(p) for p in prompt_value if p) 

1016 else: 

1017 prompt_text = "" 

1018 

1019 return prompt_text 

1020 

1021 @log_guardrail_information 

1022 async def async_pre_call_hook( 

1023 self, 

1024 user_api_key_dict: UserAPIKeyAuth, 

1025 cache: DualCache, 

1026 data: dict[str, Any], 

1027 call_type: CallTypesLiteral, 

1028 ) -> dict[str, object] | None: 

1029 """ 

1030 Pre-call hook to scan user prompts before sending to LLM. 

1031 

1032 Raises HTTPException if content should be blocked. 

1033 """ 

1034 verbose_proxy_logger.info("PANW Prisma AIRS: Running pre-call prompt scan") 

1035 

1036 # Check if guardrail should run for this request 

1037 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.pre_call 

1038 if self.should_run_guardrail(data=data, event_type=event_type) is not True: 

1039 return data 

1040 

1041 # Prevent duplicate scans by checking if already processed 

1042 if self._check_and_mark_scanned(data, "pre"): 

1043 return data 

1044 

1045 try: 

1046 start_time: Final = datetime.now() 

1047 

1048 # Extract prompt text from request 

1049 prompt_text: Final = self._extract_prompt_from_request(data) 

1050 messages: Final = data.get("messages", []) # Keep for masking operations 

1051 

1052 if not prompt_text: 

1053 verbose_proxy_logger.warning( 

1054 "PANW Prisma AIRS: No user prompt found in request (checked 'messages' and 'prompt' fields)" 

1055 ) 

1056 return None 

1057 

1058 # Prepare metadata - include user's metadata for profile override 

1059 metadata: Final = self._prepare_metadata_from_request(data) 

1060 

1061 # Scan prompt with PANW Prisma AIRS 

1062 scan_result: Final = await self._call_panw_api( 

1063 content=prompt_text, 

1064 is_response=False, 

1065 metadata=metadata, 

1066 call_id=data.get("litellm_call_id"), 

1067 ) 

1068 

1069 if scan_result.get("_is_transient") or scan_result.get("_always_block"): 

1070 return self._handle_api_error_with_logging( 

1071 scan_result, 

1072 data, 

1073 start_time, 

1074 is_response=False, 

1075 event_type=GuardrailEventHooks.pre_call, 

1076 ) 

1077 

1078 end_time: Final = datetime.now() 

1079 self.add_standard_logging_guardrail_information_to_request_data( 

1080 guardrail_provider=self._PROVIDER_NAME, 

1081 guardrail_json_response=scan_result, 

1082 request_data=data, 

1083 guardrail_status=("success" if scan_result.get("action") == "allow" else "guardrail_intervened"), 

1084 start_time=start_time.timestamp(), 

1085 end_time=end_time.timestamp(), 

1086 duration=(end_time - start_time).total_seconds(), 

1087 event_type=GuardrailEventHooks.pre_call, 

1088 ) 

1089 self._record_scan_id(data, scan_result, GuardrailEventHooks.pre_call) 

1090 

1091 action: Final = scan_result.get("action", "block") 

1092 category: Final = scan_result.get("category", "unknown") 

1093 masked_text: Final = self._get_masked_text(scan_result, is_response=False) 

1094 

1095 # If action is "allow", apply masking if available and allow through 

1096 if action == "allow": 

1097 if masked_text: 

1098 if messages: 

1099 data["messages"] = self._apply_masking_to_messages(messages, masked_text) 

1100 elif "prompt" in data: 

1101 data["prompt"] = masked_text 

1102 verbose_proxy_logger.info("PANW Prisma AIRS: Prompt allowed with masking (Category: %s)", category) 

1103 else: 

1104 verbose_proxy_logger.info("PANW Prisma AIRS: Prompt allowed (Category: %s)", category) 

1105 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

1106 return None 

1107 

1108 # Action is "block" - check if we should mask instead of blocking 

1109 if masked_text and self.mask_request_content: 

1110 if messages: 

1111 data["messages"] = self._apply_masking_to_messages(messages, masked_text) 

1112 elif "prompt" in data: 

1113 data["prompt"] = masked_text 

1114 verbose_proxy_logger.warning( 

1115 "PANW Prisma AIRS: Prompt blocked but masked instead (mask_request_content=True)" 

1116 ) 

1117 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

1118 return None 

1119 

1120 # Block the request 

1121 error_detail: Final = self._build_error_detail(scan_result, is_response=False) 

1122 verbose_proxy_logger.warning("PANW Prisma AIRS: %s", error_detail["error"]["message"]) 

1123 raise HTTPException(status_code=400, detail=error_detail) 

1124 

1125 except HTTPException: 

1126 raise 

1127 except Exception as e: 

1128 verbose_proxy_logger.error("PANW Prisma AIRS scan failed: %s", e) 

1129 raise HTTPException( 

1130 status_code=500, 

1131 detail={ 

1132 "error": { 

1133 "message": "Security scan failed - request blocked for safety", 

1134 "type": "guardrail_scan_error", 

1135 "code": "panw_prisma_airs_scan_failed", 

1136 "guardrail": self.guardrail_name, 

1137 } 

1138 }, 

1139 ) 

1140 

1141 @log_guardrail_information 

1142 async def async_post_call_success_hook( 

1143 self, 

1144 data: dict[str, object], 

1145 user_api_key_dict: UserAPIKeyAuth, 

1146 response: object, 

1147 ) -> object: 

1148 """ 

1149 Post-call hook to scan LLM responses before returning to user. 

1150 

1151 Raises HTTPException if response should be blocked. 

1152 """ 

1153 # Only process ModelResponse objects 

1154 if not isinstance(response, ModelResponse): 

1155 return response 

1156 

1157 verbose_proxy_logger.info("PANW Prisma AIRS: Running post-call response scan") 

1158 

1159 # Check if guardrail should run for this request 

1160 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.post_call 

1161 if self.should_run_guardrail(data=data, event_type=event_type) is not True: 

1162 return response 

1163 

1164 # Prevent duplicate scans by checking if already processed 

1165 if self._check_and_mark_scanned(data, "post"): 

1166 return response 

1167 

1168 try: 

1169 start_time: Final = datetime.now() 

1170 

1171 # Extract response text 

1172 response_text: Final = self._extract_response_text(response) 

1173 

1174 if not response_text: 

1175 verbose_proxy_logger.warning("PANW Prisma AIRS: No response content found to scan") 

1176 return response 

1177 

1178 # Prepare metadata - include user's metadata for profile override 

1179 metadata: Final = self._prepare_metadata_from_request(data) 

1180 

1181 # Scan response with PANW Prisma AIRS 

1182 scan_result: Final = await self._call_panw_api( 

1183 content=response_text, 

1184 is_response=True, 

1185 metadata=metadata, 

1186 call_id=data.get("litellm_call_id"), 

1187 ) 

1188 

1189 if scan_result.get("_is_transient") or scan_result.get("_always_block"): 

1190 self._handle_api_error_with_logging( 

1191 scan_result, 

1192 data, 

1193 start_time, 

1194 is_response=True, 

1195 event_type=GuardrailEventHooks.post_call, 

1196 ) 

1197 return response 

1198 

1199 end_time: Final = datetime.now() 

1200 self.add_standard_logging_guardrail_information_to_request_data( 

1201 guardrail_provider=self._PROVIDER_NAME, 

1202 guardrail_json_response=scan_result, 

1203 request_data=data, 

1204 guardrail_status=("success" if scan_result.get("action") == "allow" else "guardrail_intervened"), 

1205 start_time=start_time.timestamp(), 

1206 end_time=end_time.timestamp(), 

1207 duration=(end_time - start_time).total_seconds(), 

1208 event_type=GuardrailEventHooks.post_call, 

1209 ) 

1210 self._record_scan_id(data, scan_result, GuardrailEventHooks.post_call) 

1211 

1212 action: Final = scan_result.get("action", "block") 

1213 category: Final = scan_result.get("category", "unknown") 

1214 masked_text: Final = self._get_masked_text(scan_result, is_response=True) 

1215 

1216 # If action is "allow", apply masking if available and allow through 

1217 if action == "allow": 

1218 if masked_text: 

1219 self._apply_masking_to_response(response, masked_text) 

1220 verbose_proxy_logger.info( 

1221 "PANW Prisma AIRS: Response allowed with masking (Category: %s)", category 

1222 ) 

1223 else: 

1224 verbose_proxy_logger.info("PANW Prisma AIRS: Response allowed (Category: %s)", category) 

1225 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

1226 return response 

1227 

1228 # Action is "block" - check if we should mask instead of blocking 

1229 if masked_text and self.mask_response_content: 

1230 self._apply_masking_to_response(response, masked_text) 

1231 verbose_proxy_logger.warning( 

1232 "PANW Prisma AIRS: Response blocked but masked instead (mask_response_content=True)" 

1233 ) 

1234 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

1235 return response 

1236 

1237 # Block the response 

1238 error_detail: Final = self._build_error_detail(scan_result, is_response=True) 

1239 verbose_proxy_logger.warning("PANW Prisma AIRS: %s", error_detail["error"]["message"]) 

1240 raise HTTPException(status_code=400, detail=error_detail) 

1241 

1242 except HTTPException: 

1243 raise 

1244 except Exception as e: 

1245 verbose_proxy_logger.error("PANW Prisma AIRS scan failed: %s", e) 

1246 raise HTTPException( 

1247 status_code=500, 

1248 detail={ 

1249 "error": { 

1250 "message": "Security scan failed - response blocked for safety", 

1251 "type": "guardrail_scan_error", 

1252 "code": "panw_prisma_airs_scan_failed", 

1253 "guardrail": self.guardrail_name, 

1254 } 

1255 }, 

1256 ) 

1257 

1258 async def _scan_and_process_streaming_response( 

1259 self, 

1260 assembled_model_response: ModelResponse, 

1261 request_data: dict, 

1262 start_time: datetime, 

1263 ) -> tuple[bool, ModelResponse, dict[str, object]]: 

1264 """ 

1265 Scan assembled streaming response and apply masking if needed. 

1266 Returns (content_was_modified, response, scan_result). 

1267 """ 

1268 content_was_modified = False 

1269 response_text: Final = self._extract_response_text(assembled_model_response) 

1270 

1271 if not response_text or not response_text.strip(): 

1272 verbose_proxy_logger.info("PANW Prisma AIRS: No content to scan in streaming response") 

1273 return ( 

1274 content_was_modified, 

1275 assembled_model_response, 

1276 {"action": "allow", "category": "no_content"}, 

1277 ) 

1278 

1279 # Prepare metadata - include user's metadata for profile override 

1280 metadata: Final = self._prepare_metadata_from_request(request_data) 

1281 

1282 scan_result: Final = await self._call_panw_api( 

1283 content=response_text, 

1284 is_response=True, 

1285 metadata=metadata, 

1286 call_id=request_data.get("litellm_call_id"), 

1287 ) 

1288 

1289 # Early return for transient/always-block results — let the 

1290 # streaming iterator hook handle fallback_on_error semantics. 

1291 if scan_result.get("_is_transient") or scan_result.get("_always_block"): 

1292 return (content_was_modified, assembled_model_response, scan_result) 

1293 

1294 action: Final = scan_result.get("action", "block") 

1295 category: Final = scan_result.get("category", "unknown") 

1296 masked_text: Final = self._get_masked_text(scan_result, is_response=True) 

1297 

1298 # Handle scan results 

1299 if action == "allow": 

1300 if masked_text: 

1301 self._apply_masking_to_response(assembled_model_response, masked_text) 

1302 content_was_modified = True 

1303 verbose_proxy_logger.info( 

1304 "PANW Prisma AIRS: Streaming response allowed with masking (Category: %s)", category 

1305 ) 

1306 else: 

1307 verbose_proxy_logger.info("PANW Prisma AIRS: Streaming response allowed (Category: %s)", category) 

1308 elif masked_text and self.mask_response_content: 

1309 self._apply_masking_to_response(assembled_model_response, masked_text) 

1310 content_was_modified = True 

1311 verbose_proxy_logger.warning( 

1312 "PANW Prisma AIRS: Streaming response blocked but masked instead (mask_response_content=True)" 

1313 ) 

1314 else: 

1315 error_detail: Final = self._build_error_detail(scan_result, is_response=True) 

1316 verbose_proxy_logger.warning("PANW Prisma AIRS: %s", error_detail["error"]["message"]) 

1317 raise HTTPException(status_code=400, detail=error_detail) 

1318 

1319 return content_was_modified, assembled_model_response, scan_result 

1320 

1321 @log_guardrail_information 

1322 async def async_post_call_streaming_iterator_hook( 

1323 self, 

1324 user_api_key_dict: UserAPIKeyAuth, 

1325 response: AsyncIterable[object], 

1326 request_data: dict[str, object], 

1327 ): 

1328 """ 

1329 Process streaming response chunks and scan the assembled response. 

1330 """ 

1331 from litellm.llms.base_llm.base_model_iterator import MockResponseIterator 

1332 from litellm.main import stream_chunk_builder 

1333 

1334 # Check if guardrail should run for this request 

1335 

1336 if not self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call): 

1337 async for chunk in response: 

1338 yield chunk 

1339 return 

1340 

1341 # Prevent duplicate scans by checking if already processed 

1342 if self._check_and_mark_scanned(request_data, "streaming"): 

1343 async for chunk in response: 

1344 yield chunk 

1345 return 

1346 

1347 verbose_proxy_logger.info("PANW Prisma AIRS: Running post-call streaming scan") 

1348 

1349 all_chunks: Final = [] 

1350 content_was_modified = False 

1351 

1352 try: 

1353 start_time: Final = datetime.now() 

1354 

1355 # Collect all chunks 

1356 async for chunk in response: 

1357 all_chunks.append(chunk) 

1358 

1359 # Handle /v1/messages streaming: chunks are raw bytes (Anthropic SSE) 

1360 if all_chunks and isinstance(all_chunks[0], bytes): 

1361 text = self._extract_text_from_sse_bytes(all_chunks) 

1362 await self._scan_raw_streaming_text(text, request_data, start_time) 

1363 for chunk in all_chunks: 

1364 yield chunk 

1365 return 

1366 

1367 # Handle /v1/responses streaming: chunks are Pydantic events (not ModelResponse/ModelResponseStream) 

1368 if all_chunks and not isinstance(all_chunks[0], (ModelResponse, ModelResponseStream)): 

1369 text = self._extract_text_from_streaming_events(all_chunks) 

1370 await self._scan_raw_streaming_text(text, request_data, start_time) 

1371 for chunk in all_chunks: 

1372 yield chunk 

1373 return 

1374 

1375 # Assemble complete response from chunks 

1376 assembled_model_response = stream_chunk_builder(chunks=all_chunks) 

1377 

1378 if isinstance(assembled_model_response, ModelResponse): 

1379 # Scan and process the assembled response 

1380 ( 

1381 content_was_modified, 

1382 assembled_model_response, 

1383 scan_result, 

1384 ) = await self._scan_and_process_streaming_response(assembled_model_response, request_data, start_time) 

1385 

1386 if scan_result.get("_is_transient") or scan_result.get("_always_block"): 

1387 self._handle_api_error_with_logging( 

1388 scan_result, 

1389 request_data, 

1390 start_time, 

1391 is_response=True, 

1392 event_type=GuardrailEventHooks.post_call, 

1393 ) 

1394 # Control only reaches here for _is_transient errors with 

1395 # fallback_on_error="allow"; _always_block and fail-closed 

1396 # paths raise inside _handle_api_error_with_logging above. 

1397 for chunk in all_chunks: 

1398 yield chunk 

1399 return 

1400 

1401 end_time: Final = datetime.now() 

1402 self.add_standard_logging_guardrail_information_to_request_data( 

1403 guardrail_provider=self._PROVIDER_NAME, 

1404 guardrail_json_response=scan_result, 

1405 request_data=request_data, 

1406 guardrail_status=("success" if scan_result.get("action") == "allow" else "guardrail_intervened"), 

1407 start_time=start_time.timestamp(), 

1408 end_time=end_time.timestamp(), 

1409 duration=(end_time - start_time).total_seconds(), 

1410 event_type=GuardrailEventHooks.post_call, 

1411 ) 

1412 self._record_scan_id(request_data, scan_result, GuardrailEventHooks.post_call) 

1413 

1414 # Add guardrail to applied guardrails header for observability 

1415 add_guardrail_to_applied_guardrails_header( 

1416 request_data=request_data, guardrail_name=self.guardrail_name 

1417 ) 

1418 

1419 # Only use MockResponseIterator if content was modified 

1420 # Otherwise, yield original chunks to preserve streaming behavior 

1421 if content_was_modified: 

1422 mock_response: Final = MockResponseIterator(model_response=assembled_model_response) 

1423 async for chunk in mock_response: 

1424 yield chunk 

1425 else: 

1426 for chunk in all_chunks: 

1427 yield chunk 

1428 else: 

1429 # stream_chunk_builder returned None; yield original chunks unmodified 

1430 for chunk in all_chunks: 

1431 yield chunk 

1432 

1433 except HTTPException as e: 

1434 # Yield error as SSE event so create_response() detects it and 

1435 # returns a proper JSON error response with the correct status code. 

1436 # (Raising from a generator hits create_response's generic except → 500.) 

1437 detail: Final = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)} 

1438 error_obj: Final[dict[str, object]] = dict(detail.get("error", detail)) 

1439 error_obj["code"] = e.status_code 

1440 yield f"data: {json.dumps({'error': error_obj})}\n\n" 

1441 except Exception as e: 

1442 verbose_proxy_logger.error("PANW Prisma AIRS streaming error: %s", e) 

1443 yield f"data: {json.dumps({'error': {'message': 'Security scan failed - streaming response blocked for safety', 'type': 'guardrail_scan_error', 'code': 500, 'guardrail': self.guardrail_name}})}\n\n" 

1444 

1445 async def _scan_tool_calls_for_guardrail( 

1446 self, 

1447 tool_calls: list, 

1448 is_response: bool, 

1449 metadata: Mapping[str, object], 

1450 call_id: object, 

1451 request_data: dict, 

1452 start_time: datetime, 

1453 ) -> None: 

1454 """Scan tool calls with allow/block/mask treatment (in-place modification). 

1455 

1456 Tool name and arguments go out as plain prompt/response text, newline separated: 

1457 the AIRS ``tool_event`` schema only accepts ``ecosystem: "mcp"``, which 

1458 OpenAI-format tool calls are not. A name-only call is still scanned so 

1459 tool-name policies keep firing on empty arguments. 

1460 """ 

1461 for tool_call in tool_calls: 

1462 tool_name, args_text = self._get_tool_call_function(tool_call) 

1463 scanned_args = args_text if args_text and args_text.strip() else None 

1464 scan_text = "\n".join(part for part in (tool_name, scanned_args) if part) 

1465 if not scan_text.strip(): 

1466 continue 

1467 

1468 scan_result = await self._call_panw_api( 

1469 content=scan_text, 

1470 is_response=is_response, 

1471 metadata=metadata, 

1472 call_id=call_id, 

1473 ) 

1474 

1475 if scan_result.get("_is_transient") or scan_result.get("_always_block"): 

1476 event_type = GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call 

1477 self._handle_api_error_with_logging( 

1478 scan_result=scan_result, 

1479 data=request_data, 

1480 start_time=start_time, 

1481 event_type=event_type, 

1482 is_response=is_response, 

1483 ) 

1484 continue 

1485 

1486 self._record_scan_id( 

1487 request_data, 

1488 scan_result, 

1489 GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call, 

1490 ) 

1491 

1492 action = scan_result.get("action", "block") 

1493 masked_args = self._masked_tool_call_arguments( 

1494 self._get_masked_text(scan_result, is_response=is_response), 

1495 scanned_name=bool(tool_name), 

1496 scanned_args=scanned_args, 

1497 ) 

1498 

1499 if action == "allow": 

1500 if masked_args: 

1501 self._set_tool_call_arguments(tool_call, masked_args) 

1502 elif masked_args and ( 

1503 (is_response and self.mask_response_content) or (not is_response and self.mask_request_content) 

1504 ): 

1505 self._set_tool_call_arguments(tool_call, masked_args) 

1506 else: 

1507 # Tool calls now go out as ordinary prompt/response text, so a 

1508 # response-side scan reports the model's arguments under 

1509 # response_masked_data, which _CLIENT_HIDDEN_SCAN_FIELDS already 

1510 # withholds. prompt_masked_data is the caller's own input again and 

1511 # must keep reaching them -- it is one of the fields LIT-5638 asks for. 

1512 error_detail = self._build_error_detail(scan_result, is_response=is_response) 

1513 raise HTTPException(status_code=400, detail=error_detail) 

1514 

1515 @staticmethod 

1516 def _masked_tool_call_arguments( 

1517 masked_text: str | None, 

1518 *, 

1519 scanned_name: bool, 

1520 scanned_args: str | None, 

1521 ) -> str | None: 

1522 """Recover the arguments slice of a masked scan, or None when it cannot be applied.""" 

1523 if masked_text is None or scanned_args is None: 

1524 return None 

1525 if not scanned_name: 

1526 return masked_text 

1527 _, separator, masked_args = masked_text.partition("\n") 

1528 return masked_args if separator else None 

1529 

1530 @staticmethod 

1531 def _get_tool_call_function(tool_call: ToolCallLike) -> tuple[str | None, str | None]: 

1532 """Read a tool call's function name and arguments; (None, None) for non-function shapes.""" 

1533 try: 

1534 parsed: Final = _ToolCallSlice.model_validate(tool_call, from_attributes=True) 

1535 except ValidationError: 

1536 return (None, None) 

1537 if parsed.function is None: 

1538 return (None, None) 

1539 return (parsed.function.name, parsed.function.arguments) 

1540 

1541 @staticmethod 

1542 def _set_tool_call_arguments(tool_call: ToolCallLike, masked_text: str) -> None: 

1543 """Set masked text on the function arguments of a call that _get_tool_call_function accepted.""" 

1544 if isinstance(tool_call, dict): 

1545 tool_call["function"]["arguments"] = masked_text 

1546 return 

1547 if isinstance(tool_call, ChatCompletionMessageCustomToolCall | ChatCompletionDeltaCustomToolCall): 

1548 return 

1549 tool_call.function.arguments = masked_text 

1550 

1551 @staticmethod 

1552 def _is_anthropic_request( 

1553 request_data: Mapping[str, object], 

1554 logging_obj: Optional["LiteLLMLoggingObj"] = None, 

1555 ) -> bool: 

1556 """Detect if the current request is an Anthropic /v1/messages call.""" 

1557 if logging_obj: 

1558 call_type: Final = getattr(logging_obj, "call_type", None) 

1559 if call_type in ( 

1560 CallTypes.anthropic_messages.value, 

1561 CallTypes.anthropic_messages, 

1562 ): 

1563 return True 

1564 psr: Final = request_data.get("proxy_server_request") or {} 

1565 if not isinstance(psr, dict): 

1566 return False 

1567 url: Final = psr.get("url") or "" 

1568 if not isinstance(url, str): 

1569 return False 

1570 # Match exact path segments, not substring (avoid matching e.g. /v1/messages_batch) 

1571 path: Final = urlparse(url).path.rstrip("/") 

1572 if path.endswith("/v1/messages"): 

1573 return True 

1574 return False 

1575 

1576 def _use_latest_user_only( 

1577 self, 

1578 request_data: Mapping[str, object], 

1579 logging_obj: Optional["LiteLLMLoggingObj"] = None, 

1580 ) -> bool: 

1581 """Resolve whether to scan only the latest user message. 

1582 

1583 - Non-Anthropic requests: always False (existing behavior) 

1584 - Anthropic requests: 

1585 - Flag explicitly True/False: respect it 

1586 - Flag None (not set): default to True 

1587 """ 

1588 if not self._is_anthropic_request(request_data, logging_obj): 

1589 return False 

1590 if self.experimental_use_latest_role_message_only is None: 

1591 return True # Default-on for Anthropic 

1592 return self.experimental_use_latest_role_message_only 

1593 

1594 @staticmethod 

1595 def _get_latest_user_text_indices( 

1596 texts: Sequence[str], 

1597 messages: Sequence[object], 

1598 ) -> set | None: 

1599 """Return text indices belonging to only the latest scannable human-authored (user or developer) message. 

1600 

1601 Args: 

1602 texts: Flattened text entries from the framework. 

1603 messages: The structured messages the framework flattened into ``texts``, 

1604 hoisted top-level system prompt included, so positions line up. 

1605 

1606 Returns a set of scannable indices, or None on count mismatch or no user/developer 

1607 message (safety fallback to existing role-filter behavior). 

1608 """ 

1609 last_human_msg_idx: int | None = None 

1610 for idx in range(len(messages) - 1, -1, -1): 

1611 msg = messages[idx] 

1612 if isinstance(msg, dict) and msg.get("role") in ("user", "developer"): 

1613 last_human_msg_idx = idx 

1614 break 

1615 

1616 if last_human_msg_idx is None: 

1617 return None # No user/developer message → fallback to existing role-filter scan 

1618 

1619 scannable: Final[set] = set() 

1620 text_idx = 0 

1621 for msg_idx, msg in enumerate(messages): 

1622 if not isinstance(msg, dict): 

1623 continue 

1624 content = msg.get("content") 

1625 is_latest_human = msg_idx == last_human_msg_idx 

1626 

1627 if content is None: 

1628 pass 

1629 elif isinstance(content, str): 

1630 if is_latest_human: 

1631 scannable.add(text_idx) 

1632 text_idx += 1 

1633 elif isinstance(content, list): 

1634 for item in content: 

1635 if isinstance(item, dict) and item.get("text") is not None: 

1636 if is_latest_human: 

1637 scannable.add(text_idx) 

1638 text_idx += 1 

1639 

1640 if text_idx != len(texts): 

1641 return None # Count mismatch → safety fallback 

1642 

1643 return scannable 

1644 

1645 def supports_scan_only_tool_results(self) -> bool: 

1646 return False 

1647 

1648 @staticmethod 

1649 def _get_scannable_text_indices( 

1650 texts: Sequence[str], 

1651 structured_messages: Sequence[object], 

1652 ) -> set | None: 

1653 """Derive which ``texts`` indices originate from user/system messages. 

1654 

1655 The unified guardrail framework flattens message content into ``texts`` 

1656 without preserving role info. This helper re-walks 

1657 ``structured_messages`` using the **same** extraction logic the 

1658 framework uses (string content → 1 entry, list content → 1 per text 

1659 item, None → 0) and records the running text index for each entry 

1660 whose source role is ``"user"``, ``"system"``, or ``"developer"``. 

1661 

1662 Returns a set of scannable indices, or ``None`` if the count doesn't 

1663 match ``len(texts)`` (safety fallback → scan everything). 

1664 """ 

1665 scannable: Final[set] = set() 

1666 text_idx = 0 

1667 for msg in structured_messages: 

1668 if not isinstance(msg, dict): 

1669 continue 

1670 role = msg.get("role", "") 

1671 content = msg.get("content") 

1672 is_scannable = role in ("user", "system", "developer") 

1673 

1674 if content is None: 

1675 # No content → 0 text entries 

1676 pass 

1677 elif isinstance(content, str): 

1678 if is_scannable: 

1679 scannable.add(text_idx) 

1680 text_idx += 1 

1681 elif isinstance(content, list): 

1682 for item in content: 

1683 if isinstance(item, dict) and item.get("text") is not None: 

1684 if is_scannable: 

1685 scannable.add(text_idx) 

1686 text_idx += 1 

1687 # Ignore other content types (shouldn't happen) 

1688 

1689 if text_idx != len(texts): 

1690 # Count mismatch → safety fallback: scan all 

1691 return None 

1692 

1693 return scannable 

1694 

1695 @staticmethod 

1696 def _mcp_name_fallback(rd: dict) -> str | None: 

1697 """Return rd['name'] only when 'arguments' or 'mcp_arguments' co-occurs (MCP shape). 

1698 

1699 A bare 'name' key without 'arguments' is NOT an MCP request — it's a 

1700 stray field from the chat completion body that should be ignored. 

1701 """ 

1702 return rd.get("name") if ("arguments" in rd or "mcp_arguments" in rd) else None 

1703 

1704 @log_guardrail_information 

1705 async def apply_guardrail( 

1706 self, 

1707 inputs: GenericGuardrailAPIInputs, 

1708 request_data: dict[str, object], 

1709 input_type: Literal["request", "response"], 

1710 logging_obj: Optional["LiteLLMLoggingObj"] = None, 

1711 ) -> GenericGuardrailAPIInputs: 

1712 """ 

1713 Unified guardrail method for the apply_guardrail framework. 

1714 

1715 Called by the UI "Test Guardrail" endpoint, UnifiedLLMGuardrails orchestrator, 

1716 and MCP tool input scanning. 

1717 """ 

1718 texts: Final = inputs.get("texts", []) 

1719 is_response: Final = input_type == "response" 

1720 

1721 # Resolve litellm_call_id: request_data first, then logging_obj fallback. 

1722 # Post-call path reconstructs request_data as {"response": ...} without 

1723 # litellm_call_id, but logging_obj.litellm_call_id is available. 

1724 call_id = request_data.get("litellm_call_id") 

1725 if not call_id and logging_obj: 

1726 call_id = getattr(logging_obj, "litellm_call_id", None) 

1727 if not call_id: 

1728 # Use MCP name fallback: mcp_tool_name (canonical) or name (/mcp-rest path) 

1729 _mcp_tool = str(request_data.get("mcp_tool_name") or self._mcp_name_fallback(request_data) or "").strip() 

1730 if input_type == "request" and logging_obj is None and _mcp_tool: 

1731 # Synthesize a tool-prefixed call_id for AIRS grouping. 

1732 # Slug: lowercase, non-alphanum → "-", truncate to 40 chars. 

1733 slug = re.sub(r"[^a-z0-9]+", "-", _mcp_tool.lower()).strip("-")[:40] 

1734 if not slug: 

1735 slug = "mcp-tool" 

1736 call_id = f"{slug}-{uuid.uuid4()}" 

1737 request_data["litellm_call_id"] = call_id 

1738 verbose_proxy_logger.debug( 

1739 "PANW Prisma AIRS: synthesized MCP tr_id=%s for tool=%s", 

1740 call_id, 

1741 _mcp_tool, 

1742 ) 

1743 elif not request_data and logging_obj is None and input_type == "request": 

1744 # Direct /apply_guardrail endpoint — empty request_data, no 

1745 # logging_obj. Existing behavior: synthesize UUID. 

1746 call_id = str(uuid.uuid4()) 

1747 request_data["litellm_call_id"] = call_id 

1748 verbose_proxy_logger.warning( 

1749 "PANW Prisma AIRS: litellm_call_id missing from empty " 

1750 "request_data, synthesized %s (direct /apply_guardrail?)", 

1751 call_id, 

1752 ) 

1753 else: 

1754 call_id = str(uuid.uuid4()) 

1755 request_data["litellm_call_id"] = call_id 

1756 verbose_proxy_logger.warning( 

1757 "PANW Prisma AIRS: litellm_call_id missing, synthesized %s (input_type=%s)", 

1758 call_id, 

1759 input_type, 

1760 ) 

1761 

1762 # Enrich request_data with model if missing (post-call metadata loss) 

1763 if not request_data.get("model"): 

1764 if inputs.get("model"): 

1765 request_data["model"] = inputs["model"] 

1766 elif logging_obj: 

1767 request_data["model"] = getattr(logging_obj, "model", None) 

1768 

1769 # Enrich request_data with metadata from logging_obj (post-call metadata loss). 

1770 # Merge: logging_obj provides the base, request_data keys win on conflict. 

1771 if logging_obj: 

1772 _lp: Final = (getattr(logging_obj, "model_call_details", {}) or {}).get("litellm_params", {}) or {} 

1773 _orig_meta: Final = _lp.get("metadata") or {} 

1774 if _orig_meta: 

1775 existing_meta = request_data.get("metadata") 

1776 if not isinstance(existing_meta, dict): 

1777 existing_meta = {} 

1778 request_data["metadata"] = {**_orig_meta, **existing_meta} 

1779 

1780 metadata: Final = self._prepare_metadata_from_request(request_data) 

1781 start_time: Final = datetime.now() 

1782 new_texts: Final[list[str]] = [] 

1783 

1784 # On request side, determine which text indices correspond to scannable 

1785 # messages so we can skip scanning assistant/tool history text. 

1786 scannable_indices: set | None = None 

1787 if input_type == "request": 

1788 structured_messages: Final = inputs.get("structured_messages") 

1789 if structured_messages: 

1790 # For Anthropic /v1/messages: default to latest-user-only scanning. 

1791 if self._use_latest_user_only(request_data, logging_obj): 

1792 scannable_indices = self._get_latest_user_text_indices(texts, structured_messages) 

1793 # Fall through to existing role filtering if: 

1794 # - not Anthropic, OR flag explicitly False, OR 

1795 # - latest-user extraction returned None (no user / count mismatch) 

1796 if scannable_indices is None: 

1797 scannable_indices = self._get_scannable_text_indices(texts, structured_messages) 

1798 if ( 

1799 scannable_indices is not None 

1800 and not scannable_indices 

1801 and effective_scan_only_tool_results_for_guardrail(self) 

1802 ): 

1803 verbose_proxy_logger.warning( 

1804 "PANW Prisma AIRS scans only user, system, and developer messages, " 

1805 "so scan_only_tool_results leaves nothing to scan for this request" 

1806 ) 

1807 

1808 for i, text in enumerate(texts): 

1809 if not text or not text.strip(): 

1810 new_texts.append(text) 

1811 continue 

1812 

1813 # Skip non-user/system texts on request side 

1814 if scannable_indices is not None and i not in scannable_indices: 

1815 new_texts.append(text) 

1816 continue 

1817 

1818 scan_result = await self._call_panw_api( 

1819 content=text, 

1820 is_response=is_response, 

1821 metadata=metadata, 

1822 call_id=call_id, 

1823 ) 

1824 

1825 # Handle API errors (transient/config) 

1826 if scan_result.get("_is_transient") or scan_result.get("_always_block"): 

1827 event_type = GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call 

1828 self._handle_api_error_with_logging( 

1829 scan_result=scan_result, 

1830 data=request_data, 

1831 start_time=start_time, 

1832 event_type=event_type, 

1833 is_response=is_response, 

1834 ) 

1835 # If we reach here, fallback_on_error="allow" 

1836 new_texts.append(text) 

1837 continue 

1838 

1839 self._record_scan_id( 

1840 request_data, 

1841 scan_result, 

1842 GuardrailEventHooks.post_call if is_response else GuardrailEventHooks.pre_call, 

1843 ) 

1844 

1845 action = scan_result.get("action", "block") 

1846 masked_text = self._get_masked_text(scan_result, is_response=is_response) 

1847 

1848 if action == "allow": 

1849 new_texts.append(masked_text if masked_text else text) 

1850 elif masked_text and ( 

1851 (is_response and self.mask_response_content) or (not is_response and self.mask_request_content) 

1852 ): 

1853 new_texts.append(masked_text) 

1854 else: 

1855 error_detail = self._build_error_detail(scan_result, is_response=is_response) 

1856 raise HTTPException(status_code=400, detail=error_detail) 

1857 

1858 # Scan tool call arguments — same masking policy as texts. 

1859 # In-place modifications propagate for pre-call and OpenAI post-call. 

1860 # Anthropic post-call drops tool_call modifications (framework limitation). 

1861 tool_calls: Final = inputs.get("tool_calls", []) 

1862 if tool_calls: 

1863 await self._scan_tool_calls_for_guardrail( 

1864 tool_calls=tool_calls, 

1865 is_response=is_response, 

1866 metadata=metadata, 

1867 call_id=call_id, 

1868 request_data=request_data, 

1869 start_time=start_time, 

1870 ) 

1871 

1872 # MCP REST tool invocation scan (request-side only). 

1873 # When an MCP tool is being invoked via /mcp-rest/tools/call, the 

1874 # proxy sets mcp_tool_name (and optional mcp_arguments) on request_data. 

1875 # We send a tool_event so AIRS can apply tool-aware policies. 

1876 # REST MCP path sets "name"/"arguments"; canonical keys are 

1877 # "mcp_tool_name"/"mcp_arguments". Check canonical first, then fallback. 

1878 mcp_tool_name: Final = request_data.get("mcp_tool_name") or self._mcp_name_fallback(request_data) 

1879 if mcp_tool_name and input_type == "request": 

1880 mcp_tool_event: Final[dict[str, object]] = { 

1881 "metadata": { 

1882 "ecosystem": "mcp", 

1883 "method": "tools/call", 

1884 "server_name": self._get_mcp_server_name(request_data, mcp_tool_name), 

1885 "tool_invoked": mcp_tool_name, 

1886 }, 

1887 } 

1888 mcp_arguments = request_data.get("mcp_arguments") 

1889 if mcp_arguments is None: 

1890 mcp_arguments = request_data.get("arguments") 

1891 if mcp_arguments is not None and mcp_arguments != "": 

1892 if isinstance(mcp_arguments, (dict, list)): 

1893 serialized_args = json.dumps(mcp_arguments) 

1894 else: 

1895 serialized_args = str(mcp_arguments) 

1896 if serialized_args.strip(): 

1897 mcp_tool_event["input"] = serialized_args 

1898 

1899 mcp_scan_result: Final = await self._call_panw_api( 

1900 tool_event=mcp_tool_event, 

1901 metadata=metadata, 

1902 call_id=call_id, 

1903 ) 

1904 

1905 if mcp_scan_result.get("_is_transient") or mcp_scan_result.get("_always_block"): 

1906 self._handle_api_error_with_logging( 

1907 scan_result=mcp_scan_result, 

1908 data=request_data, 

1909 start_time=start_time, 

1910 event_type=GuardrailEventHooks.pre_call, 

1911 is_response=False, 

1912 ) 

1913 # If we reach here, fallback_on_error="allow" 

1914 else: 

1915 self._record_scan_id(request_data, mcp_scan_result, GuardrailEventHooks.pre_call) 

1916 action = mcp_scan_result.get("action", "block") 

1917 masked_text = self._get_masked_text(mcp_scan_result, is_response=False) 

1918 if action == "allow": 

1919 # PANW says OK — apply PII scrubbing if present (unconditional, 

1920 # matching _scan_tool_calls_for_guardrail behavior). 

1921 if masked_text: 

1922 self._apply_mcp_masking( 

1923 request_data, 

1924 mcp_arguments, 

1925 masked_text, 

1926 is_blocked=False, 

1927 ) 

1928 elif masked_text and self.mask_request_content: 

1929 self._apply_mcp_masking(request_data, mcp_arguments, masked_text) 

1930 else: 

1931 error_detail = self._build_error_detail(mcp_scan_result, is_response=False) 

1932 raise HTTPException(status_code=400, detail=error_detail) 

1933 

1934 inputs["texts"] = new_texts 

1935 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) 

1936 return inputs 

1937 

1938 @staticmethod 

1939 def get_config_model() -> type["GuardrailConfigModel"] | None: 

1940 from litellm.types.proxy.guardrails.guardrail_hooks.panw_prisma_airs import ( 

1941 PanwPrismaAirsGuardrailConfigModel, 

1942 ) 

1943 

1944 return PanwPrismaAirsGuardrailConfigModel 

1945 

1946 @classmethod 

1947 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

1948 return [ 

1949 GuardrailEventHooks.pre_call, 

1950 GuardrailEventHooks.during_call, 

1951 GuardrailEventHooks.post_call, 

1952 GuardrailEventHooks.logging_only, 

1953 GuardrailEventHooks.pre_mcp_call, 

1954 GuardrailEventHooks.during_mcp_call, 

1955 ]