Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/tool_permission.py: 11%
473 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
1import json
2import re
3from collections.abc import AsyncGenerator, AsyncIterable, Mapping, Sequence
4from typing import Any, Final, Literal, TypedDict
6from fastapi import HTTPException
7from typing_extensions import ReadOnly, Required
9from litellm import ChatCompletionToolParam
10from litellm._logging import verbose_proxy_logger
11from litellm.caching.dual_cache import DualCache
12from litellm.exceptions import GuardrailRaisedException
13from litellm.integrations.custom_guardrail import (
14 CustomGuardrail,
15 log_guardrail_information,
16)
17from litellm.proxy._types import UserAPIKeyAuth
18from litellm.proxy.common_utils.callback_utils import (
19 add_guardrail_to_applied_guardrails_header,
20)
21from litellm.proxy.guardrails.anthropic_sse import (
22 anthropic_sse_chunks_from_response,
23 assemble_anthropic_sse_stream,
24 is_raw_sse_stream,
25)
26from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
27from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import (
28 PermissionError,
29 ToolPermissionRule,
30 ToolResult,
31)
32from litellm.types.utils import (
33 CallTypesLiteral,
34 ChatCompletionMessageToolCall,
35 Choices,
36 Function,
37 LLMResponseTypes,
38 ModelResponse,
39 ModelResponseStream,
40)
42GUARDRAIL_NAME: Final = "tool_permission"
45def _object_mapping(value: object) -> Mapping[str, object] | None:
46 """Return ``value`` as an opaque mapping when it is a dict."""
47 return value if isinstance(value, dict) else None
50def _object_list(value: object) -> Sequence[object] | None:
51 """Return ``value`` as an opaque sequence when it is a list."""
52 return value if isinstance(value, list) else None
55class _ToolPermissionRuleFields(TypedDict, total=False):
56 """The config-file shape a :class:`ToolPermissionRule` is built from."""
58 id: ReadOnly[Required[str]]
59 tool_name: ReadOnly[str | None]
60 tool_type: ReadOnly[str | None]
61 decision: ReadOnly[Required[Literal["allow", "deny"]]]
62 allowed_param_patterns: ReadOnly[dict[str, str] | None]
65def _rule_from_fields(fields: _ToolPermissionRuleFields) -> ToolPermissionRule:
66 """Validate one config-file rule entry into a :class:`ToolPermissionRule`."""
67 return ToolPermissionRule(**fields)
70def _is_tool_use_block(block: object) -> bool:
71 """Whether ``block`` is an Anthropic ``tool_use`` content block."""
72 fields: Final = _object_mapping(block)
73 return fields is not None and fields.get("type") == "tool_use"
76class ToolPermissionGuardrail(CustomGuardrail):
77 def __init__(
78 self,
79 rules: list[dict] | None = None,
80 default_action: Literal["deny", "allow"] = "deny",
81 on_disallowed_action: Literal["block", "rewrite"] = "block",
82 **kwargs,
83 ):
84 """
85 Initialize the Tool Permission Guardrail
87 Args:
88 rules: List of permission rules
89 default_action: Default action when no rule matches ("allow" or "deny")
90 on_disallowed_action:
91 **kwargs: Additional arguments passed to CustomGuardrail
92 """
93 # Set supported event hooks - this guardrail only works on post_call
94 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
96 super().__init__(**kwargs)
98 self._load_rules(rules)
100 # Normalize to lowercase for case-insensitive handling
101 self.default_action = default_action.lower() if isinstance(default_action, str) else default_action
102 self.on_disallowed_action = (
103 on_disallowed_action.lower() if isinstance(on_disallowed_action, str) else on_disallowed_action
104 )
106 verbose_proxy_logger.debug(
107 "Tool Permission Guardrail initialized with %d rules, default_action: %s",
108 len(self.rules),
109 self.default_action,
110 )
112 def _load_rules(self, rules: list[Any] | None) -> None:
113 """Parse ``rules`` and (re)build the compiled target/pattern lookups.
115 ``self.rules`` plus ``_compiled_rule_targets`` / ``_compiled_rule_patterns``
116 are the state every matching path reads. Centralizing the build here lets
117 both ``__init__`` and ``update_in_memory_litellm_params`` recompile from a
118 single source of truth, so an in-place update (PUT /guardrails, immediate
119 sync) reflects rule changes instead of keeping the construction-time maps.
120 """
121 parsed_rules: Final[list[ToolPermissionRule]] = []
122 compiled_targets: Final[dict[str, dict[str, re.Pattern | None]]] = {}
123 compiled_patterns: Final[dict[str, dict[str, re.Pattern]]] = {}
125 for rule_item in rules or []:
126 rule = rule_item if isinstance(rule_item, ToolPermissionRule) else _rule_from_fields(rule_item)
128 target_patterns: dict[str, re.Pattern | None] = {
129 "tool_name": None,
130 "tool_type": None,
131 }
132 if rule.tool_name is not None:
133 try:
134 target_patterns["tool_name"] = re.compile(rule.tool_name)
135 except re.error as exc:
136 raise ValueError(f"Invalid regex for tool_name in rule '{rule.id}': {exc}") from exc
137 if rule.tool_type is not None:
138 try:
139 target_patterns["tool_type"] = re.compile(rule.tool_type)
140 except re.error as exc:
141 raise ValueError(f"Invalid regex for tool_type in rule '{rule.id}': {exc}") from exc
143 rule_patterns: dict[str, re.Pattern] = {}
144 for path, pattern in (rule.allowed_param_patterns or {}).items():
145 try:
146 rule_patterns[path] = re.compile(pattern)
147 except re.error as exc:
148 raise ValueError(f"Invalid regex in allowed_param_patterns for rule '{rule.id}': {exc}") from exc
150 parsed_rules.append(rule)
151 compiled_targets[rule.id] = target_patterns
152 if rule_patterns:
153 compiled_patterns[rule.id] = rule_patterns
155 # Swap in the fully-built maps only after every rule compiles, so an
156 # invalid regex raises without leaving a partially-built ruleset (a
157 # missing compiled target is read as a match-all wildcard).
158 self.rules = parsed_rules
159 self._compiled_rule_targets = compiled_targets
160 self._compiled_rule_patterns = compiled_patterns
162 def update_in_memory_litellm_params(self, litellm_params: LitellmParams | dict) -> None:
163 """Apply updated params in place, rebuilding the compiled rule state.
165 The base implementation only ``setattr``s raw fields, which would leave
166 ``_compiled_rule_targets`` / ``_compiled_rule_patterns`` (built in
167 ``__init__``) stale, so a guardrail updated without reinitialization would
168 keep enforcing the old ruleset. Recompile here so PUT /guardrails and the
169 immediate in-memory sync take effect, mirroring the PresidioGuardrail
170 override of this method.
171 """
172 # ``litellm_params`` may arrive as the raw DB dict (the proxy ``cast()``s
173 # it to ``LitellmParams`` without converting), so handle both shapes. The
174 # base ``setattr`` loop is model-only, so apply the dict case here.
175 previous_rules: Final = self.rules
176 if isinstance(litellm_params, dict):
177 params = litellm_params
178 for key, value in params.items():
179 setattr(self, key, value)
180 else:
181 super().update_in_memory_litellm_params(litellm_params)
182 params = vars(litellm_params)
184 # The generic update above sets ``self.rules`` from the incoming value
185 # (None on a partial update that omits rules), but never rebuilds the
186 # compiled maps. Rebuild them when rules are provided; otherwise restore
187 # the previous ruleset so a partial update doesn't silently wipe it. An
188 # explicit empty list still clears the rules.
189 rules: Final = params.get("rules")
190 if rules is not None:
191 try:
192 self._load_rules(rules)
193 except Exception:
194 # The generic update above may have overwritten self.rules with
195 # the raw payload; restore the prior consistent ruleset so a
196 # rejected update can't leave the live guardrail enforcing a
197 # broken policy.
198 self.rules = previous_rules
199 raise
200 else:
201 self.rules = previous_rules
202 default_action: Final = params.get("default_action")
203 if isinstance(default_action, str):
204 self.default_action = default_action.lower()
205 on_disallowed_action: Final = params.get("on_disallowed_action")
206 if isinstance(on_disallowed_action, str):
207 self.on_disallowed_action = on_disallowed_action.lower()
209 @staticmethod
210 def get_config_model():
211 from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import (
212 ToolPermissionGuardrailConfigModel,
213 )
215 return ToolPermissionGuardrailConfigModel
217 @classmethod
218 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
219 return [
220 GuardrailEventHooks.pre_call,
221 GuardrailEventHooks.post_call,
222 ]
224 def _matches_regex(self, pattern: re.Pattern | None, value: str | None) -> bool:
225 if pattern is None:
226 return True
227 if value is None:
228 return False
229 return bool(pattern.fullmatch(value))
231 def _rule_matches_tool(
232 self,
233 rule: ToolPermissionRule,
234 *,
235 tool_name: str | None,
236 tool_type: str | None = None,
237 ) -> tuple[bool, bool]:
238 target_patterns: Final = self._compiled_rule_targets.get(rule.id, {})
239 name_pattern: Final = target_patterns.get("tool_name")
240 type_pattern: Final = target_patterns.get("tool_type")
242 name_required: Final = rule.tool_name is not None
243 type_required: Final = rule.tool_type is not None
245 name_matched: Final = self._matches_regex(name_pattern, tool_name) if name_required else True
246 type_matched: Final = self._matches_regex(type_pattern, tool_type) if type_required else True
248 overall_match: Final = name_matched and type_matched
249 should_check_params: Final = name_required and name_matched
251 return overall_match, should_check_params
253 def _check_tool_permission(
254 self,
255 tool_name: str | None,
256 tool_type: str | None = None,
257 ) -> tuple[bool, str | None, str | None]:
258 """
259 Check if a tool is allowed based on the configured rules
261 Args:
262 tool_name: Name of the tool to check
263 tool_type: Type of the tool to check
265 Returns:
266 Tuple of (is_allowed, rule_id, message)
267 """
268 verbose_proxy_logger.debug("Checking permission for tool: %s", tool_name or tool_type)
270 # Check each rule in order
271 for rule in self.rules:
272 matches, _ = self._rule_matches_tool(
273 rule,
274 tool_name=tool_name,
275 tool_type=tool_type,
276 )
277 if matches:
278 is_allowed = rule.decision == "allow"
279 tool_identifier = tool_name or tool_type or "unknown_tool"
280 default_message = (
281 f"Tool '{tool_identifier}' {'allowed' if is_allowed else 'denied'} by rule '{rule.id}'"
282 )
283 message = self.render_violation_message(
284 default=default_message,
285 context={
286 "tool_name": tool_name or tool_identifier,
287 "rule_id": rule.id,
288 },
289 )
290 verbose_proxy_logger.debug(message)
291 return is_allowed, rule.id, message
293 # No rule matched, use default action
294 is_allowed = self.default_action == "allow"
295 tool_identifier = tool_name or tool_type or "unknown_tool"
296 default_message = f"Tool '{tool_identifier}' {'allowed' if is_allowed else 'denied'} by default action"
297 message = self.render_violation_message(
298 default=default_message,
299 context={
300 "tool_name": tool_name or tool_identifier,
301 "rule_id": None,
302 },
303 )
304 verbose_proxy_logger.debug(message)
305 return is_allowed, None, message
307 def _parse_tool_call_arguments(
308 self, tool_call: ChatCompletionMessageToolCall
309 ) -> tuple[Mapping[str, object] | None, str | None]:
310 arguments: Final = getattr(tool_call.function, "arguments", None)
311 if not arguments:
312 return None, "missing arguments"
314 parsed_arguments: object = {}
315 try:
316 if isinstance(arguments, str):
317 parsed_arguments = json.loads(arguments)
318 elif isinstance(arguments, dict):
319 parsed_arguments = arguments
320 else:
321 return None, "arguments must be a JSON object"
322 except (json.JSONDecodeError, TypeError) as exc:
323 verbose_proxy_logger.warning(
324 "Tool Permission Guardrail: Failed to decode arguments for tool %s: %s",
325 tool_call.function.name,
326 exc,
327 )
328 return None, "arguments could not be parsed"
330 if isinstance(parsed_arguments, dict):
331 return parsed_arguments, None
333 verbose_proxy_logger.debug(
334 "Tool Permission Guardrail: Rejecting non-dict arguments for tool %s",
335 tool_call.function.name,
336 )
337 return None, "arguments must be a JSON object"
339 def _collect_argument_paths(
340 self,
341 value: object,
342 current_path: str,
343 collected: dict[str, list[object]],
344 depth: int = 0,
345 ) -> None:
346 from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
348 if depth > DEFAULT_MAX_RECURSE_DEPTH:
349 return
351 mapping_value: Final = _object_mapping(value)
352 list_value: Final = _object_list(value)
353 if mapping_value is not None:
354 for key, sub_value in mapping_value.items():
355 next_path = f"{current_path}.{key}" if current_path else key
356 self._collect_argument_paths(sub_value, next_path, collected, depth + 1)
357 elif list_value is not None:
358 list_path: Final = f"{current_path}[]" if current_path else "[]"
359 for item in list_value:
360 self._collect_argument_paths(item, list_path, collected, depth + 1)
361 else:
362 if not current_path:
363 return
364 collected.setdefault(current_path, []).append(value)
366 def _patterns_match_for_rule(
367 self,
368 *,
369 arguments: Mapping[str, object],
370 rule: ToolPermissionRule,
371 tool_name: str | None,
372 ) -> tuple[bool, str | None]:
373 compiled_patterns: Final = self._compiled_rule_patterns.get(rule.id)
374 if not compiled_patterns:
375 return True, None
377 path_value_map: Final[dict[str, list[object]]] = {}
378 self._collect_argument_paths(arguments, "", path_value_map)
380 for path, compiled_pattern in compiled_patterns.items():
381 values = path_value_map.get(path)
382 if not values:
383 return (
384 False,
385 f"Missing value for path '{path}' required by rule '{rule.id}'",
386 )
387 for raw_value in values:
388 if not compiled_pattern.fullmatch(str(raw_value)):
389 return (
390 False,
391 f"Value '{raw_value}' for path '{path}' does not match allowed pattern"
392 f" '{compiled_pattern.pattern}' for tool '{tool_name or 'unknown_tool'}'",
393 )
395 return True, None
397 def _get_permission_for_tool_call(
398 self, tool_call: ChatCompletionMessageToolCall
399 ) -> tuple[bool, str | None, str | None]:
400 tool_name: Final = tool_call.function.name if tool_call.function else None
401 tool_type: Final = getattr(tool_call, "type", None)
402 if not tool_name and not tool_type:
403 return self.default_action == "allow", None, None
405 tool_identifier: Final = tool_name or tool_type or "unknown_tool"
407 last_pattern_failure_msg: str | None = None
409 for rule in self.rules:
410 matches, should_check_params = self._rule_matches_tool(
411 rule,
412 tool_name=tool_name,
413 tool_type=tool_type,
414 )
415 if not matches:
416 continue
418 if rule.allowed_param_patterns and should_check_params:
419 arguments, parse_error = self._parse_tool_call_arguments(tool_call)
420 if parse_error:
421 default_message = f"Tool '{tool_identifier}' {parse_error} required by rule '{rule.id}'"
422 message = self.render_violation_message(
423 default=default_message,
424 context={"tool_name": tool_identifier, "rule_id": rule.id},
425 )
426 return False, rule.id, message
427 if not arguments:
428 default_message = f"Tool '{tool_identifier}' is missing arguments required by rule '{rule.id}'"
429 message = self.render_violation_message(
430 default=default_message,
431 context={"tool_name": tool_identifier, "rule_id": rule.id},
432 )
433 return False, rule.id, message
435 patterns_match, failure_message = self._patterns_match_for_rule(
436 arguments=arguments,
437 rule=rule,
438 tool_name=tool_name,
439 )
440 if not patterns_match:
441 last_pattern_failure_msg = failure_message
442 continue
444 is_allowed = rule.decision == "allow"
445 default_message = f"Tool '{tool_identifier}' {'allowed' if is_allowed else 'denied'} by rule '{rule.id}'"
446 message = self.render_violation_message(
447 default=default_message,
448 context={"tool_name": tool_identifier, "rule_id": rule.id},
449 )
450 return is_allowed, rule.id, message
452 is_allowed = self.default_action == "allow"
453 default_message = (
454 last_pattern_failure_msg
455 if (last_pattern_failure_msg and not is_allowed)
456 else f"Tool '{tool_identifier}' {'allowed' if is_allowed else 'denied'} by default action"
457 )
458 message = self.render_violation_message(
459 default=default_message,
460 context={"tool_name": tool_identifier, "rule_id": None},
461 )
462 return is_allowed, None, message
464 @staticmethod
465 def _get_mapping_value(item: object, key: str) -> Any:
466 if isinstance(item, dict):
467 return item.get(key)
468 return getattr(item, key, None)
470 @staticmethod
471 def _legacy_function_call_id(choice_index: int) -> str:
472 return f"legacy_function_call_{choice_index}"
474 def _legacy_function_call_to_tool_call(
475 self, function_call: object, choice_index: int
476 ) -> ChatCompletionMessageToolCall | None:
477 if function_call is None:
478 return None
480 function_name: Final = self._get_mapping_value(function_call, "name")
481 arguments: Final = self._get_mapping_value(function_call, "arguments") or ""
482 if not function_name:
483 return None
485 return ChatCompletionMessageToolCall(
486 id=self._legacy_function_call_id(choice_index),
487 type="function",
488 function={"name": function_name, "arguments": arguments},
489 )
491 def _extract_tool_calls_from_response(self, response: ModelResponse) -> list[ChatCompletionMessageToolCall]:
492 """
493 Extract tool_calls from all choices in a model response.
495 Args:
496 response: The model response to analyze
498 Returns:
499 List of tool_calls blocks found in the response
500 """
501 tool_calls: Final = []
503 for choice_index, choice in enumerate(response.choices):
504 if isinstance(choice, Choices):
505 for tool in choice.message.tool_calls or []:
506 tool_calls.append(tool)
507 legacy_tool_call = self._legacy_function_call_to_tool_call(
508 getattr(choice.message, "function_call", None), choice_index
509 )
510 if legacy_tool_call is not None:
511 tool_calls.append(legacy_tool_call)
513 return tool_calls
515 @staticmethod
516 def _anthropic_tool_use_to_tool_call(block: object) -> ChatCompletionMessageToolCall | None:
517 if not isinstance(block, dict) or block.get("type") != "tool_use":
518 return None
519 name: Final = block.get("name")
520 if not isinstance(name, str) or not name:
521 return None
522 tool_input: Final[object] = block.get("input")
523 return ChatCompletionMessageToolCall(
524 id=str(block.get("id") or ""),
525 function=Function(name=name, arguments=json.dumps(tool_input) if isinstance(tool_input, dict) else "{}"),
526 type="function",
527 )
529 @staticmethod
530 def _get_anthropic_content_blocks(response: object) -> tuple[object, ...] | None:
531 if not isinstance(response, dict):
532 return None
533 content: Final[object] = response.get("content")
534 return tuple(content) if isinstance(content, list) else None
536 def _extract_tool_calls_from_anthropic_content(
537 self, content: tuple[object, ...]
538 ) -> tuple[ChatCompletionMessageToolCall, ...]:
539 return tuple(
540 tool_call for block in content if (tool_call := self._anthropic_tool_use_to_tool_call(block)) is not None
541 )
543 def _evaluate_tool_calls(
544 self, tool_calls: Sequence[ChatCompletionMessageToolCall]
545 ) -> tuple[tuple[ChatCompletionMessageToolCall, PermissionError], ...]:
546 checked: Final = tuple((tool_call, *self._get_permission_for_tool_call(tool_call)) for tool_call in tool_calls)
548 for _tool_call, is_allowed, _rule_id, message in checked:
549 if not is_allowed and message is not None:
550 verbose_proxy_logger.info("Tool Permission Guardrail: %s", message)
551 if self.on_disallowed_action == "block":
552 raise GuardrailRaisedException(
553 guardrail_name=self.guardrail_name, message=message, blocked_content=True
554 )
556 return tuple(
557 (
558 tool_call,
559 PermissionError(
560 tool_name=(
561 tool_call.function.name if tool_call.function and tool_call.function.name else "unknown_tool"
562 ),
563 rule_id=rule_id,
564 message=message,
565 ),
566 )
567 for tool_call, is_allowed, rule_id, message in checked
568 if not is_allowed and message is not None
569 )
571 def _modify_anthropic_content_with_permission_errors(
572 self,
573 response: object,
574 content: tuple[object, ...],
575 denied_tools: tuple[tuple[ChatCompletionMessageToolCall, PermissionError], ...],
576 ) -> None:
577 if not denied_tools or not isinstance(response, dict):
578 return
580 verbose_proxy_logger.info("Blocking %s unauthorized tool uses", len(denied_tools))
582 error_by_tool_use_id: Final[
583 Mapping[object, str]
584 ] = { # mutable-ok: read-only lookup, never mutated after construction
585 tool_call.id: self._create_permission_error_result(tool_call, error).content
586 for tool_call, error in denied_tools
587 }
589 def _denied_message(block: object) -> str | None:
590 fields: Final = _object_mapping(block)
591 if fields is None or fields.get("type") != "tool_use":
592 return None
593 return error_by_tool_use_id.get(fields.get("id"))
595 error_messages: Final = tuple(
596 message for message in (_denied_message(block) for block in content) if message is not None
597 )
598 kept_blocks: Final = tuple(block for block in content if _denied_message(block) is None)
599 new_content: Final = [ # mutable-ok: response content is a JSON array on the wire
600 *kept_blocks,
601 {"type": "text", "text": "\n".join(error_messages)}, # mutable-ok: content block is a JSON object
602 ]
604 response["content"] = new_content # rebind-ok: the guardrail rewrites the provider response in place
605 if not any(_is_tool_use_block(block) for block in kept_blocks):
606 response["stop_reason"] = "end_turn" # rebind-ok: dropping every tool_use ends the turn
608 def _get_request_tool_name(self, tool: object) -> tuple[str | None, str | None]:
609 tool_type: Final = self._get_mapping_value(tool, "type")
610 if tool_type != "function":
611 return None, tool_type
613 function: Final = self._get_mapping_value(tool, "function")
614 tool_name: Final = self._get_mapping_value(function, "name")
615 return tool_name, tool_type
617 def _get_legacy_function_name(self, function: object) -> str | None:
618 return self._get_mapping_value(function, "name")
620 def _get_named_tool_choice(self, data: dict) -> str | None:
621 tool_choice: Final = data.get("tool_choice")
622 if not tool_choice or tool_choice in ("auto", "none", "required"):
623 return None
624 if isinstance(tool_choice, str):
625 return tool_choice
626 if self._get_mapping_value(tool_choice, "type") != "function":
627 return None
628 return self._get_mapping_value(self._get_mapping_value(tool_choice, "function"), "name")
630 def _get_named_function_call(self, data: dict) -> str | None:
631 function_call: Final = data.get("function_call")
632 if not function_call or function_call in ("auto", "none"):
633 return None
634 if isinstance(function_call, str):
635 return function_call
636 return self._get_mapping_value(function_call, "name")
638 def _collect_request_tools(self, data: dict) -> list[tuple[str, str | None]]:
639 request_tools: Final[list[tuple[str, str | None]]] = []
641 for tool in data.get("tools") or []:
642 tool_name, tool_type = self._get_request_tool_name(tool)
643 if tool_name is not None:
644 request_tools.append((tool_name, tool_type))
646 for function in data.get("functions") or []:
647 function_name = self._get_legacy_function_name(function)
648 if function_name is not None:
649 request_tools.append((function_name, "function"))
651 for forced_tool_name in (
652 self._get_named_tool_choice(data),
653 self._get_named_function_call(data),
654 ):
655 if forced_tool_name is not None:
656 request_tools.append((forced_tool_name, "function"))
658 return request_tools
660 def _modify_request_with_permission_errors(
661 self,
662 data: dict,
663 denied_tool_names: list[str],
664 ):
665 """
666 Modify the request to replace denied tool_calls blocks with error results
668 Args:
669 data: The model request to modify
670 denied_tools: List of (tool_use, error) tuples for denied tools
671 """
672 if not denied_tool_names:
673 return data
675 verbose_proxy_logger.info("Blocking %s unauthorized tool uses", len(denied_tool_names))
677 # Create a mapping of tool_use_id to error result
678 error_tool_names: Final = set()
679 for tool_use in denied_tool_names:
680 error_tool_names.add(tool_use)
682 tools: Final[list[ChatCompletionToolParam] | None] = data.get("tools")
683 if tools is not None:
684 new_tools: Final = []
685 for tool in tools:
686 tool_name, tool_type = self._get_request_tool_name(tool)
687 if tool_type == "function" and tool_name in error_tool_names:
688 continue
689 new_tools.append(tool)
690 data["tools"] = new_tools
692 functions: Final = data.get("functions")
693 if functions is not None:
694 data["functions"] = [
695 function for function in functions if self._get_legacy_function_name(function) not in error_tool_names
696 ]
698 named_tool_choice: Final = self._get_named_tool_choice(data)
699 if named_tool_choice in error_tool_names:
700 data["tool_choice"] = "none"
702 named_function_call: Final = self._get_named_function_call(data)
703 if named_function_call in error_tool_names:
704 data["function_call"] = "none"
706 return data
708 def _create_permission_error_result(
709 self, tool_call: ChatCompletionMessageToolCall, error: PermissionError
710 ) -> ToolResult:
711 """
712 Create a tool_result block for a permission error
714 Args:
715 tool_use: The tool use that was denied
716 error: The permission error details
718 Returns:
719 A tool_result block with the error message
720 """
721 error_message = f"Permission denied: {error.message}"
722 if error.rule_id:
723 error_message += f" (Rule: {error.rule_id})"
725 return ToolResult(tool_use_id=tool_call.id, content=error_message, is_error=True)
727 def _modify_response_with_permission_errors(
728 self,
729 response: ModelResponse,
730 denied_tools: Sequence[tuple[ChatCompletionMessageToolCall, PermissionError]],
731 ) -> None:
732 """
733 Modify the response to replace denied tool_calls blocks with error results
735 Args:
736 response: The model response to modify
737 denied_tools: List of (tool_use, error) tuples for denied tools
738 """
739 if not denied_tools:
740 return
742 verbose_proxy_logger.info("Blocking %s unauthorized tool uses", len(denied_tools))
744 # Create a mapping of tool_use_id to error result
745 error_results: Final = {}
746 for tool_use, error in denied_tools:
747 error_result = self._create_permission_error_result(tool_use, error)
748 error_results[tool_use.id] = error_result
750 # Modify the response content
751 for choice_index, choice in enumerate(response.choices):
752 if isinstance(choice, Choices):
753 filtered_tool_calls = []
754 error_messages = []
756 # Rewrite tool_calls
757 for tool_call in choice.message.tool_calls or []:
758 tool_call_id = tool_call.id
759 if tool_call_id in error_results:
760 error_result = error_results[tool_call_id]
761 error_messages.append(error_result.content)
762 else:
763 filtered_tool_calls.append(tool_call)
765 choice.message.tool_calls = filtered_tool_calls if filtered_tool_calls else None
767 legacy_tool_call = self._legacy_function_call_to_tool_call(
768 getattr(choice.message, "function_call", None), choice_index
769 )
770 if legacy_tool_call is not None:
771 legacy_error_result = error_results.get(legacy_tool_call.id)
772 if legacy_error_result is not None:
773 choice.message.function_call = None
774 error_messages.append(legacy_error_result.content)
776 # Add error messages to content
777 if error_messages:
778 existing_content = choice.message.content
779 if existing_content:
780 choice.message.content = existing_content + "\n\n" + "\n".join(error_messages)
781 else:
782 choice.message.content = "\n".join(error_messages)
784 if (
785 not choice.message.tool_calls
786 and getattr(choice.message, "function_call", None) is None
787 and choice.finish_reason in ("tool_calls", "function_call")
788 ):
789 choice.finish_reason = "stop"
791 @log_guardrail_information
792 async def async_pre_call_hook(
793 self,
794 user_api_key_dict: UserAPIKeyAuth,
795 cache: DualCache,
796 data: dict,
797 call_type: CallTypesLiteral,
798 ) -> Exception | str | dict | None:
799 """ """
800 verbose_proxy_logger.debug("Tool Permission Guardrail Pre-Call Hook")
802 from litellm.proxy.common_utils.callback_utils import (
803 add_guardrail_to_applied_guardrails_header,
804 )
806 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.pre_call
807 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
808 return data
810 new_tools: Final = self._collect_request_tools(data)
811 if not new_tools:
812 verbose_proxy_logger.debug(
813 "Tool Permission Guardrail: not running guardrail. No tools or functions in data"
814 )
815 return data
817 # Check permissions for each tool
818 denied_tool_names: Final = []
819 for tool_name, tool_type in new_tools:
820 is_allowed, _, message = self._check_tool_permission(tool_name, tool_type)
822 if not is_allowed and message is not None:
823 verbose_proxy_logger.info("Tool Permission Guardrail: %s", message)
824 if self.on_disallowed_action == "block":
825 raise HTTPException(
826 status_code=400,
827 detail={
828 "error": "Violated guardrail policy",
829 "detection_message": message,
830 },
831 )
832 denied_tool_names.append(tool_name)
834 if denied_tool_names:
835 data = self._modify_request_with_permission_errors(data, denied_tool_names)
837 verbose_proxy_logger.debug("Tool Permission Guardrail Pre-Call Hook: All tools allowed")
839 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
840 return data
842 @log_guardrail_information
843 async def async_post_call_success_hook(
844 self,
845 data: dict,
846 user_api_key_dict: UserAPIKeyAuth,
847 response: LLMResponseTypes,
848 ):
849 """
850 Check tool usage permissions after the LLM call
852 Args:
853 data: Request data
854 user_api_key_dict: User API key information (unused but required by interface)
855 response: The model response to check
856 """
857 anthropic_content: Final = (
858 None if isinstance(response, ModelResponse) else self._get_anthropic_content_blocks(response)
859 )
860 if not isinstance(response, ModelResponse) and anthropic_content is None:
861 return response
863 verbose_proxy_logger.debug("Tool Permission Guardrail Post-Call Hook: Checking response")
865 if not self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.post_call):
866 verbose_proxy_logger.debug("Tool Permission Guardrail: Skipping check (not enabled)")
867 return response
869 # Extract tool_calls from the response
870 tool_calls: Final = (
871 self._extract_tool_calls_from_response(response)
872 if isinstance(response, ModelResponse)
873 else self._extract_tool_calls_from_anthropic_content(anthropic_content or ())
874 )
876 if not tool_calls:
877 verbose_proxy_logger.debug("Tool Permission Guardrail: No tool uses found")
878 return response
880 verbose_proxy_logger.debug("Tool Permission Guardrail: Found %s tool calls", len(tool_calls))
882 denied_tools: Final = self._evaluate_tool_calls(tool_calls)
884 if not denied_tools:
885 verbose_proxy_logger.debug("Tool Permission Guardrail Post-Call Hook: All tools allowed")
886 elif isinstance(response, ModelResponse):
887 self._modify_response_with_permission_errors(response, denied_tools)
888 else:
889 self._modify_anthropic_content_with_permission_errors(response, anthropic_content or (), denied_tools)
891 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
892 return response
894 async def async_post_call_streaming_iterator_hook(
895 self,
896 user_api_key_dict: UserAPIKeyAuth,
897 response: AsyncIterable[ModelResponseStream],
898 request_data: dict,
899 ) -> AsyncGenerator[ModelResponseStream, None]:
900 """
901 Check tool usage permissions after the LLM stream call
903 Args:
904 user_api_key_dict: User API key information (unused but required by interface)
905 response: The model response to check
906 request_data: The model request (unused but required by interface)
907 """
909 # Import here to avoid circular imports
910 from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
911 from litellm.main import stream_chunk_builder
912 from litellm.types.utils import TextCompletionResponse
914 # Collect all chunks to process them together
915 all_chunks: Final[list[ModelResponseStream]] = []
916 async for chunk in response:
917 all_chunks.append(chunk)
919 assembled_model_response: Final[ModelResponse | TextCompletionResponse | None] = (
920 stream_chunk_builder(chunks=all_chunks) if not is_raw_sse_stream(all_chunks) else None
921 )
922 if isinstance(assembled_model_response, ModelResponse):
923 denied_tools = self._check_assembled_stream(assembled_model_response)
924 if denied_tools:
925 self._modify_response_with_permission_errors(assembled_model_response, denied_tools)
927 mock_response: Final = MockResponseIterator(model_response=assembled_model_response)
928 # Return the reconstructed stream
929 async for chunk in mock_response:
930 yield chunk
931 return
933 anthropic_response: Final = assemble_anthropic_sse_stream(all_chunks)
934 if anthropic_response is None:
935 if is_raw_sse_stream(all_chunks):
936 raise GuardrailRaisedException(
937 guardrail_name=self.guardrail_name,
938 message=(
939 "Streamed response could not be verified for tool permissions "
940 "(not a parseable Anthropic SSE stream), blocking it"
941 ),
942 )
943 for chunk in all_chunks:
944 yield chunk
945 return
947 anthropic_denials: Final = self._check_assembled_stream(anthropic_response)
948 if not anthropic_denials:
949 for chunk in all_chunks:
950 yield chunk
951 return
953 self._modify_response_with_permission_errors(anthropic_response, anthropic_denials)
954 for sse_chunk in anthropic_sse_chunks_from_response(anthropic_response):
955 yield sse_chunk
957 def _check_assembled_stream(
958 self, assembled: ModelResponse
959 ) -> tuple[tuple[ChatCompletionMessageToolCall, PermissionError], ...]:
960 verbose_proxy_logger.debug("Tool Permission Guardrail: Checking response")
961 tool_calls: Final = self._extract_tool_calls_from_response(assembled)
962 if not tool_calls:
963 verbose_proxy_logger.debug("Tool Permission Guardrail: No tool uses found")
964 return ()
965 verbose_proxy_logger.debug("Tool Permission Guardrail: Found %s tool calls", len(tool_calls))
966 denied_tools: Final = self._evaluate_tool_calls(tool_calls)
967 if not denied_tools:
968 verbose_proxy_logger.debug("Tool Permission Guardrail Post-Call Hook: All tools allowed")
969 return denied_tools