Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense_mcp.py: 7%
347 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"""MCP-specific inspection logic for the Cisco AI Defense guardrail.
3The public guardrail class imports this private mixin from
4``cisco_ai_defense.py``. Keeping MCP logic here avoids circular imports
5while preserving the existing public import path.
6"""
8from collections.abc import Sequence
9from datetime import datetime
10from typing import TYPE_CHECKING, Final, Optional
12from fastapi import HTTPException
14from litellm._logging import verbose_proxy_logger
15from litellm.proxy._types import UserAPIKeyAuth
16from litellm.proxy.common_utils.callback_utils import (
17 add_guardrail_to_applied_guardrails_header,
18)
19from litellm.types.guardrails import GuardrailEventHooks
21if TYPE_CHECKING: 21 ↛ 22line 21 didn't jump to line 22 because the condition on line 21 was never true
22 from litellm.types.mcp import MCPPostCallResponseObject
24 from .cisco_ai_defense import _ScanContext
27def _serialize_mcp_content_item(item: object) -> dict[str, object]:
28 """Serialize an MCP content item to a JSON-friendly dict.
30 Handles raw dicts, MCP SDK Pydantic models, and simple ``.text`` objects.
31 """
32 if isinstance(item, dict):
33 return dict(item)
34 model_dump: Final = getattr(item, "model_dump", None)
35 if callable(model_dump):
36 try:
37 dumped: Final[dict[str, object]] = model_dump(exclude_none=True, by_alias=True)
38 return dict(dumped)
39 except TypeError:
40 dumped_fallback: Final[dict[str, object]] = model_dump()
41 return dict(dumped_fallback)
42 text: Final = getattr(item, "text", None)
43 if isinstance(text, str):
44 return {"type": getattr(item, "type", "text"), "text": text}
45 return {"type": "text", "text": str(item)}
48def _source_field(source: object, key: str, snake_key: str) -> object:
49 if isinstance(source, dict):
50 for candidate in (key, snake_key):
51 if candidate in source:
52 return source[candidate] # pyright: ignore[reportUnknownVariableType] # dict-shaped sources arrive untyped
53 return None
54 return getattr(source, snake_key, None)
57class _CiscoAIDefenseMcpMixin:
58 """MCP-specific instance methods for ``CiscoAIDefenseGuardrail``.
60 Holds the MCP hooks, JSON-RPC payload builders, and redaction helpers.
61 """
63 if TYPE_CHECKING: 63 ↛ 64line 63 didn't jump to line 64 because the condition on line 63 was never true
64 api_base: str
65 inspect_path: str
66 inspection_type: str
67 _PROVIDER_NAME: str
68 guardrail_name: str | None
70 def should_run_guardrail(self, data: dict, event_type: GuardrailEventHooks) -> bool: ...
72 async def _post_inspection(self, url: str, payload: dict[str, object], surface: str) -> dict[str, object]: ...
74 def _handle_api_error(
75 self,
76 error: Exception,
77 *,
78 request_data: dict | None = ...,
79 start_time: datetime | None = ...,
80 surface: str = ...,
81 direction: str = ...,
82 ) -> dict[str, object]: ...
84 def _finalize_inspection(
85 self,
86 inspect_response: dict[str, object],
87 request_data: dict,
88 context: "_ScanContext",
89 start_time: datetime,
90 response_obj: object = ...,
91 ) -> dict[str, object]: ...
93 # ------------------------------------------------------------------
94 # MCP post-tool hook (dispatcher contract)
95 # ------------------------------------------------------------------
97 async def async_post_mcp_tool_call_hook(
98 self,
99 kwargs: dict,
100 response_obj: "MCPPostCallResponseObject",
101 start_time: datetime,
102 end_time: datetime,
103 ) -> Optional["MCPPostCallResponseObject"]:
104 """Scan MCP tool output and return a replacement object on block."""
105 del start_time, end_time
107 if self.inspection_type != "mcp":
108 return None
110 request_data: Final[dict[str, object]] = {}
111 for key in (
112 "name",
113 "litellm_call_id",
114 "id",
115 "user",
116 "mcp_tool_name",
117 "tool_name",
118 "mcp_arguments",
119 "arguments",
120 "mcp_server_name",
121 "server_name",
122 "metadata",
123 "litellm_metadata",
124 "mcp_tool_call_metadata",
125 "guardrails",
126 ):
127 if key in kwargs and kwargs[key] is not None:
128 request_data[key] = kwargs[key]
129 self._hydrate_mcp_tool_context(request_data)
131 if not (
132 self.should_run_guardrail(
133 data=request_data,
134 event_type=GuardrailEventHooks.during_mcp_call,
135 )
136 or self.should_run_guardrail(
137 data=request_data,
138 event_type=GuardrailEventHooks.pre_mcp_call,
139 )
140 ):
141 verbose_proxy_logger.debug(
142 "Cisco AI Defense guardrail (%s): no MCP mode configured — skipping MCP response scan.",
143 self.guardrail_name,
144 )
145 return None
147 mcp_tool_response: Final = self._extract_mcp_tool_call_response(response_obj)
148 if mcp_tool_response is None:
149 verbose_proxy_logger.debug("Cisco AI Defense guardrail: no MCP tool response payload to scan, skipping")
150 return None
152 original_response: Final = kwargs.get("original_response")
153 try:
154 await self._inspect_mcp_response(
155 request_data=request_data,
156 response=mcp_tool_response,
157 redact_response_obj=(original_response if original_response is not None else mcp_tool_response),
158 )
159 except HTTPException as exc:
160 blocking_response = self._build_blocking_mcp_response(detail=exc.detail, original_response_obj=response_obj)
161 self._replace_mcp_tool_response(response_obj, blocking_response)
162 if original_response is not None:
163 self._replace_mcp_tool_response(original_response, blocking_response)
164 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
165 verbose_proxy_logger.warning(
166 "Cisco AI Defense guardrail (%s): MCP response blocked — "
167 "tool output replaced with synthesized violation message.",
168 self.guardrail_name,
169 )
170 return blocking_response
172 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
173 return None
175 def _build_blocking_mcp_response(
176 self,
177 detail: object,
178 original_response_obj: object,
179 ) -> "MCPPostCallResponseObject":
180 """Build a synthetic MCPPostCallResponseObject for blocked output."""
181 import json as _json
183 from mcp.types import TextContent
185 from litellm.types.llms.base import HiddenParams
186 from litellm.types.mcp import MCPPostCallResponseObject
188 if isinstance(detail, dict):
189 payload = detail
190 else:
191 payload = {
192 "error": "Blocked by Cisco AI Defense Guardrail",
193 "message": (str(detail) if detail else "Blocked by Cisco AI Defense Guardrail"),
194 "provider": self._PROVIDER_NAME,
195 "guardrail": self.guardrail_name,
196 "surface": "mcp",
197 "direction": "output",
198 "action": "block",
199 }
201 original_hidden: Final = getattr(original_response_obj, "hidden_params", None)
202 if isinstance(original_hidden, HiddenParams):
203 hidden_params: HiddenParams = original_hidden
204 else:
205 response_cost: Final[float | None] = getattr(original_hidden, "response_cost", None)
206 hidden_params = HiddenParams(response_cost=response_cost) if response_cost is not None else HiddenParams()
208 return MCPPostCallResponseObject(
209 mcp_tool_call_response=[TextContent(type="text", text=_json.dumps(payload))],
210 hidden_params=hidden_params,
211 )
213 @staticmethod
214 def _replace_mcp_tool_response(response_obj: object, replacement_obj: object) -> bool:
215 replacement: Final[list[object] | None] = getattr(replacement_obj, "mcp_tool_call_response", None)
216 if replacement is None:
217 return False
219 inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
220 if inner is not None:
221 if _CiscoAIDefenseMcpMixin._replace_mcp_tool_response(inner, replacement_obj):
222 return True
223 try:
224 setattr(response_obj, "mcp_tool_call_response", replacement)
225 return True
226 except (AttributeError, TypeError, ValueError):
227 return False
229 content: Final = getattr(response_obj, "content", None)
230 if isinstance(content, list):
231 content[:] = replacement
232 structured_replacement: Final = _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement)
233 if hasattr(response_obj, "structured_content"):
234 try:
235 setattr(response_obj, "structured_content", structured_replacement)
236 except (AttributeError, TypeError, ValueError):
237 pass
238 if hasattr(response_obj, "is_error"):
239 try:
240 setattr(response_obj, "is_error", True)
241 except (AttributeError, TypeError, ValueError):
242 pass
243 return True
245 if isinstance(response_obj, list):
246 response_obj[:] = replacement
247 return True
249 if isinstance(response_obj, dict):
250 result: Final = response_obj.get("result")
251 if isinstance(result, dict):
252 result["content"] = replacement
253 result["structuredContent"] = _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement)
254 result["isError"] = True
255 return True
256 response_obj["result"] = {
257 "content": replacement,
258 "structuredContent": _CiscoAIDefenseMcpMixin._replacement_structured_content(replacement),
259 "isError": True,
260 }
261 return True
263 return False
265 @staticmethod
266 def _replacement_structured_content(
267 replacement: object,
268 ) -> dict[str, str] | None:
269 if not isinstance(replacement, list) or not replacement:
270 return None
271 first: Final = replacement[0]
272 text: Final = first.get("text") if isinstance(first, dict) else getattr(first, "text", None)
273 return {"result": text} if isinstance(text, str) else None
275 @staticmethod
276 def _extract_mcp_tool_call_response(response_obj: object) -> object:
277 """Pull the raw tool-call response off a MCPPostCallResponseObject."""
278 inner = getattr(response_obj, "mcp_tool_call_response", None)
279 if inner is None and isinstance(response_obj, dict):
280 inner = response_obj.get("mcp_tool_call_response")
281 return inner if inner is not None else response_obj
283 # ------------------------------------------------------------------
284 # MCP request / response inspection
285 # ------------------------------------------------------------------
287 async def _inspect_mcp_request(
288 self,
289 data: dict,
290 user_api_key_dict: UserAPIKeyAuth,
291 ) -> dict[str, object]:
292 del user_api_key_dict # carried via logging metadata, not the wire payload
293 url: Final = f"{self.api_base}{self.inspect_path}"
294 payload: Final = self._build_mcp_request_payload(data=data)
295 if payload is None:
296 verbose_proxy_logger.debug("Cisco AI Defense guardrail: could not build MCP request payload, skipping")
297 return {}
298 start_time: Final = datetime.now()
299 try:
300 inspect_response: Final = await self._post_inspection(url=url, payload=payload, surface="mcp")
301 except HTTPException:
302 raise
303 except Exception as exc:
304 return self._handle_api_error(
305 exc,
306 request_data=data,
307 start_time=start_time,
308 surface="mcp",
309 direction="input",
310 )
312 from .cisco_ai_defense import _ScanContext
314 return self._finalize_inspection(
315 inspect_response=inspect_response,
316 request_data=data,
317 context=_ScanContext(surface="mcp", direction="input"),
318 start_time=start_time,
319 )
321 async def _inspect_mcp_response(
322 self,
323 request_data: dict,
324 response: object,
325 user_api_key_dict: UserAPIKeyAuth | None = None,
326 redact_response_obj: object = None,
327 ) -> dict[str, object]:
328 del user_api_key_dict # carried via logging metadata, not the wire payload
329 url: Final = f"{self.api_base}{self.inspect_path}"
330 payload: Final = self._build_mcp_response_payload(
331 request_data=request_data,
332 response=response,
333 )
334 if payload is None:
335 verbose_proxy_logger.debug("Cisco AI Defense guardrail: could not build MCP response payload, skipping")
336 return {}
337 start_time: Final = datetime.now()
338 try:
339 inspect_response: Final = await self._post_inspection(url=url, payload=payload, surface="mcp")
340 except HTTPException:
341 raise
342 except Exception as exc:
343 return self._handle_api_error(
344 exc,
345 request_data=request_data,
346 start_time=start_time,
347 surface="mcp",
348 direction="output",
349 )
351 from .cisco_ai_defense import _ScanContext
353 return self._finalize_inspection(
354 inspect_response=inspect_response,
355 request_data=request_data,
356 context=_ScanContext(surface="mcp", direction="output"),
357 start_time=start_time,
358 response_obj=(response if redact_response_obj is None else redact_response_obj),
359 )
361 def _build_mcp_request_payload(
362 self,
363 data: dict,
364 ) -> dict[str, object] | None:
365 """Build the JSON-RPC ``tools/call`` envelope sent to ``/inspect/mcp``.
367 The Cisco AI Defense MCP inspect endpoint expects the JSON-RPC
368 envelope itself as the request body — *not* wrapped under a
369 ``request`` key with sibling ``metadata`` / ``config`` keys. Policies
370 are applied based on the API key linked to the request. Operator
371 metadata (user, call id, src/dst app, etc.) is carried out-of-band
372 via the standard logging payload so the wire contract stays
373 identical to a hand-rolled ``curl`` against ``/inspect/mcp``.
374 """
375 if data.get("jsonrpc") == "2.0":
376 return {
377 "jsonrpc": "2.0",
378 "id": (data.get("id") or data.get("litellm_call_id") or "litellm-mcp"),
379 "method": data.get("method") or "tools/call",
380 "params": data.get("params") or {},
381 }
383 tool_name: Final = data.get("mcp_tool_name") or data.get("tool_name") or data.get("name")
384 if not tool_name:
385 return None
387 arguments = data.get("mcp_arguments")
388 if arguments is None:
389 arguments = data.get("arguments")
391 return {
392 "jsonrpc": "2.0",
393 "id": data.get("litellm_call_id") or "litellm-mcp",
394 "method": "tools/call",
395 "params": {
396 "name": tool_name,
397 "arguments": (arguments if isinstance(arguments, dict) else {}),
398 },
399 }
401 def _build_mcp_response_payload(
402 self,
403 request_data: dict,
404 response: object,
405 ) -> dict[str, object] | None:
406 """Build the MCP response-inspection body sent to ``/inspect/mcp``."""
407 request_payload: Final = self._build_mcp_request_payload(data=request_data)
408 if request_payload is None:
409 return None
410 normalized: Final = self._normalize_mcp_response(response)
411 if normalized is None:
412 return None
414 payload: Final = dict(request_payload)
415 response_id: Final = normalized.get("id")
416 if response_id not in (None, "litellm-mcp"):
417 payload["id"] = response_id
418 elif payload.get("id") in (None, "litellm-mcp"):
419 request_id: Final = request_data.get("litellm_call_id") or request_data.get("id")
420 if request_id:
421 payload["id"] = request_id
423 if "result" in normalized:
424 payload["result"] = normalized["result"]
425 if "error" in normalized:
426 payload["error"] = normalized["error"]
427 return payload
429 @staticmethod
430 def _hydrate_mcp_tool_context(request_data: dict[str, object]) -> None:
431 metadata = request_data.get("mcp_tool_call_metadata")
432 if metadata is None:
433 nested: Final = request_data.get("metadata") or request_data.get("litellm_metadata")
434 if isinstance(nested, dict):
435 metadata = nested.get("mcp_tool_call_metadata")
436 if not isinstance(metadata, dict):
437 return
439 name: Final = metadata.get("name")
440 arguments: Final = metadata.get("arguments")
441 server_name: Final = metadata.get("mcp_server_name")
443 if name:
444 request_data.setdefault("mcp_tool_name", name)
445 request_data.setdefault("tool_name", name)
446 request_data.setdefault("name", name)
447 if arguments is not None:
448 request_data.setdefault("mcp_arguments", arguments)
449 request_data.setdefault("arguments", arguments)
450 if server_name:
451 request_data.setdefault("mcp_server_name", server_name)
452 request_data.setdefault("server_name", server_name)
454 @staticmethod
455 def _normalize_mcp_response(response: object) -> dict[str, object] | None:
456 """Normalize an MCP tool response into a JSON-RPC envelope.
458 Handles JSON-RPC dicts, raw content lists, MCP SDK models, and
459 Pydantic-coerced ``[(field_name, value)]`` lists.
460 """
461 if isinstance(response, dict):
462 if response.get("jsonrpc") == "2.0":
463 return dict(response)
464 if isinstance(response.get("result"), dict):
465 return {
466 "jsonrpc": "2.0",
467 "id": response.get("id") or "litellm-mcp",
468 "result": response["result"],
469 }
470 content = response.get("content")
471 if isinstance(content, list):
472 return {
473 "jsonrpc": "2.0",
474 "id": response.get("id") or "litellm-mcp",
475 "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response),
476 }
477 if isinstance(response, list):
478 if response and all(
479 isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str) for item in response
480 ):
481 response_fields: Final = dict(response)
482 inner_content: Final = response_fields.get("content")
483 if isinstance(inner_content, list):
484 return {
485 "jsonrpc": "2.0",
486 "id": "litellm-mcp",
487 "result": _CiscoAIDefenseMcpMixin._build_mcp_result(
488 content=inner_content, source=response_fields
489 ),
490 }
491 else:
492 return None
493 return {
494 "jsonrpc": "2.0",
495 "id": "litellm-mcp",
496 "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=response),
497 }
498 model_dump: Final = getattr(response, "model_dump", None)
499 if callable(model_dump):
500 try:
501 dumped = model_dump(exclude_none=True, by_alias=True)
502 except TypeError:
503 dumped = model_dump()
504 if isinstance(dumped, dict):
505 return _CiscoAIDefenseMcpMixin._normalize_mcp_response(dumped)
506 content = getattr(response, "content", None)
507 if isinstance(content, list):
508 return {
509 "jsonrpc": "2.0",
510 "id": "litellm-mcp",
511 "result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=content, source=response),
512 }
513 return None
515 @staticmethod
516 def _build_mcp_result(
517 content: Sequence[object],
518 source: object = None,
519 ) -> dict[str, object]:
520 result: Final[dict[str, object]] = {"content": [_serialize_mcp_content_item(item) for item in content]}
521 for key, snake_key in (("structuredContent", "structured_content"), ("isError", "is_error")):
522 value = _source_field(source, key, snake_key)
523 if value is not None and (key != "isError" or isinstance(value, bool)):
524 result[key] = value
525 return result
527 # ------------------------------------------------------------------
528 # MCP redact (in-place rewrite of tool output)
529 # ------------------------------------------------------------------
531 @staticmethod
532 def _set_mcp_tool_response_text(response_obj: object, text: str) -> bool:
533 """Replace text content in any supported MCP response shape."""
534 if response_obj is None:
535 return False
537 inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
538 if inner is not None:
539 return _CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text)
541 content_list: Final = _CiscoAIDefenseMcpMixin._coerce_to_content_list(response_obj)
543 replaced = False
544 if isinstance(content_list, list):
545 for item in content_list:
546 if isinstance(item, dict) and item.get("type") == "text":
547 item["text"] = text
548 replaced = True
549 elif hasattr(item, "type") and getattr(item, "type", None) == "text":
550 try:
551 setattr(item, "text", text)
552 replaced = True
553 except (AttributeError, TypeError, ValueError):
554 continue
556 replacement: Final = {"result": text}
557 if (
558 isinstance(response_obj, list)
559 and response_obj
560 and all(isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str) for item in response_obj)
561 ):
562 for index, item in enumerate(response_obj):
563 if item[0] in ("structuredContent", "structured_content"):
564 response_obj[index] = (item[0], replacement)
565 replaced = True
566 elif hasattr(response_obj, "structured_content"):
567 try:
568 setattr(response_obj, "structured_content", replacement)
569 replaced = True
570 except (AttributeError, TypeError, ValueError):
571 pass
572 elif isinstance(response_obj, dict):
573 result: Final = response_obj.get("result")
574 target: Final[dict[object, object]] = result if isinstance(result, dict) else response_obj
575 structured_key: Final = "structured_content" if "structured_content" in target else "structuredContent"
576 if structured_key in target:
577 target[structured_key] = replacement
578 replaced = True
580 return replaced
582 @staticmethod
583 def _coerce_to_content_list(response_obj: object) -> list[object] | None:
584 """Find the MCP content list inside supported response shapes."""
585 if response_obj is None:
586 return None
587 inner: Final[object | None] = getattr(response_obj, "mcp_tool_call_response", None)
588 if inner is not None:
589 return _CiscoAIDefenseMcpMixin._coerce_to_content_list(inner)
590 content: Final = getattr(response_obj, "content", None)
591 if isinstance(content, list):
592 return content
593 if isinstance(response_obj, list):
594 if response_obj and all(
595 isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str) for item in response_obj
596 ):
597 inner_content: Final = dict(response_obj).get("content")
598 if isinstance(inner_content, list):
599 return inner_content
600 return None
601 return response_obj
602 return None
604 # ------------------------------------------------------------------
605 # MCP-specific verdict extraction
606 # ------------------------------------------------------------------
608 @staticmethod
609 def _extract_sanitized_mcp_arguments(
610 inspect_response: dict[str, object],
611 ) -> dict[str, object] | None:
612 """Pull sanitized MCP tool-call arguments off the verdict.
614 Cisco can return them at the top level (``params.arguments``) or
615 under ``sanitized_payload`` / ``modified_payload``.
616 """
617 containers: Final = [inspect_response]
618 for container_key in ("result", "data"):
619 container = inspect_response.get(container_key)
620 if isinstance(container, dict):
621 containers.append(container)
623 for container in containers:
624 params = container.get("params")
625 if isinstance(params, dict):
626 args = params.get("arguments")
627 if isinstance(args, dict) and args:
628 return dict(args)
629 for key in (
630 "sanitized_payload",
631 "sanitizedPayload",
632 "modified_payload",
633 "modifiedPayload",
634 ):
635 payload = container.get(key)
636 if isinstance(payload, dict):
637 inner_params = payload.get("params")
638 if isinstance(inner_params, dict):
639 args = inner_params.get("arguments")
640 if isinstance(args, dict) and args:
641 return dict(args)
642 direct = payload.get("arguments")
643 if isinstance(direct, dict) and direct:
644 return dict(direct)
645 return None