Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py: 11%
1085 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Cisco AI Defense guardrail integration for LiteLLM.
4Cisco AI Defense exposes two distinct inspection surfaces, each with its own
5endpoint:
7* Chat inspection: POST <base>/api/v1/inspect/chat — LLM conversations
8* MCP inspection: POST <base>/api/v1/inspect/mcp — MCP tool calls
10Each guardrail instance targets exactly one surface, chosen via the
11``inspection_type`` dropdown:
13* ``chat`` — scan LLM model traffic only
14* ``mcp`` — scan MCP tool-call traffic only
16Configure two separate guardrails if you need both surfaces scanned. Each
17request is sent with the ``X-Cisco-AI-Defense-API-Key`` header.
18"""
20import json
21import os
22from collections.abc import AsyncIterator, Mapping, Sequence
23from dataclasses import dataclass, replace
24from datetime import datetime
25from typing import TYPE_CHECKING, Any, Final, Literal
27import httpx
28from fastapi import HTTPException
29from typing_extensions import TypedDict, Unpack
31from litellm import DualCache
32from litellm._logging import verbose_proxy_logger
33from litellm._version import version as litellm_version
34from litellm.integrations.custom_guardrail import (
35 CustomGuardrail,
36 log_guardrail_information,
37)
38from litellm.llms.custom_httpx.http_handler import (
39 get_async_httpx_client,
40 httpxSpecialProvider,
41)
42from litellm.proxy._types import UserAPIKeyAuth
43from litellm.proxy.common_utils.callback_utils import (
44 add_guardrail_to_applied_guardrails_header,
45)
46from litellm.types.guardrails import GuardrailEventHooks
47from litellm.types.utils import (
48 Choices,
49 LLMResponseTypes,
50 ModelResponse,
51 ModelResponseStream,
52 TextCompletionResponse,
53)
55from .cisco_ai_defense_mcp import _CiscoAIDefenseMcpMixin
57if TYPE_CHECKING: 57 ↛ 58line 57 didn't jump to line 58 because the condition on line 57 was never true
58 from litellm.types.proxy.guardrails.guardrail_hooks.base import (
59 GuardrailConfigModel,
60 )
63CISCO_DEFAULT_API_BASE: Final = "https://us.api.inspect.aidefense.security.cisco.com"
64CISCO_CHAT_INSPECT_PATH: Final = "/api/v1/inspect/chat"
65CISCO_MCP_INSPECT_PATH: Final = "/api/v1/inspect/mcp"
66CISCO_API_KEY_HEADER: Final = "X-Cisco-AI-Defense-API-Key"
67DEFAULT_TIMEOUT_SECONDS: Final = 10.0
69SUPPORTED_INSPECTION_TYPES: Final[tuple[str, ...]] = ("chat", "mcp")
70DEFAULT_INSPECTION_TYPE: Final = "chat"
72# LiteLLM marks MCP guardrail calls with these call_type values; the proxy
73# routes pre_mcp_call / during_mcp_call events through async_pre_call_hook /
74# async_moderation_hook with the call_type set accordingly.
75_MCP_CALL_TYPES: Final[tuple[str, ...]] = ("mcp_call", "call_mcp_tool")
77# Action vocabulary Cisco AI Defense can return.
78_ACTION_BLOCK: Final = "block"
79_ACTION_REDACT: Final = "redact"
80_ACTION_ALLOW: Final = "allow"
83@dataclass(frozen=True, slots=True)
84class _ScanContext:
85 """The surface (``chat`` / ``mcp``) and direction (``input`` / ``output``) a scan targets."""
87 surface: str
88 direction: str
91@dataclass(frozen=True, slots=True)
92class _CiscoVerdict:
93 """Parsed Cisco AI Defense decision plus any sanitized rewrites it carries."""
95 is_safe: bool | None
96 classifications: list[str]
97 severity: str | None
98 rules: list[dict[str, object]]
99 explanation: str | None
100 event_id: str | None
101 action: str | None = None
102 sanitized_text: str | None = None
103 sanitized_messages: list[dict[str, object]] | None = None
104 sanitized_mcp_arguments: dict[str, object] | None = None
107class CiscoAIDefenseGuardrailMissingSecrets(Exception):
108 """Raised when the Cisco AI Defense API key is missing."""
111class CiscoAIDefenseGuardrailAPIError(Exception):
112 """Raised when there is an error talking to the Cisco AI Defense API."""
115class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
116 """Base-class constructor options this guardrail forwards untouched to CustomGuardrail."""
119class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail):
120 """
121 Cisco AI Defense guardrail integration.
123 Each instance scans exactly one inspection surface (``chat`` or ``mcp``)
124 via the corresponding Cisco AI Defense Inspection API endpoint.
126 MCP-specific hooks and helpers live on ``_CiscoAIDefenseMcpMixin`` in
127 ``cisco_ai_defense_mcp.py``.
128 """
130 SUPPORTED_ON_FLAGGED_ACTIONS: tuple[str, ...] = ("block", "monitor")
131 DEFAULT_ON_FLAGGED_ACTION: str = "block"
132 SUPPORTED_FALLBACK_ACTIONS: tuple[str, ...] = ("allow", "block")
133 DEFAULT_FALLBACK_ON_ERROR: str = "block"
135 _PROVIDER_NAME = "cisco_ai_defense"
137 def __init__(
138 self,
139 guardrail_name: str | None = "cisco-ai-defense",
140 api_key: str | None = None,
141 api_base: str | None = None,
142 inspection_type: str | None = None,
143 inspect_path: str | None = None,
144 enabled_rules: Sequence[object] | None = None,
145 integration_profile_id: str | None = None,
146 integration_profile_version: str | None = None,
147 integration_tenant_id: str | None = None,
148 integration_type: str | None = None,
149 on_flagged_action: str | None = None,
150 fallback_on_error: str | None = None,
151 timeout: float | None = None,
152 **kwargs: Unpack[_CustomGuardrailOptions],
153 ) -> None:
154 resolved_api_key: Final = api_key or os.environ.get("CISCO_AI_DEFENSE_API_KEY")
155 if not resolved_api_key:
156 raise CiscoAIDefenseGuardrailMissingSecrets(
157 "Cisco AI Defense API key is required. Set "
158 "`CISCO_AI_DEFENSE_API_KEY` in the environment or pass "
159 "`api_key` in the guardrail config."
160 )
161 self.api_key: str = resolved_api_key
163 self.api_base: str = (api_base or os.environ.get("CISCO_AI_DEFENSE_API_BASE") or CISCO_DEFAULT_API_BASE).rstrip(
164 "/"
165 )
167 self.inspection_type: str = self._resolve_choice(
168 value=inspection_type,
169 env_var="CISCO_AI_DEFENSE_INSPECTION_TYPE",
170 allowed=SUPPORTED_INSPECTION_TYPES,
171 default=DEFAULT_INSPECTION_TYPE,
172 setting_name="inspection_type",
173 )
175 inferred: Final = self._infer_inspection_type_from_mode(kwargs.get("event_hook"), self.inspection_type)
176 if inferred != self.inspection_type:
177 verbose_proxy_logger.info(
178 "Cisco AI Defense: inferred inspection_type=%s from MCP-only event_hook configuration (was %s)",
179 inferred,
180 self.inspection_type,
181 )
182 self.inspection_type = inferred
184 if inspect_path:
185 self.inspect_path = inspect_path if inspect_path.startswith("/") else f"/{inspect_path}"
186 else:
187 self.inspect_path = CISCO_MCP_INSPECT_PATH if self.inspection_type == "mcp" else CISCO_CHAT_INSPECT_PATH
189 self.enabled_rules = [self._normalize_rule(rule) for rule in enabled_rules] if enabled_rules else None
190 self.integration_profile_id = integration_profile_id
191 self.integration_profile_version = integration_profile_version
192 self.integration_tenant_id = integration_tenant_id
193 self.integration_type = integration_type
195 self.on_flagged_action = self._resolve_choice(
196 value=on_flagged_action,
197 env_var="CISCO_AI_DEFENSE_ON_FLAGGED_ACTION",
198 allowed=self.SUPPORTED_ON_FLAGGED_ACTIONS,
199 default=self.DEFAULT_ON_FLAGGED_ACTION,
200 setting_name="on_flagged_action",
201 )
203 self.fallback_on_error = self._resolve_choice(
204 value=fallback_on_error,
205 env_var="CISCO_AI_DEFENSE_FALLBACK_ON_ERROR",
206 allowed=self.SUPPORTED_FALLBACK_ACTIONS,
207 default=self.DEFAULT_FALLBACK_ON_ERROR,
208 setting_name="fallback_on_error",
209 )
211 resolved_timeout: float | None
212 if timeout is not None:
213 resolved_timeout = self._coerce_timeout(timeout)
214 else:
215 env_timeout: Final = os.environ.get("CISCO_AI_DEFENSE_TIMEOUT")
216 resolved_timeout = self._coerce_timeout(env_timeout) if env_timeout is not None else None
217 self.timeout: float = resolved_timeout if resolved_timeout is not None else DEFAULT_TIMEOUT_SECONDS
219 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
221 # Register broadly; runtime filtering happens in ``_surface_matches``.
222 super().__init__(
223 guardrail_name=guardrail_name,
224 supported_event_hooks=list(self.get_supported_event_hooks()),
225 **kwargs,
226 )
228 self._warn_if_mode_surface_mismatch(kwargs.get("event_hook"))
230 verbose_proxy_logger.debug(
231 "Cisco AI Defense guardrail initialized: name=%s, "
232 "inspection_type=%s, url=%s%s, on_flagged_action=%s, "
233 "fallback_on_error=%s, timeout=%ss",
234 guardrail_name,
235 self.inspection_type,
236 self.api_base,
237 self.inspect_path,
238 self.on_flagged_action,
239 self.fallback_on_error,
240 self.timeout,
241 )
243 # ------------------------------------------------------------------
244 # Configuration helpers
245 # ------------------------------------------------------------------
247 @staticmethod
248 def _resolve_choice(
249 value: str | None,
250 env_var: str,
251 allowed: tuple[str, ...],
252 default: str,
253 setting_name: str,
254 ) -> str:
255 candidate: Final = value if value is not None else os.environ.get(env_var)
256 if candidate is None:
257 return default
258 if candidate in allowed:
259 return candidate
260 verbose_proxy_logger.warning(
261 "Cisco AI Defense guardrail: invalid value '%s' for %s, falling back to default '%s'. Allowed values: %s",
262 candidate,
263 setting_name,
264 default,
265 ", ".join(allowed),
266 )
267 return default
269 @staticmethod
270 def _coerce_timeout(value: str | float) -> float | None:
271 try:
272 parsed: Final = float(value)
273 except (TypeError, ValueError):
274 verbose_proxy_logger.warning(
275 "Cisco AI Defense guardrail: invalid timeout value '%s', using default %ss",
276 value,
277 DEFAULT_TIMEOUT_SECONDS,
278 )
279 return None
280 if parsed < 1.0:
281 return 1.0
282 if parsed > 60.0:
283 return 60.0
284 return parsed
286 @staticmethod
287 def _is_mcp_call_type(call_type: str | None) -> bool:
288 return bool(call_type) and call_type in _MCP_CALL_TYPES
290 # ------------------------------------------------------------------
291 # Hook methods
292 # ------------------------------------------------------------------
294 @log_guardrail_information
295 async def async_pre_call_hook(
296 self,
297 user_api_key_dict: UserAPIKeyAuth,
298 cache: DualCache,
299 data: dict,
300 call_type: Literal[
301 "completion",
302 "text_completion",
303 "embeddings",
304 "image_generation",
305 "moderation",
306 "audio_transcription",
307 "pass_through_endpoint",
308 "rerank",
309 "mcp_call",
310 "anthropic_messages",
311 ],
312 ) -> Exception | str | dict | None:
313 # Trust proxy call_type, not caller-controlled request shape.
314 is_mcp: Final = self._is_mcp_call_type(call_type)
316 if not self._surface_matches(is_mcp):
317 verbose_proxy_logger.debug(
318 "Cisco AI Defense guardrail: call_type=%s does not match configured inspection_type=%s, skipping",
319 call_type,
320 self.inspection_type,
321 )
322 return data
324 event_type: Final = GuardrailEventHooks.pre_mcp_call if is_mcp else GuardrailEventHooks.pre_call
325 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
326 return data
328 if is_mcp:
329 await self._inspect_mcp_request(data=data, user_api_key_dict=user_api_key_dict)
330 else:
331 messages: Final = self._extract_inspect_messages_from_request(data)
332 if not messages:
333 verbose_proxy_logger.debug(
334 "Cisco AI Defense guardrail: no scannable messages in pre-call request, skipping"
335 )
336 return data
337 await self._inspect_chat(
338 messages=messages,
339 request_data=data,
340 user_api_key_dict=user_api_key_dict,
341 )
343 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
344 return data
346 @log_guardrail_information
347 async def async_moderation_hook(
348 self,
349 data: dict,
350 user_api_key_dict: UserAPIKeyAuth,
351 call_type: Literal[
352 "completion",
353 "embeddings",
354 "image_generation",
355 "moderation",
356 "audio_transcription",
357 "responses",
358 "mcp_call",
359 "anthropic_messages",
360 ],
361 ) -> Exception | str | dict | None:
362 is_mcp: Final = self._is_mcp_call_type(call_type)
364 if not self._surface_matches(is_mcp):
365 return data
367 event_type: Final = GuardrailEventHooks.during_mcp_call if is_mcp else GuardrailEventHooks.during_call
368 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
369 return data
371 if is_mcp:
372 await self._inspect_mcp_request(data=data, user_api_key_dict=user_api_key_dict)
373 else:
374 messages: Final = self._extract_inspect_messages_from_request(data)
375 if not messages:
376 return data
377 await self._inspect_chat(
378 messages=messages,
379 request_data=data,
380 user_api_key_dict=user_api_key_dict,
381 )
383 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
384 return data
386 @log_guardrail_information
387 async def async_post_call_success_hook(
388 self,
389 data: dict,
390 user_api_key_dict: UserAPIKeyAuth,
391 response: LLMResponseTypes,
392 ) -> LLMResponseTypes:
393 if self.inspection_type != "chat":
394 return response
396 if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call) is not True:
397 return response
399 response_messages: Final = self._extract_response_messages(response)
400 if not response_messages:
401 verbose_proxy_logger.debug(
402 "Cisco AI Defense guardrail: no response content to scan, skipping post-call analysis"
403 )
404 return response
406 request_messages: Final = self._extract_inspect_messages_from_request(data)
407 conversation: Final = request_messages + response_messages
409 await self._inspect_chat(
410 messages=conversation,
411 request_data=data,
412 user_api_key_dict=user_api_key_dict,
413 direction="output",
414 response_obj=response,
415 )
417 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
418 return response
420 async def async_post_call_streaming_iterator_hook(
421 self,
422 user_api_key_dict: UserAPIKeyAuth,
423 response: AsyncIterator[object],
424 request_data: dict,
425 ):
426 """Buffer and inspect streaming chat output before delivery."""
427 from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
428 from litellm.main import stream_chunk_builder
430 if self.inspection_type != "chat":
431 async for chunk in response:
432 yield chunk
433 return
435 if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True:
436 async for chunk in response:
437 yield chunk
438 return
440 verbose_proxy_logger.debug(
441 "Cisco AI Defense guardrail (%s): scanning streaming chat response.",
442 self.guardrail_name,
443 )
445 all_chunks: Final[list[object]] = []
446 try:
447 async for chunk in response:
448 all_chunks.append(chunk)
449 except Exception as exc:
450 verbose_proxy_logger.error(
451 "Cisco AI Defense guardrail: upstream streaming failed: %s",
452 exc,
453 )
454 raise
456 if not all_chunks:
457 return
459 if not isinstance(all_chunks[0], (ModelResponse, ModelResponseStream)):
460 verbose_proxy_logger.warning(
461 "Cisco AI Defense guardrail (%s): unsupported streaming chunk shape (%s) — failing closed.",
462 self.guardrail_name,
463 type(all_chunks[0]).__name__,
464 )
465 yield f"data: {json.dumps({'error': {'message': 'Cisco AI Defense: unsupported streaming format — response withheld for safety', 'type': 'guardrail_unsupported_stream', 'code': 400, 'guardrail': self.guardrail_name}})}\n\n"
466 return
468 assembled: Final = stream_chunk_builder(chunks=all_chunks)
469 if assembled is None:
470 for chunk in all_chunks:
471 yield chunk
472 return
473 if not isinstance(assembled, ModelResponse):
474 verbose_proxy_logger.warning(
475 "Cisco AI Defense guardrail (%s): assembled streaming "
476 "response has unsupported shape (%s) — failing closed.",
477 self.guardrail_name,
478 type(assembled).__name__,
479 )
480 yield f"data: {json.dumps({'error': {'message': 'Cisco AI Defense: unsupported streaming format — response withheld for safety', 'type': 'guardrail_unsupported_stream', 'code': 400, 'guardrail': self.guardrail_name}})}\n\n"
481 return
483 response_messages: Final = self._extract_response_messages(assembled)
484 original_stream_text: Final = self._extract_streaming_chunk_scan_text(all_chunks)
485 assembled_text: Final = " ".join(m.get("content", "") for m in response_messages if isinstance(m, dict))
486 if original_stream_text and original_stream_text not in assembled_text:
487 response_messages.append({"role": "assistant", "content": original_stream_text})
488 if not response_messages:
489 for chunk in all_chunks:
490 yield chunk
491 return
493 request_messages: Final = self._extract_inspect_messages_from_request(request_data)
494 conversation: Final = request_messages + response_messages
496 try:
497 await self._inspect_chat(
498 messages=conversation,
499 request_data=request_data,
500 user_api_key_dict=user_api_key_dict,
501 direction="output",
502 response_obj=assembled,
503 )
504 except HTTPException as exc:
505 error_obj: dict[str, object] = self._http_exception_to_error_obj(exc)
506 verbose_proxy_logger.warning(
507 "Cisco AI Defense guardrail (%s): streaming response "
508 "blocked — emitting SSE error event instead of "
509 "delivering buffered chunks.",
510 self.guardrail_name,
511 )
512 yield f"data: {json.dumps({'error': error_obj})}\n\n"
513 return
514 except Exception as exc:
515 verbose_proxy_logger.error(
516 "Cisco AI Defense guardrail (%s): streaming response scan failed: %s",
517 self.guardrail_name,
518 exc,
519 )
520 error_obj = {
521 "message": ("Cisco AI Defense streaming scan failed — response withheld."),
522 "type": "guardrail_scan_error",
523 "code": 500,
524 "guardrail": self.guardrail_name,
525 }
526 yield f"data: {json.dumps({'error': error_obj})}\n\n"
527 return
529 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
531 if self._streaming_content_was_modified(all_chunks, assembled):
532 mock_iterator: Final = MockResponseIterator(model_response=assembled)
533 async for chunk in mock_iterator:
534 yield chunk
535 else:
536 for chunk in all_chunks:
537 yield chunk
539 def _build_block_payload(self, context: _ScanContext, verdict: _CiscoVerdict) -> dict[str, object]:
540 """Canonical block payload used across all four block paths.
542 Same dict is the ``HTTPException.detail`` for chat / MCP request
543 and chat response blocks, the ``error`` value in the streaming
544 SSE event, and (JSON-encoded) the text content of the synthetic
545 MCP response object. Keeps the customer-facing format identical
546 regardless of which transport carries the block.
547 """
548 return {
549 "error": "Blocked by Cisco AI Defense Guardrail",
550 "message": "Blocked by Cisco AI Defense Guardrail",
551 "provider": self._PROVIDER_NAME,
552 "guardrail": self.guardrail_name,
553 "surface": context.surface,
554 "direction": context.direction,
555 "action": "block",
556 "classifications": list(verdict.classifications),
557 "severity": verdict.severity,
558 "rules": [r.get("rule_name") for r in verdict.rules if isinstance(r, dict)],
559 "explanation": verdict.explanation,
560 "event_id": verdict.event_id,
561 }
563 def _http_exception_to_error_obj(self, exc: HTTPException) -> dict[str, object]:
564 """Wrap an ``HTTPException`` detail into the SSE ``error`` payload.
566 For Cisco's own blocks the detail is already the canonical block
567 payload, so this is a near-passthrough that just adds ``code``
568 / ``guardrail`` defaults for non-Cisco / unstructured details.
569 """
570 error_obj: dict[str, object] = {**exc.detail} if isinstance(exc.detail, dict) else {"message": str(exc.detail)}
571 error_obj.setdefault("message", error_obj.get("error", "Guardrail block"))
572 error_obj.setdefault("code", exc.status_code)
573 error_obj.setdefault("guardrail", self.guardrail_name)
574 return error_obj
576 @classmethod
577 def _streaming_content_was_modified(cls, original_chunks: Sequence[object], assembled: ModelResponse) -> bool:
578 """Decide whether redact changed content or tool/function arguments."""
579 original_text: Final = cls._extract_streaming_chunk_scan_text(original_chunks)
580 assembled_text: Final = " ".join(m.get("content", "") for m in cls._extract_response_messages(assembled))
581 return original_text != assembled_text
583 @classmethod
584 def _extract_streaming_chunk_scan_text(cls, chunks: Sequence[object]) -> str:
585 original_text = ""
586 argument_text = ""
587 for chunk in chunks:
588 choices = getattr(chunk, "choices", None) or []
589 for c in choices:
590 delta: object | None = getattr(c, "delta", None)
591 if delta is None:
592 continue
593 text = getattr(delta, "content", None)
594 if isinstance(text, str):
595 original_text += text
596 reasoning_text = " ".join(cls._extract_message_reasoning_parts(delta))
597 if reasoning_text:
598 original_text += reasoning_text
599 for tc in getattr(delta, "tool_calls", None) or []:
600 args = cls._extract_tool_call_arguments(tc)
601 if args:
602 argument_text += args
603 fc: object | None = getattr(delta, "function_call", None)
604 if fc is not None:
605 args = cls._extract_function_call_arguments(fc)
606 if args:
607 argument_text += args
608 return " ".join(part for part in (original_text, argument_text) if part)
610 # ------------------------------------------------------------------
611 # MCP post-tool-call hook lives on ``_CiscoAIDefenseMcpMixin`` in
612 # ``cisco_ai_defense_mcp.py``. The mixin's methods are inherited via
613 # the class declaration above (multiple-inheritance with
614 # ``_CiscoAIDefenseMcpMixin`` placed first).
615 # ------------------------------------------------------------------
617 def _surface_matches(self, is_mcp_traffic: bool) -> bool:
618 """Return True when the traffic surface matches the configured type."""
619 if self.inspection_type == "mcp":
620 return is_mcp_traffic
621 return not is_mcp_traffic
623 @staticmethod
624 def _normalize_event_hooks(event_hook: object) -> set:
625 """Coerce a ``mode`` arg (str, enum, or list of either) to a set of values."""
627 def _norm(hook: object) -> str | None:
628 value: Final = getattr(hook, "value", None)
629 if isinstance(value, str):
630 return value
631 if isinstance(hook, str):
632 return hook
633 return None
635 if event_hook is None:
636 return set()
637 if isinstance(event_hook, list):
638 values = {_norm(h) for h in event_hook}
639 else:
640 values = {_norm(event_hook)}
641 values.discard(None)
642 return values
644 @staticmethod
645 def _infer_inspection_type_from_mode(event_hook: object, current: str) -> str:
646 """Return ``mcp`` when ``event_hook`` is exclusively MCP-typed.
648 ``pre_mcp_call`` and ``during_mcp_call`` only fire for MCP traffic,
649 so a user who picks them clearly wants MCP inspection — auto-flip
650 the surface so they don't also have to toggle ``inspection_type``.
651 """
652 configured: Final = CiscoAIDefenseGuardrail._normalize_event_hooks(event_hook)
653 if not configured:
654 return current
655 mcp_hooks: Final = {"pre_mcp_call", "during_mcp_call"}
656 chat_hooks: Final = {"pre_call", "during_call", "post_call"}
657 has_mcp: Final = bool(configured & mcp_hooks)
658 has_chat: Final = bool(configured & chat_hooks)
659 # Exclusively MCP → mcp; exclusively chat → chat; mixed → keep
660 # current so the user retains control over the dual-surface case.
661 if has_mcp and not has_chat:
662 return "mcp"
663 if has_chat and not has_mcp:
664 return "chat"
665 return current
667 def _log_decision(
668 self,
669 context: _ScanContext,
670 verdict: _CiscoVerdict,
671 duration_ms: float,
672 request_data: dict,
673 ) -> None:
674 """Emit a single visible log line per scan.
676 Mirrors the reference plugin's ``AI_DEFENSE_DECISION`` line so
677 operators can observe scans without bumping log levels. INFO for
678 allow, WARNING for intervened/redacted, ERROR is left for
679 upstream API failures.
680 """
681 fields: Final[dict[str, object]] = {
682 "guardrail": self.guardrail_name,
683 "surface": context.surface,
684 "direction": context.direction,
685 "action": verdict.action,
686 "is_safe": verdict.is_safe,
687 "severity": verdict.severity,
688 "classifications": (list(verdict.classifications) if verdict.classifications else []),
689 "rule_violations": sorted(
690 {
691 rule.get("rule_name")
692 for rule in verdict.rules
693 if isinstance(rule, dict)
694 and rule.get("rule_name")
695 and rule.get("classification") not in (None, "NONE_VIOLATION")
696 }
697 ),
698 "event_id": verdict.event_id,
699 "duration_ms": round(duration_ms, 1),
700 }
701 # Best-effort request context — useful when correlating with model
702 # / MCP-tool calls. None values are dropped for log-line brevity.
703 for source_key, target_key in (
704 ("model", "model"),
705 ("litellm_call_id", "call_id"),
706 ("mcp_tool_name", "mcp_tool"),
707 ("mcp_server_name", "mcp_server"),
708 ):
709 value = request_data.get(source_key)
710 if value:
711 fields[target_key] = value
713 payload: Final = {k: v for k, v in fields.items() if v not in (None, [], "")}
714 line = "CISCO_AI_DEFENSE_DECISION " + json.dumps(payload, default=str, sort_keys=True, separators=(",", ":"))
716 if verdict.action == _ACTION_ALLOW:
717 verbose_proxy_logger.info(line)
718 else:
719 verbose_proxy_logger.warning(line)
721 def _warn_if_mode_surface_mismatch(self, event_hook: object) -> None:
722 """Log a warning only when ``mode`` mixes both surfaces.
724 Auto-inference in ``_infer_inspection_type_from_mode`` handles the
725 "exclusively MCP" and "exclusively chat" cases, so this warning
726 fires only for genuinely mixed configurations where we can't tell
727 which surface the user wants and have to honour their explicit
728 ``inspection_type``.
729 """
730 configured: Final = self._normalize_event_hooks(event_hook)
731 mcp_hooks: Final = configured & {"pre_mcp_call", "during_mcp_call"}
732 chat_hooks: Final = configured & {"pre_call", "during_call", "post_call"}
733 if not (mcp_hooks and chat_hooks):
734 return
736 unused_hooks: Final = mcp_hooks if self.inspection_type == "chat" else chat_hooks
737 verbose_proxy_logger.warning(
738 "Cisco AI Defense guardrail '%s' (inspection_type=%s) has mixed "
739 "mode %s — the %s event hooks won't fire because this guardrail "
740 "only inspects %s traffic. Configure two guardrails (one per "
741 "surface) for full coverage, or drop the cross-surface modes.",
742 self.guardrail_name,
743 self.inspection_type,
744 sorted(configured),
745 sorted(unused_hooks),
746 self.inspection_type,
747 )
749 # ------------------------------------------------------------------
750 # Chat inspection
751 # ------------------------------------------------------------------
753 async def _inspect_chat(
754 self,
755 messages: list[dict[str, str]],
756 request_data: dict,
757 user_api_key_dict: UserAPIKeyAuth,
758 direction: str = "input",
759 response_obj: object = None,
760 ) -> dict[str, object]:
761 url: Final = f"{self.api_base}{self.inspect_path}"
762 payload: Final = self._build_chat_payload(messages, request_data, user_api_key_dict)
763 start_time: Final = datetime.now()
764 try:
765 inspect_response: Final = await self._post_inspection(url=url, payload=payload, surface="chat")
766 except HTTPException:
767 # Re-raise; _post_inspection only raises CiscoAIDefenseGuardrailAPIError,
768 # but be defensive in case downstream evolves.
769 raise
770 except Exception as exc:
771 return self._handle_api_error(
772 exc,
773 request_data=request_data,
774 start_time=start_time,
775 surface="chat",
776 direction=direction,
777 )
779 return self._finalize_inspection(
780 inspect_response=inspect_response,
781 request_data=request_data,
782 context=_ScanContext(surface="chat", direction=direction),
783 start_time=start_time,
784 response_obj=response_obj,
785 )
787 def _build_chat_payload(
788 self,
789 messages: list[dict[str, str]],
790 request_data: dict,
791 user_api_key_dict: UserAPIKeyAuth,
792 ) -> dict[str, object]:
793 return {
794 "messages": messages,
795 "metadata": self._build_metadata(request_data, user_api_key_dict),
796 "config": self._build_config(),
797 }
799 # ------------------------------------------------------------------
800 # Shared HTTP / metadata helpers
801 # ------------------------------------------------------------------
803 async def _post_inspection(
804 self,
805 url: str,
806 payload: dict[str, object],
807 surface: str,
808 ) -> dict[str, object]:
809 headers: Final = self._build_headers()
810 verbose_proxy_logger.debug(
811 "Cisco AI Defense guardrail: posting %s inspection to %s",
812 surface,
813 url,
814 )
815 try:
816 request: Final = self.async_handler.client.build_request(
817 "POST",
818 url,
819 headers=headers,
820 json=payload,
821 timeout=self.timeout,
822 )
823 response: Final = await self.async_handler.client.send(
824 request,
825 follow_redirects=False,
826 )
827 response.raise_for_status()
828 except httpx.HTTPStatusError as exc:
829 status_code: Final = exc.response.status_code if exc.response is not None else 0
830 body_snippet = ""
831 try:
832 body_snippet = exc.response.text[:500] if exc.response else ""
833 except Exception:
834 body_snippet = ""
835 raise CiscoAIDefenseGuardrailAPIError(
836 f"Cisco AI Defense {surface} API returned HTTP {status_code}: {body_snippet}"
837 ) from exc
838 except httpx.TimeoutException as exc:
839 raise CiscoAIDefenseGuardrailAPIError(
840 f"Cisco AI Defense {surface} API call timed out after {self.timeout}s"
841 ) from exc
842 except httpx.RequestError as exc:
843 raise CiscoAIDefenseGuardrailAPIError(f"Cisco AI Defense {surface} API request failed: {exc}") from exc
845 try:
846 return response.json()
847 except ValueError as exc:
848 raise CiscoAIDefenseGuardrailAPIError(
849 f"Cisco AI Defense {surface} API returned a non-JSON response"
850 ) from exc
852 def _build_headers(self) -> dict[str, str]:
853 return {
854 CISCO_API_KEY_HEADER: self.api_key,
855 "Content-Type": "application/json",
856 "Accept": "application/json",
857 "User-Agent": f"litellm/{litellm_version}",
858 }
860 def _build_metadata(
861 self,
862 request_data: dict,
863 user_api_key_dict: UserAPIKeyAuth,
864 ) -> dict[str, object]:
865 metadata: Final[dict[str, object]] = {}
867 user: Final = request_data.get("user") or getattr(user_api_key_dict, "user_id", None)
868 if user:
869 metadata["user"] = str(user)
871 litellm_call_id: Final = request_data.get("litellm_call_id")
872 if litellm_call_id:
873 metadata["client_transaction_id"] = str(litellm_call_id)
875 request_metadata: Final = request_data.get("metadata") or {}
876 if isinstance(request_metadata, dict):
877 for src_key in (
878 "src_app",
879 "dst_app",
880 "src_ip",
881 "dst_ip",
882 "dst_host",
883 "sni",
884 "user_agent",
885 ):
886 value = request_metadata.get(src_key)
887 if value:
888 metadata[src_key] = str(value)
890 return metadata
892 def _build_config(self) -> dict[str, object]:
893 config: Final[dict[str, object]] = {}
894 if self.enabled_rules:
895 config["enabled_rules"] = self.enabled_rules
896 if self.integration_profile_id:
897 config["integration_profile_id"] = self.integration_profile_id
898 if self.integration_profile_version:
899 config["integration_profile_version"] = self.integration_profile_version
900 if self.integration_tenant_id:
901 config["integration_tenant_id"] = self.integration_tenant_id
902 if self.integration_type:
903 config["integration_type"] = self.integration_type
904 return config
906 @staticmethod
907 def _normalize_rule(rule: object) -> dict[str, object]:
908 """Coerce a user-supplied rule into the wire-shape dict Cisco expects.
910 Accepts ``str``, ``dict``, and Pydantic model inputs.
911 """
912 if isinstance(rule, str):
913 return {"rule_name": rule}
915 if not isinstance(rule, dict):
916 # Pydantic BaseModel (CiscoAIDefenseRule and friends): dump
917 # to a dict and re-enter the dict branch. Anything else
918 # falls through to the explicit raise so misconfig still
919 # surfaces clearly at startup instead of mid-request.
920 model_dump: Final = getattr(rule, "model_dump", None)
921 if callable(model_dump):
922 try:
923 dumped = model_dump(exclude_none=True)
924 except TypeError:
925 dumped = model_dump()
926 if isinstance(dumped, dict):
927 rule = dumped
929 if isinstance(rule, dict):
930 normalized: Final[dict[str, object]] = {}
931 rule_name: Final = rule.get("rule_name")
932 if rule_name:
933 normalized["rule_name"] = rule_name
934 entity_types: Final = rule.get("entity_types")
935 if entity_types:
936 normalized["entity_types"] = list(entity_types)
937 rule_id: Final = rule.get("rule_id")
938 if rule_id is not None:
939 normalized["rule_id"] = rule_id
940 classification: Final = rule.get("classification")
941 if classification:
942 normalized["classification"] = classification
943 return normalized
945 raise ValueError(f"Cisco AI Defense guardrail: invalid rule definition: {rule!r}")
947 # ------------------------------------------------------------------
948 # Response processing
949 # ------------------------------------------------------------------
951 def _finalize_inspection(
952 self,
953 inspect_response: dict[str, Any],
954 request_data: dict,
955 context: _ScanContext,
956 start_time: datetime,
957 response_obj: object = None,
958 ) -> dict[str, object]:
959 """Parse, log, and (optionally) raise/redact on the Cisco verdict.
961 ``context.direction`` is ``"input"`` for request scans and ``"output"``
962 for response scans (used for metadata namespacing and response headers).
963 ``response_obj`` is the LiteLLM response object (or MCP tool-call
964 response) used when applying a ``redact`` action to outputs.
966 Cisco AI Defense returns two different envelope shapes depending on
967 the endpoint:
969 * ``/api/v1/inspect/chat`` — top-level verdict
970 ``{"is_safe": ..., "classifications": [...], "action": ..., ...}``
971 * ``/api/v1/inspect/mcp`` — JSON-RPC wrapper
972 ``{"jsonrpc": "2.0", "id": ..., "result": {<same verdict>}}``
974 We unwrap the JSON-RPC ``result`` so both endpoints feed the same
975 downstream code path. The error envelope detection below already
976 handles ``error`` at either level.
977 """
978 # Surface JSON-RPC error envelopes (HTTP 200 + Cisco-side error) the
979 # same way as transport errors: fail-open or fail-closed.
980 jsonrpc_error: Final = self._extract_jsonrpc_error(inspect_response)
981 if jsonrpc_error is not None:
982 verbose_proxy_logger.warning(
983 "Cisco AI Defense guardrail: API returned JSON-RPC error envelope (code=%s message=%s)",
984 jsonrpc_error.get("code"),
985 jsonrpc_error.get("message"),
986 )
987 return self._handle_api_error(
988 CiscoAIDefenseGuardrailAPIError(
989 f"AI Defense error code={jsonrpc_error.get('code')} message={jsonrpc_error.get('message')}"
990 ),
991 request_data=request_data,
992 start_time=start_time,
993 surface=context.surface,
994 direction=context.direction,
995 )
997 # Unwrap the JSON-RPC ``result`` envelope used by the MCP inspect
998 # endpoint. The chat endpoint returns the verdict at the top
999 # level and isn't wrapped, so this is a no-op there.
1000 verdict_dict: Final = self._unwrap_verdict_envelope(inspect_response)
1002 # OpenAPI spec lists `classification` as required (singular) but
1003 # examples & SDK return `classifications` (plural). Accept both.
1004 classifications: Final = (
1005 verdict_dict.get("classifications")
1006 or ([verdict_dict["classification"]] if verdict_dict.get("classification") else [])
1007 or []
1008 )
1009 verdict = _CiscoVerdict(
1010 is_safe=verdict_dict.get("is_safe"),
1011 classifications=classifications,
1012 severity=verdict_dict.get("severity"),
1013 rules=verdict_dict.get("rules") or [],
1014 explanation=verdict_dict.get("explanation"),
1015 event_id=verdict_dict.get("event_id"),
1016 sanitized_text=self._extract_sanitized_text(verdict_dict),
1017 sanitized_messages=self._extract_sanitized_messages(verdict_dict),
1018 sanitized_mcp_arguments=self._extract_sanitized_mcp_arguments(verdict_dict),
1019 )
1021 action_raw: Final = verdict_dict.get("action")
1022 if isinstance(action_raw, str) and action_raw.strip():
1023 action = self._normalize_action(action_raw)
1024 else:
1025 action = _ACTION_ALLOW
1026 verdict = replace(verdict, action=action)
1028 end_time: Final = datetime.now()
1029 duration: Final = (end_time - start_time).total_seconds()
1031 if context.surface == "mcp":
1032 logging_event_type = (
1033 GuardrailEventHooks.during_mcp_call
1034 if context.direction == "output"
1035 else GuardrailEventHooks.pre_mcp_call
1036 )
1037 else:
1038 logging_event_type = (
1039 GuardrailEventHooks.post_call if context.direction == "output" else GuardrailEventHooks.pre_call
1040 )
1042 self.add_standard_logging_guardrail_information_to_request_data(
1043 guardrail_provider=self._PROVIDER_NAME,
1044 guardrail_json_response=self._sanitize_response_for_logging(
1045 inspect_response, surface=context.surface, action=action
1046 ),
1047 request_data=request_data,
1048 guardrail_status=("guardrail_intervened" if action in (_ACTION_BLOCK, _ACTION_REDACT) else "success"),
1049 start_time=start_time.timestamp(),
1050 end_time=end_time.timestamp(),
1051 duration=duration,
1052 masked_entity_count=self._extract_masked_entity_count(verdict.rules),
1053 event_type=logging_event_type,
1054 )
1056 self._stash_verdict_on_request(request_data, context, verdict)
1058 self._log_decision(context, verdict, duration * 1000, request_data)
1060 if action == _ACTION_ALLOW:
1061 return inspect_response
1063 if action == _ACTION_REDACT:
1064 redacted: Final = self._apply_redaction(request_data, response_obj, context, verdict)
1065 if redacted:
1066 verbose_proxy_logger.info(
1067 "Cisco AI Defense guardrail (%s): redaction applied (event_id=%s)",
1068 context.surface,
1069 verdict.event_id,
1070 )
1071 return inspect_response
1072 verbose_proxy_logger.warning(
1073 "Cisco AI Defense guardrail (%s): redact requested but no "
1074 "rewritable surface found — falling through to "
1075 "on_flagged_action=%s",
1076 context.surface,
1077 self.on_flagged_action,
1078 )
1080 if self.on_flagged_action == "block":
1081 raise HTTPException(
1082 status_code=400,
1083 detail=self._build_block_payload(context, verdict),
1084 )
1086 verbose_proxy_logger.info(
1087 "Cisco AI Defense guardrail (%s): violation in monitor mode — request allowed to proceed (event_id=%s)",
1088 context.surface,
1089 verdict.event_id,
1090 )
1091 return inspect_response
1093 @staticmethod
1094 def _stash_verdict_on_request(request_data: dict, context: _ScanContext, verdict: _CiscoVerdict) -> None:
1095 """Surface the Cisco verdict on the request metadata for observability."""
1096 metadata_store: Final = request_data.setdefault("metadata", {})
1097 if not isinstance(metadata_store, dict):
1098 return
1099 prefix: Final = f"cisco_ai_defense_{context.surface}_{context.direction}"
1100 metadata_store[f"{prefix}_is_safe"] = verdict.is_safe
1101 if verdict.action:
1102 metadata_store[f"{prefix}_action"] = verdict.action
1103 if verdict.classifications:
1104 metadata_store[f"{prefix}_classifications"] = list(verdict.classifications)
1105 if verdict.severity:
1106 metadata_store[f"{prefix}_severity"] = verdict.severity
1107 if verdict.rules:
1108 metadata_store[f"{prefix}_rules"] = [
1109 rule.get("rule_name") for rule in verdict.rules if isinstance(rule, dict)
1110 ]
1111 if verdict.event_id:
1112 metadata_store[f"{prefix}_event_id"] = verdict.event_id
1114 _REDACTED_LOG_KEYS = frozenset(
1115 {
1116 "raw_request",
1117 "sanitized_payload",
1118 "sanitizedPayload",
1119 "modified_payload",
1120 "modifiedPayload",
1121 }
1122 )
1124 @classmethod
1125 def _sanitize_response_for_logging(
1126 cls,
1127 inspect_response: Mapping[str, object],
1128 surface: str,
1129 action: str | None = None,
1130 ) -> dict[str, object]:
1131 """Drop bulky / privacy-sensitive fields, recursing into nested dicts.
1133 MCP verdicts are commonly nested under ``result``, so a
1134 top-level-only strip would leave ``result.raw_request`` or
1135 ``result.sanitized_payload`` in the logging metadata.
1136 """
1137 if not isinstance(inspect_response, dict):
1138 return {"surface": surface, **({"action": action} if action else {})}
1139 sanitized: Final = cls._strip_sensitive_keys(inspect_response)
1140 sanitized["surface"] = surface
1141 if action:
1142 sanitized["action"] = action
1143 return sanitized
1145 @classmethod
1146 def _strip_sensitive_keys(cls, d: Mapping[str, object]) -> dict[str, object]:
1147 """Recursively strip privacy-sensitive keys from a verdict dict."""
1148 out: Final[dict[str, object]] = {}
1149 for key, value in d.items():
1150 if key.startswith("_") or key in cls._REDACTED_LOG_KEYS:
1151 continue
1152 if isinstance(value, dict):
1153 out[key] = cls._strip_sensitive_keys(value)
1154 else:
1155 out[key] = value
1156 return out
1158 # ------------------------------------------------------------------
1159 # Verdict extraction helpers (sanitized content + JSON-RPC errors)
1160 # ------------------------------------------------------------------
1162 _DECISION_FIELDS: tuple[str, ...] = (
1163 "action",
1164 "allowed",
1165 "blocked",
1166 "safe",
1167 "is_safe",
1168 "decision",
1169 "verdict",
1170 "status",
1171 "score",
1172 "risk_score",
1173 "confidence",
1174 "categories",
1175 "classifications",
1176 "violations",
1177 "threats",
1178 "policies",
1179 "reason",
1180 "rules",
1181 "sanitized_text",
1182 "sanitizedText",
1183 "sanitized_payload",
1184 )
1186 @classmethod
1187 def _has_decision_fields(cls, payload: object) -> bool:
1188 if not isinstance(payload, dict):
1189 return False
1190 return any(key in payload for key in cls._DECISION_FIELDS)
1192 @classmethod
1193 def _unwrap_verdict_envelope(cls, inspect_response: dict[str, Any]) -> dict[str, Any]:
1194 """Return the dict that actually holds is_safe / action / rules.
1196 Cisco AI Defense returns the verdict at different nesting depths
1197 depending on the endpoint and SDK version:
1199 * ``/api/v1/inspect/chat`` — verdict is at the top level.
1200 * ``/api/v1/inspect/mcp`` — JSON-RPC envelope wraps the verdict
1201 under ``result``.
1202 * Some SDKs nest under ``data`` / ``inspection`` / ``ai_defense``.
1204 Mirrors the reference plugin's ``_decision_payload`` so the
1205 handler tolerates every shape Cisco's own tested integration
1206 already supports.
1207 """
1208 if not isinstance(inspect_response, dict):
1209 return {}
1211 if cls._has_decision_fields(inspect_response):
1212 return inspect_response
1214 for key in ("result", "data", "inspection", "ai_defense", "aiDefense"):
1215 value = inspect_response.get(key)
1216 if cls._has_decision_fields(value):
1217 return value
1219 result: Final = inspect_response.get("result")
1220 if isinstance(result, dict):
1221 for key in ("data", "inspection", "ai_defense", "aiDefense"):
1222 value = result.get(key)
1223 if cls._has_decision_fields(value):
1224 return value
1226 return inspect_response
1228 @staticmethod
1229 def _extract_jsonrpc_error(
1230 inspect_response: Mapping[str, object],
1231 ) -> dict[str, object] | None:
1232 """Detect a JSON-RPC error envelope inside an HTTP 200 response.
1234 The Cisco Inspect API can return ``{"error": {...}}`` (or nest one
1235 under ``"result"``) inside a 200. We treat that the same as a
1236 transport error so the configured ``fallback_on_error`` policy
1237 applies.
1238 """
1239 if not isinstance(inspect_response, dict):
1240 return None
1241 error: Final = inspect_response.get("error")
1242 if isinstance(error, dict):
1243 return error
1244 result: Final = inspect_response.get("result")
1245 if isinstance(result, dict):
1246 inner: Final = result.get("error")
1247 if isinstance(inner, dict):
1248 return inner
1249 return None
1251 @staticmethod
1252 def _normalize_action(raw_action: str) -> str:
1253 """Map Cisco/reference-plugin action vocabulary to ours."""
1254 normalized: Final = raw_action.strip().lower()
1255 if normalized in {
1256 "deny",
1257 "denied",
1258 "block",
1259 "blocked",
1260 "reject",
1261 "rejected",
1262 "unsafe",
1263 "malicious",
1264 }:
1265 return _ACTION_BLOCK
1266 if normalized in {"redact", "redacted", "sanitize", "sanitized", "mask"}:
1267 return _ACTION_REDACT
1268 if normalized in {"allow", "allowed", "safe", "ok"}:
1269 return _ACTION_ALLOW
1270 verbose_proxy_logger.warning(
1271 "Cisco AI Defense guardrail: unrecognized action %r treated as block",
1272 raw_action,
1273 )
1274 return _ACTION_BLOCK
1276 @staticmethod
1277 def _extract_sanitized_text(
1278 inspect_response: Mapping[str, object],
1279 ) -> str | None:
1280 """Pull ``sanitized_text`` (or camelCase variant) off the verdict."""
1281 for key in ("sanitized_text", "sanitizedText"):
1282 value = inspect_response.get(key)
1283 if isinstance(value, str) and value:
1284 return value
1285 result: Final = inspect_response.get("result")
1286 if isinstance(result, dict):
1287 for key in ("sanitized_text", "sanitizedText"):
1288 value = result.get(key)
1289 if isinstance(value, str) and value:
1290 return value
1291 return None
1293 @staticmethod
1294 def _extract_sanitized_messages(
1295 inspect_response: Mapping[str, object],
1296 ) -> list[dict[str, object]] | None:
1297 """Pull a sanitized OpenAI-format messages array off the verdict.
1299 Cisco can return the rewrite under several keys; we accept any of
1300 the common variants and stop at the first non-empty match.
1301 """
1302 containers: Final = [inspect_response]
1303 for container_key in ("result", "data"):
1304 container = inspect_response.get(container_key)
1305 if isinstance(container, dict):
1306 containers.append(container)
1308 for container in containers:
1309 for key in (
1310 "sanitized_messages",
1311 "sanitizedMessages",
1312 "modified_messages",
1313 "modifiedMessages",
1314 ):
1315 value = container.get(key)
1316 if isinstance(value, list) and value:
1317 return [m for m in value if isinstance(m, dict)]
1318 for key in (
1319 "sanitized_payload",
1320 "sanitizedPayload",
1321 "modified_payload",
1322 "modifiedPayload",
1323 ):
1324 payload = container.get(key)
1325 if isinstance(payload, dict):
1326 messages = payload.get("messages")
1327 if isinstance(messages, list) and messages:
1328 return [m for m in messages if isinstance(m, dict)]
1329 return None
1331 def _apply_redaction(
1332 self,
1333 request_data: dict,
1334 response_obj: object,
1335 context: _ScanContext,
1336 verdict: _CiscoVerdict,
1337 ) -> bool:
1338 """Apply a Cisco-supplied rewrite to the request/response in place.
1340 Returns True when a rewrite was applied; False when there was no
1341 suitable surface to rewrite (caller then falls back to
1342 ``on_flagged_action``).
1343 """
1344 if context.surface == "mcp" and context.direction == "input":
1345 return self._redact_mcp_input(request_data, verdict.sanitized_text, verdict.sanitized_mcp_arguments)
1346 if context.surface == "mcp" and context.direction == "output":
1347 if response_obj is None:
1348 return False
1349 if verdict.sanitized_text:
1350 return self._set_mcp_tool_response_text(response_obj, verdict.sanitized_text)
1351 return False
1352 if context.surface == "chat" and context.direction == "input":
1353 return self._redact_chat_input(request_data, verdict.sanitized_text, verdict.sanitized_messages)
1354 if context.surface == "chat" and context.direction == "output":
1355 return self._redact_chat_output(response_obj, verdict.sanitized_text, verdict.sanitized_messages)
1356 return False
1358 @staticmethod
1359 def _redact_mcp_input(
1360 request_data: dict,
1361 sanitized_text: str | None,
1362 sanitized_mcp_arguments: dict[str, object] | None,
1363 ) -> bool:
1364 """Rewrite MCP request arguments in all locations the proxy reads."""
1365 if sanitized_mcp_arguments is not None:
1366 request_data["mcp_arguments"] = sanitized_mcp_arguments
1367 request_data["modified_arguments"] = sanitized_mcp_arguments
1368 params: Final = request_data.get("params")
1369 if isinstance(params, dict):
1370 params["arguments"] = sanitized_mcp_arguments
1371 if isinstance(request_data.get("arguments"), dict):
1372 request_data["arguments"] = sanitized_mcp_arguments
1373 return True
1374 if sanitized_text:
1375 applied = False
1376 for args_path in (
1377 request_data.get("mcp_arguments"),
1378 request_data.get("arguments"),
1379 (request_data.get("params") or {}).get("arguments"),
1380 ):
1381 if not isinstance(args_path, dict):
1382 continue
1383 string_keys = [key for key, value in args_path.items() if isinstance(value, str)]
1384 if len(string_keys) != 1:
1385 continue
1386 args_path[string_keys[0]] = sanitized_text
1387 request_data["modified_arguments"] = args_path
1388 applied = True
1389 return applied
1390 return False
1392 def _redact_chat_input(
1393 self,
1394 request_data: dict,
1395 sanitized_text: str | None,
1396 sanitized_messages: list[dict[str, object]] | None,
1397 ) -> bool:
1398 """Rewrite chat request input (``messages`` or ``input``)."""
1399 if sanitized_messages and self._extract_tool_definition_text(request_data):
1400 # We append one synthetic message carrying the tool/function
1401 # definitions for inspection; Cisco echoes it back in
1402 # ``sanitized_messages``, but it maps to no structured request
1403 # field, so drop it before rewriting the real conversation.
1404 sanitized_messages = sanitized_messages[:-1] or None
1405 uses_input: Final = "input" in request_data and "messages" not in request_data
1406 has_instructions: Final = request_data.get("instructions") is not None
1407 instructions_redacted = False
1408 if has_instructions:
1409 instructions_redacted = self._redact_responses_instructions(
1410 request_data, sanitized_text, sanitized_messages
1411 )
1412 sanitized_messages = self._non_instruction_messages(sanitized_messages)
1413 if not sanitized_messages:
1414 return instructions_redacted
1415 if sanitized_messages:
1416 if uses_input:
1417 rewritten: Final = self._sanitized_messages_to_responses_input(sanitized_messages)
1418 if rewritten is not None:
1419 request_data["input"] = rewritten
1420 return True
1421 return False
1422 request_data["messages"] = sanitized_messages
1423 return True
1424 if sanitized_text:
1425 if uses_input:
1426 rewritten_input: Final = self._rewrite_responses_input_text(request_data.get("input"), sanitized_text)
1427 if rewritten_input is not None:
1428 request_data["input"] = rewritten_input
1429 return True
1430 return False
1431 redacted_arguments: Final = self._clear_chat_input_tool_arguments(request_data)
1432 messages: Final = request_data.get("messages")
1433 redacted_content = False
1434 if isinstance(messages, list) and messages:
1435 for message in reversed(messages):
1436 if (
1437 isinstance(message, dict)
1438 and message.get("role") == "user"
1439 and isinstance(message.get("content"), str)
1440 ):
1441 message["content"] = sanitized_text
1442 redacted_content = True
1443 break
1444 return redacted_content or redacted_arguments
1445 return False
1447 @classmethod
1448 def _redact_responses_instructions(
1449 cls,
1450 request_data: dict,
1451 sanitized_text: str | None,
1452 sanitized_messages: list[dict[str, object]] | None,
1453 ) -> bool:
1454 if sanitized_messages:
1455 instruction_text: Final = cls._instruction_text_from_messages(sanitized_messages)
1456 if instruction_text:
1457 request_data["instructions"] = instruction_text
1458 return True
1459 if sanitized_text and not any(key in request_data for key in ("input", "messages", "prompt")):
1460 request_data["instructions"] = sanitized_text
1461 return True
1462 return False
1464 @classmethod
1465 def _instruction_text_from_messages(cls, messages: list[dict[str, object]]) -> str | None:
1466 for message in messages:
1467 if not isinstance(message, dict):
1468 continue
1469 if cls._is_instruction_role(message.get("role")):
1470 text = cls._normalize_message_content(message.get("content"))
1471 if text:
1472 return text
1473 return None
1475 @classmethod
1476 def _non_instruction_messages(cls, messages: list[dict[str, object]] | None) -> list[dict[str, object]] | None:
1477 if messages is None:
1478 return None
1479 return [
1480 message
1481 for message in messages
1482 if not (isinstance(message, dict) and cls._is_instruction_role(message.get("role")))
1483 ]
1485 @staticmethod
1486 def _is_instruction_role(role: object) -> bool:
1487 return isinstance(role, str) and role.lower() in {"system", "developer"}
1489 @classmethod
1490 def _clear_chat_input_tool_arguments(cls, request_data: dict) -> bool:
1491 messages: Final = request_data.get("messages")
1492 if not isinstance(messages, list):
1493 return False
1494 applied = False
1495 for message in messages:
1496 if not isinstance(message, dict):
1497 continue
1498 if cls._extract_message_tool_argument_parts(message):
1499 cls._clear_tool_call_arguments(message)
1500 applied = True
1501 return applied
1503 def _redact_chat_output(
1504 self,
1505 response_obj: object,
1506 sanitized_text: str | None,
1507 sanitized_messages: list[dict[str, object]] | None,
1508 ) -> bool:
1509 """Rewrite chat response (``ModelResponse`` or ``ResponsesAPIResponse``)."""
1510 if response_obj is None:
1511 return False
1513 if isinstance(response_obj, TextCompletionResponse):
1514 return self._redact_text_completion_choices(
1515 getattr(response_obj, "choices", None) or [],
1516 sanitized_text,
1517 sanitized_messages,
1518 )
1520 choices: Final = getattr(response_obj, "choices", None)
1521 if isinstance(choices, list):
1522 return self._redact_model_response_choices(choices, sanitized_text, sanitized_messages)
1524 output_items: Final = getattr(response_obj, "output", None)
1525 if isinstance(output_items, list):
1526 return self._redact_responses_api_output(output_items, sanitized_text, sanitized_messages)
1528 return False
1530 @staticmethod
1531 def _redact_model_response_choices(
1532 choices: list,
1533 sanitized_text: str | None,
1534 sanitized_messages: list[dict[str, object]] | None,
1535 ) -> bool:
1536 """Redact every returned choice, including tool-call/reasoning fields."""
1537 if sanitized_messages:
1538 applied = False
1539 msg_iter: Final = iter(sanitized_messages)
1540 for choice in choices:
1541 if not isinstance(choice, Choices):
1542 continue
1543 replacement = next(msg_iter, None)
1544 replacement_text = sanitized_text or "[REDACTED]"
1545 if replacement is not None:
1546 text = CiscoAIDefenseGuardrail._normalize_message_content(replacement.get("content"))
1547 if text:
1548 replacement_text = text
1549 choice.message.content = text
1550 applied = True
1551 else:
1552 if getattr(choice.message, "content", None):
1553 choice.message.content = replacement_text
1554 applied = True
1555 if CiscoAIDefenseGuardrail._redact_message_reasoning_fields(choice.message, replacement_text):
1556 applied = True
1557 CiscoAIDefenseGuardrail._clear_tool_call_arguments(choice.message)
1558 return applied
1559 if sanitized_text:
1560 applied = False
1561 for choice in choices:
1562 if not isinstance(choice, Choices):
1563 continue
1564 msg = choice.message
1565 if getattr(msg, "content", None):
1566 msg.content = sanitized_text
1567 applied = True
1568 if CiscoAIDefenseGuardrail._redact_message_reasoning_fields(msg, sanitized_text):
1569 applied = True
1570 CiscoAIDefenseGuardrail._clear_tool_call_arguments(msg)
1571 return applied
1572 return False
1574 @staticmethod
1575 def _redact_text_completion_choices(
1576 choices: list,
1577 sanitized_text: str | None,
1578 sanitized_messages: list[dict[str, object]] | None,
1579 ) -> bool:
1580 """Rewrite ``/v1/completions`` text choices after Cisco redaction."""
1581 replacement = sanitized_text
1582 if not replacement and sanitized_messages:
1583 for message in sanitized_messages:
1584 if not isinstance(message, dict):
1585 continue
1586 text = CiscoAIDefenseGuardrail._normalize_message_content(message.get("content"))
1587 if text:
1588 replacement = text
1589 break
1590 if not replacement:
1591 return False
1592 applied = False
1593 for choice in choices:
1594 if getattr(choice, "text", None):
1595 choice.text = replacement
1596 applied = True
1597 return applied
1599 @classmethod
1600 def _redact_message_reasoning_fields(cls, message: object, replacement_text: str) -> bool:
1601 """Remove preserved reasoning fields and expose the sanitized text."""
1602 if not cls._extract_message_reasoning_parts(message):
1603 return False
1604 setattr(message, "content", replacement_text)
1605 for key in ("reasoning_content", "thinking_blocks", "reasoning_items"):
1606 if not hasattr(message, key):
1607 continue
1608 try:
1609 delattr(message, key)
1610 except (AttributeError, TypeError, ValueError):
1611 try:
1612 setattr(message, key, None)
1613 except (AttributeError, TypeError, ValueError):
1614 pass
1615 return True
1617 @staticmethod
1618 def _clear_arguments_field(obj: object) -> None:
1619 """Set ``obj.arguments`` (or ``obj["arguments"]``) to ``"{}"``."""
1620 if obj is None:
1621 return
1622 if isinstance(obj, dict):
1623 obj["arguments"] = "{}"
1624 return
1625 try:
1626 setattr(obj, "arguments", "{}")
1627 except (AttributeError, TypeError, ValueError):
1628 pass
1630 @classmethod
1631 def _clear_tool_call_arguments(cls, message: object) -> None:
1632 """Clear tool-call / function-call arguments after Cisco redaction."""
1633 tool_calls = message.get("tool_calls") if isinstance(message, dict) else getattr(message, "tool_calls", None)
1634 for tc in tool_calls or []:
1635 fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None)
1636 cls._clear_arguments_field(fn)
1637 function_call: Final = (
1638 message.get("function_call") if isinstance(message, dict) else getattr(message, "function_call", None)
1639 )
1640 cls._clear_arguments_field(function_call)
1642 def _redact_responses_api_output(
1643 self,
1644 output_items: list,
1645 sanitized_text: str | None,
1646 sanitized_messages: list[dict[str, object]] | None,
1647 ) -> bool:
1648 replacement_text: str | None = sanitized_text
1649 if not replacement_text and sanitized_messages:
1650 replacement_text = " ".join(
1651 self._normalize_message_content(m.get("content")) for m in sanitized_messages if isinstance(m, dict)
1652 ).strip()
1653 if not replacement_text:
1654 return False
1655 applied = False
1656 for item in output_items:
1657 content = getattr(item, "content", None) or (item.get("content") if isinstance(item, dict) else None)
1658 if isinstance(content, list):
1659 for part in content:
1660 if isinstance(part, dict):
1661 if part.get("type") in self._TEXT_PART_TYPES:
1662 part["text"] = replacement_text
1663 applied = True
1664 else:
1665 ptype = getattr(part, "type", None)
1666 if ptype in self._TEXT_PART_TYPES:
1667 try:
1668 setattr(part, "text", replacement_text)
1669 applied = True
1670 except (AttributeError, TypeError, ValueError):
1671 continue
1672 args = item.get("arguments") if isinstance(item, dict) else getattr(item, "arguments", None)
1673 if isinstance(args, str) and args:
1674 self._clear_arguments_field(item)
1675 applied = True
1676 return applied
1678 @staticmethod
1679 def _sanitized_messages_to_responses_input(
1680 sanitized_messages: list[dict[str, object]],
1681 ) -> list[dict[str, object]] | None:
1682 """Convert chat-shape sanitized_messages to Responses API ``input``.
1684 Returns ``None`` if nothing usable could be converted, so the
1685 caller falls back to ``on_flagged_action``.
1686 """
1687 out: Final[list[dict[str, object]]] = []
1688 for m in sanitized_messages:
1689 if not isinstance(m, dict):
1690 continue
1691 role = m.get("role") or "user"
1692 content = m.get("content")
1693 if isinstance(content, str):
1694 ptype = "output_text" if role == "assistant" else "input_text"
1695 out.append({"role": role, "content": [{"type": ptype, "text": content}]})
1696 elif isinstance(content, list):
1697 out.append({"role": role, "content": content})
1698 return out or None
1700 @staticmethod
1701 def _rewrite_responses_input_text(original_input: object, sanitized_text: str) -> object | None:
1702 """Apply ``sanitized_text`` to a Responses API ``input`` value.
1704 Handles plain string, list of message items (rewrites the last
1705 user item's first text part), and flat list of content parts.
1706 Returns ``None`` if no text part could be rewritten.
1707 """
1708 if isinstance(original_input, str):
1709 return sanitized_text
1710 if not isinstance(original_input, list):
1711 return None
1713 text_types: Final = CiscoAIDefenseGuardrail._TEXT_PART_TYPES
1714 has_messages: Final = any(isinstance(i, dict) and "role" in i for i in original_input)
1716 if has_messages:
1717 rewritten: Final = list(original_input)
1718 for idx in range(len(rewritten) - 1, -1, -1):
1719 item = rewritten[idx]
1720 if not (isinstance(item, dict) and item.get("role") == "user"):
1721 continue
1722 content = item.get("content")
1723 if isinstance(content, str):
1724 rewritten[idx] = {**item, "content": sanitized_text}
1725 return rewritten
1726 if isinstance(content, list):
1727 new_content = list(content)
1728 for j, part in enumerate(new_content):
1729 if isinstance(part, dict) and part.get("type") in text_types:
1730 new_content[j] = {**part, "text": sanitized_text}
1731 rewritten[idx] = {**item, "content": new_content}
1732 return rewritten
1733 return None
1735 rewritten_parts: Final = list(original_input)
1736 for j, part in enumerate(rewritten_parts):
1737 if isinstance(part, dict) and part.get("type") in text_types:
1738 rewritten_parts[j] = {**part, "text": sanitized_text}
1739 return rewritten_parts
1740 return None
1742 @staticmethod
1743 def _extract_masked_entity_count(
1744 rules: list[dict[str, Any]],
1745 ) -> dict[str, int] | None:
1746 """Count entity-type detections per Cisco rule for the logging payload."""
1747 if not rules:
1748 return None
1749 counts: Final[dict[str, int]] = {}
1750 for rule in rules:
1751 if not isinstance(rule, dict):
1752 continue
1753 entity_types = rule.get("entity_types") or []
1754 for entity_type in entity_types:
1755 if not isinstance(entity_type, str):
1756 continue
1757 counts[entity_type] = counts.get(entity_type, 0) + 1
1758 return counts or None
1760 # ------------------------------------------------------------------
1761 # Error handling
1762 # ------------------------------------------------------------------
1764 def _handle_api_error(
1765 self,
1766 error: Exception,
1767 *,
1768 request_data: dict | None = None,
1769 start_time: datetime | None = None,
1770 surface: str = "chat",
1771 direction: str = "input",
1772 ) -> dict[str, object]:
1773 verbose_proxy_logger.error(
1774 "Cisco AI Defense guardrail (%s): API communication failed: %s",
1775 surface,
1776 error,
1777 )
1779 if request_data is not None and start_time is not None:
1780 end_time: Final = datetime.now()
1781 duration: Final = (end_time - start_time).total_seconds()
1782 if surface == "mcp":
1783 evt = GuardrailEventHooks.during_mcp_call if direction == "output" else GuardrailEventHooks.pre_mcp_call
1784 else:
1785 evt = GuardrailEventHooks.post_call if direction == "output" else GuardrailEventHooks.pre_call
1786 self.add_standard_logging_guardrail_information_to_request_data(
1787 guardrail_provider=self._PROVIDER_NAME,
1788 guardrail_json_response={
1789 "error": str(error),
1790 "error_type": type(error).__name__,
1791 "surface": surface,
1792 },
1793 request_data=request_data,
1794 guardrail_status="guardrail_failed_to_respond",
1795 start_time=start_time.timestamp(),
1796 end_time=end_time.timestamp(),
1797 duration=duration,
1798 event_type=evt,
1799 )
1801 if self.fallback_on_error == "allow":
1802 verbose_proxy_logger.warning(
1803 "Cisco AI Defense guardrail: API unavailable, proceeding without scanning (fallback_on_error='allow')"
1804 )
1805 return {
1806 "is_safe": True,
1807 "classifications": [],
1808 "_unscanned": True,
1809 }
1811 raise HTTPException(
1812 status_code=503,
1813 detail={
1814 "error": "Cisco AI Defense guardrail unavailable",
1815 "message": (
1816 "Cisco AI Defense scanning service is temporarily unavailable and fallback_on_error='block'"
1817 ),
1818 "error_type": type(error).__name__,
1819 },
1820 )
1822 # ------------------------------------------------------------------
1823 # Message extraction helpers
1824 # ------------------------------------------------------------------
1826 # Content-part ``type`` values that should be flattened to text by
1827 # ``_normalize_message_content``. Covers both Chat Completions
1828 # (``text``) and the Responses API (``input_text`` for caller-side
1829 # parts, ``output_text`` for assistant turns, ``summary_text`` /
1830 # ``reasoning_text`` for reasoning summaries that may appear in
1831 # conversation history).
1832 _TEXT_PART_TYPES = frozenset({"text", "input_text", "output_text", "summary_text", "reasoning_text"})
1834 @staticmethod
1835 def _extract_inspect_messages_from_request(
1836 data: dict,
1837 ) -> list[dict[str, str]]:
1838 """Build {role, content} messages for the Cisco AI Defense chat API."""
1839 messages: Final[list[dict[str, str]]] = []
1841 instructions_text: Final = CiscoAIDefenseGuardrail._normalize_message_content(data.get("instructions"))
1842 if instructions_text:
1843 messages.append({"role": "system", "content": instructions_text})
1845 raw_messages: Final = data.get("messages") or []
1846 for message in raw_messages:
1847 if not isinstance(message, dict):
1848 continue
1849 role = message.get("role")
1850 if not role:
1851 continue
1852 parts: list[str] = []
1853 text = CiscoAIDefenseGuardrail._normalize_message_content(message.get("content"))
1854 if text:
1855 parts.append(text)
1856 parts.extend(CiscoAIDefenseGuardrail._extract_message_tool_argument_parts(message))
1857 if parts:
1858 messages.append({"role": role, "content": " ".join(parts)})
1860 if "input" in data:
1861 # Responses API ``input`` can be: a plain string, a list of
1862 # message-shaped dicts (with role + nested content array), or
1863 # a flat list of content-part dicts. Flatten properly so the
1864 # scan sees every text segment, not just the top-level ones.
1865 messages.extend(CiscoAIDefenseGuardrail._flatten_responses_input(data.get("input")))
1867 if not messages and data.get("prompt") is not None:
1868 prompt_text: Final = CiscoAIDefenseGuardrail._normalize_message_content(data.get("prompt"))
1869 if prompt_text:
1870 messages.append({"role": "user", "content": prompt_text})
1872 tool_text: Final = CiscoAIDefenseGuardrail._extract_tool_definition_text(data)
1873 if tool_text:
1874 messages.append({"role": "system", "content": tool_text})
1876 return messages
1878 @staticmethod
1879 def _extract_tool_definition_text(data: dict) -> str:
1880 """Flatten request-side tool/function definitions into scannable text.
1882 Tool definitions (names, descriptions, nested JSON-schema docs) are
1883 forwarded to the model, so attacker-controlled text placed there must
1884 be inspected too; otherwise it bypasses the guardrail by hiding in
1885 ``tools[].function.description`` and similar metadata.
1886 """
1887 parts: Final[list[str]] = []
1888 for key in ("tools", "functions"):
1889 CiscoAIDefenseGuardrail._collect_strings(data.get(key), parts)
1890 return " ".join(parts)
1892 @staticmethod
1893 def _collect_strings(value: object, out: list[str]) -> None:
1894 if isinstance(value, str):
1895 if value:
1896 out.append(value)
1897 elif isinstance(value, dict):
1898 for item in value.values():
1899 CiscoAIDefenseGuardrail._collect_strings(item, out)
1900 elif isinstance(value, list):
1901 for item in value:
1902 CiscoAIDefenseGuardrail._collect_strings(item, out)
1904 @staticmethod
1905 def _flatten_responses_input(input_value: object) -> list[dict[str, str]]:
1906 """Flatten the OpenAI Responses API ``input`` into chat-message form.
1908 Recognized shapes:
1910 1. Plain string -> one user message.
1911 2. List of message-shaped dicts
1912 ``{"role": "...", "content": [<content parts>]}`` -> one
1913 message per item, with the role preserved.
1914 3. Flat list of content-part dicts
1915 ``{"type": "input_text", "text": "..."}`` -> single user
1916 message containing the concatenated text.
1918 """
1919 if input_value is None:
1920 return []
1921 if isinstance(input_value, str):
1922 return [{"role": "user", "content": input_value}]
1923 if not isinstance(input_value, list):
1924 text = str(input_value)
1925 return [{"role": "user", "content": text}] if text else []
1927 if any(isinstance(item, dict) and "role" in item for item in input_value):
1928 result: Final[list[dict[str, str]]] = []
1929 for item in input_value:
1930 if not isinstance(item, dict):
1931 continue
1932 role = item.get("role") or "user"
1933 text = CiscoAIDefenseGuardrail._normalize_message_content([item])
1934 if text:
1935 result.append({"role": role, "content": text})
1936 return result
1938 text = CiscoAIDefenseGuardrail._normalize_message_content(input_value)
1939 return [{"role": "user", "content": text}] if text else []
1941 @staticmethod
1942 def _normalize_message_content(content: object) -> str:
1943 """Coerce OpenAI multi-modal content into a plain text string.
1945 Supports:
1947 * Plain string.
1948 * List of content-part dicts where ``type`` is one of
1949 ``text`` (Chat Completions), ``input_text`` / ``output_text`` /
1950 ``summary_text`` (Responses API).
1951 * List of message-shaped dicts with a nested ``content`` list —
1952 recurses into the nested content so a Responses API ``input``
1953 item like ``{"role":"user","content":[{"type":"input_text",...}]}``
1954 gets flattened correctly.
1955 """
1956 if content is None:
1957 return ""
1958 if isinstance(content, str):
1959 return content
1960 if isinstance(content, list):
1961 parts: Final[list[str]] = []
1962 for part in content:
1963 if not isinstance(part, dict):
1964 continue
1965 part_type = part.get("type")
1966 if part_type in CiscoAIDefenseGuardrail._TEXT_PART_TYPES and part.get("text"):
1967 parts.append(str(part["text"]))
1968 continue
1969 nested = part.get("content")
1970 if nested is not None:
1971 nested_text = CiscoAIDefenseGuardrail._normalize_message_content(nested)
1972 if nested_text:
1973 parts.append(nested_text)
1974 for key in ("arguments", "output"):
1975 value = part.get(key)
1976 if value:
1977 parts.append(CiscoAIDefenseGuardrail._normalize_message_content(value))
1978 return " ".join(parts)
1979 return str(content)
1981 @staticmethod
1982 def _extract_response_messages(response: object) -> list[dict[str, str]]:
1983 """Extract scannable assistant text from a chat response.
1985 Handles both ``ModelResponse`` (Chat Completions) and
1986 ``ResponsesAPIResponse`` (``/v1/responses``). On both shapes
1987 tool-call / function-call argument strings and reasoning fields
1988 are included alongside the main text so a model can't bypass the
1989 scan by placing content there.
1990 """
1991 if isinstance(response, ModelResponse):
1992 result: Final[list[dict[str, str]]] = []
1993 for choice in getattr(response, "choices", None) or []:
1994 if not isinstance(choice, Choices):
1995 continue
1996 parts: list[str] = []
1997 content = CiscoAIDefenseGuardrail._normalize_message_content(getattr(choice.message, "content", None))
1998 if content:
1999 parts.append(content)
2000 parts.extend(CiscoAIDefenseGuardrail._extract_message_tool_argument_parts(choice.message))
2001 parts.extend(CiscoAIDefenseGuardrail._extract_message_reasoning_parts(choice.message))
2002 if parts:
2003 result.append({"role": "assistant", "content": " ".join(parts)})
2004 return result
2006 if isinstance(response, TextCompletionResponse):
2007 text_parts: Final[list[str]] = []
2008 for choice in getattr(response, "choices", None) or []:
2009 text = getattr(choice, "text", None)
2010 if isinstance(text, str) and text:
2011 text_parts.append(text)
2012 joined = " ".join(text_parts)
2013 return [{"role": "assistant", "content": joined}] if joined else []
2015 output_items: Final = getattr(response, "output", None)
2016 if not isinstance(output_items, list):
2017 return []
2018 output_parts: Final[list[str]] = []
2019 for item in output_items:
2020 get = item.get if isinstance(item, dict) else (lambda k: getattr(item, k, None))
2021 for part in get("content") or []:
2022 pget = part.get if isinstance(part, dict) else (lambda k: getattr(part, k, None))
2023 for key in ("text", "reasoning", "thinking"):
2024 value = pget(key)
2025 if isinstance(value, str) and value:
2026 output_parts.append(value)
2027 args = get("arguments")
2028 if isinstance(args, str) and args:
2029 output_parts.append(args)
2030 direct = get("text")
2031 if isinstance(direct, str) and direct:
2032 output_parts.append(direct)
2033 joined = " ".join(output_parts)
2034 return [{"role": "assistant", "content": joined}] if joined else []
2036 @classmethod
2037 def _extract_message_reasoning_parts(cls, message: object) -> list[str]:
2038 """Extract inspectable reasoning fields from a message/delta object."""
2039 parts: Final[list[str]] = []
2040 reasoning_content: Final = cls._field(message, "reasoning_content")
2041 if isinstance(reasoning_content, str) and reasoning_content:
2042 parts.append(reasoning_content)
2043 for block in cls._field_list(message, "thinking_blocks"):
2044 # Do not forward redacted_thinking.data; it is opaque provider
2045 # metadata rather than scannable plaintext.
2046 for key in ("thinking", "reasoning", "text"):
2047 value = cls._field(block, key)
2048 if isinstance(value, str) and value:
2049 parts.append(value)
2050 for item in cls._field_list(message, "reasoning_items"):
2051 for block in cls._field_list(item, "summary"):
2052 text = cls._field(block, "text")
2053 if isinstance(text, str) and text:
2054 parts.append(text)
2055 for key in ("text", "reasoning", "reasoning_content"):
2056 value = cls._field(item, key)
2057 if isinstance(value, str) and value:
2058 parts.append(value)
2059 return parts
2061 @staticmethod
2062 def _field(obj: object, key: str) -> object:
2063 if isinstance(obj, dict):
2064 return obj.get(key)
2065 return getattr(obj, key, None)
2067 @classmethod
2068 def _field_list(cls, obj: object, key: str) -> list[object]:
2069 value: Final = cls._field(obj, key)
2070 return value if isinstance(value, list) else []
2072 @classmethod
2073 def _extract_message_tool_argument_parts(cls, message: object) -> list[str]:
2074 parts: Final[list[str]] = []
2075 tool_calls = message.get("tool_calls") if isinstance(message, dict) else getattr(message, "tool_calls", None)
2076 for tool_call in tool_calls or []:
2077 args = cls._extract_tool_call_arguments(tool_call)
2078 if args:
2079 parts.append(args)
2080 function_call: Final = (
2081 message.get("function_call") if isinstance(message, dict) else getattr(message, "function_call", None)
2082 )
2083 if function_call is not None:
2084 args = cls._extract_function_call_arguments(function_call)
2085 if args:
2086 parts.append(args)
2087 return parts
2089 @staticmethod
2090 def _extract_tool_call_arguments(tool_call: object) -> str | None:
2091 """Pull ``function.arguments`` off a tool_calls entry (dict or model)."""
2092 if tool_call is None:
2093 return None
2094 function = tool_call.get("function") if isinstance(tool_call, dict) else getattr(tool_call, "function", None)
2095 return CiscoAIDefenseGuardrail._extract_function_call_arguments(function)
2097 @staticmethod
2098 def _extract_function_call_arguments(function_call: object) -> str | None:
2099 """Pull ``arguments`` off a function_call entry (dict or model)."""
2100 if function_call is None:
2101 return None
2102 args: Final = (
2103 function_call.get("arguments")
2104 if isinstance(function_call, dict)
2105 else getattr(function_call, "arguments", None)
2106 )
2107 if args is None:
2108 return None
2109 return str(args)
2111 # ------------------------------------------------------------------
2112 # Config model surface
2113 # ------------------------------------------------------------------
2115 @staticmethod
2116 def get_config_model() -> type["GuardrailConfigModel"] | None:
2117 from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import (
2118 CiscoAIDefenseGuardrailConfigModel,
2119 )
2121 return CiscoAIDefenseGuardrailConfigModel
2123 @classmethod
2124 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
2125 return [
2126 GuardrailEventHooks.pre_call,
2127 GuardrailEventHooks.during_call,
2128 GuardrailEventHooks.post_call,
2129 GuardrailEventHooks.logging_only,
2130 GuardrailEventHooks.pre_mcp_call,
2131 GuardrailEventHooks.during_mcp_call,
2132 ]