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

1import json 

2import re 

3from collections.abc import AsyncGenerator, AsyncIterable, Mapping, Sequence 

4from typing import Any, Final, Literal, TypedDict 

5 

6from fastapi import HTTPException 

7from typing_extensions import ReadOnly, Required 

8 

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) 

41 

42GUARDRAIL_NAME: Final = "tool_permission" 

43 

44 

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 

48 

49 

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 

53 

54 

55class _ToolPermissionRuleFields(TypedDict, total=False): 

56 """The config-file shape a :class:`ToolPermissionRule` is built from.""" 

57 

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] 

63 

64 

65def _rule_from_fields(fields: _ToolPermissionRuleFields) -> ToolPermissionRule: 

66 """Validate one config-file rule entry into a :class:`ToolPermissionRule`.""" 

67 return ToolPermissionRule(**fields) 

68 

69 

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" 

74 

75 

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 

86 

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())) 

95 

96 super().__init__(**kwargs) 

97 

98 self._load_rules(rules) 

99 

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 ) 

105 

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 ) 

111 

112 def _load_rules(self, rules: list[Any] | None) -> None: 

113 """Parse ``rules`` and (re)build the compiled target/pattern lookups. 

114 

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]]] = {} 

124 

125 for rule_item in rules or []: 

126 rule = rule_item if isinstance(rule_item, ToolPermissionRule) else _rule_from_fields(rule_item) 

127 

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 

142 

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 

149 

150 parsed_rules.append(rule) 

151 compiled_targets[rule.id] = target_patterns 

152 if rule_patterns: 

153 compiled_patterns[rule.id] = rule_patterns 

154 

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 

161 

162 def update_in_memory_litellm_params(self, litellm_params: LitellmParams | dict) -> None: 

163 """Apply updated params in place, rebuilding the compiled rule state. 

164 

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) 

183 

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() 

208 

209 @staticmethod 

210 def get_config_model(): 

211 from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( 

212 ToolPermissionGuardrailConfigModel, 

213 ) 

214 

215 return ToolPermissionGuardrailConfigModel 

216 

217 @classmethod 

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

219 return [ 

220 GuardrailEventHooks.pre_call, 

221 GuardrailEventHooks.post_call, 

222 ] 

223 

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)) 

230 

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") 

241 

242 name_required: Final = rule.tool_name is not None 

243 type_required: Final = rule.tool_type is not None 

244 

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 

247 

248 overall_match: Final = name_matched and type_matched 

249 should_check_params: Final = name_required and name_matched 

250 

251 return overall_match, should_check_params 

252 

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 

260 

261 Args: 

262 tool_name: Name of the tool to check 

263 tool_type: Type of the tool to check 

264 

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) 

269 

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 

292 

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 

306 

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" 

313 

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" 

329 

330 if isinstance(parsed_arguments, dict): 

331 return parsed_arguments, None 

332 

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" 

338 

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 

347 

348 if depth > DEFAULT_MAX_RECURSE_DEPTH: 

349 return 

350 

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) 

365 

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 

376 

377 path_value_map: Final[dict[str, list[object]]] = {} 

378 self._collect_argument_paths(arguments, "", path_value_map) 

379 

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 ) 

394 

395 return True, None 

396 

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 

404 

405 tool_identifier: Final = tool_name or tool_type or "unknown_tool" 

406 

407 last_pattern_failure_msg: str | None = None 

408 

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 

417 

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 

434 

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 

443 

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 

451 

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 

463 

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) 

469 

470 @staticmethod 

471 def _legacy_function_call_id(choice_index: int) -> str: 

472 return f"legacy_function_call_{choice_index}" 

473 

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 

479 

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 

484 

485 return ChatCompletionMessageToolCall( 

486 id=self._legacy_function_call_id(choice_index), 

487 type="function", 

488 function={"name": function_name, "arguments": arguments}, 

489 ) 

490 

491 def _extract_tool_calls_from_response(self, response: ModelResponse) -> list[ChatCompletionMessageToolCall]: 

492 """ 

493 Extract tool_calls from all choices in a model response. 

494 

495 Args: 

496 response: The model response to analyze 

497 

498 Returns: 

499 List of tool_calls blocks found in the response 

500 """ 

501 tool_calls: Final = [] 

502 

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) 

512 

513 return tool_calls 

514 

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 ) 

528 

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 

535 

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 ) 

542 

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) 

547 

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 ) 

555 

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 ) 

570 

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 

579 

580 verbose_proxy_logger.info("Blocking %s unauthorized tool uses", len(denied_tools)) 

581 

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 } 

588 

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")) 

594 

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 ] 

603 

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 

607 

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 

612 

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 

616 

617 def _get_legacy_function_name(self, function: object) -> str | None: 

618 return self._get_mapping_value(function, "name") 

619 

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") 

629 

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") 

637 

638 def _collect_request_tools(self, data: dict) -> list[tuple[str, str | None]]: 

639 request_tools: Final[list[tuple[str, str | None]]] = [] 

640 

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)) 

645 

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")) 

650 

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")) 

657 

658 return request_tools 

659 

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 

667 

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 

674 

675 verbose_proxy_logger.info("Blocking %s unauthorized tool uses", len(denied_tool_names)) 

676 

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) 

681 

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 

691 

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 ] 

697 

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" 

701 

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" 

705 

706 return data 

707 

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 

713 

714 Args: 

715 tool_use: The tool use that was denied 

716 error: The permission error details 

717 

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})" 

724 

725 return ToolResult(tool_use_id=tool_call.id, content=error_message, is_error=True) 

726 

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 

734 

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 

741 

742 verbose_proxy_logger.info("Blocking %s unauthorized tool uses", len(denied_tools)) 

743 

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 

749 

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 = [] 

755 

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) 

764 

765 choice.message.tool_calls = filtered_tool_calls if filtered_tool_calls else None 

766 

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) 

775 

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) 

783 

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" 

790 

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") 

801 

802 from litellm.proxy.common_utils.callback_utils import ( 

803 add_guardrail_to_applied_guardrails_header, 

804 ) 

805 

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 

809 

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 

816 

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) 

821 

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) 

833 

834 if denied_tool_names: 

835 data = self._modify_request_with_permission_errors(data, denied_tool_names) 

836 

837 verbose_proxy_logger.debug("Tool Permission Guardrail Pre-Call Hook: All tools allowed") 

838 

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

840 return data 

841 

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 

851 

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 

862 

863 verbose_proxy_logger.debug("Tool Permission Guardrail Post-Call Hook: Checking response") 

864 

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 

868 

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 ) 

875 

876 if not tool_calls: 

877 verbose_proxy_logger.debug("Tool Permission Guardrail: No tool uses found") 

878 return response 

879 

880 verbose_proxy_logger.debug("Tool Permission Guardrail: Found %s tool calls", len(tool_calls)) 

881 

882 denied_tools: Final = self._evaluate_tool_calls(tool_calls) 

883 

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) 

890 

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

892 return response 

893 

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 

902 

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 """ 

908 

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 

913 

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) 

918 

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) 

926 

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 

932 

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 

946 

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 

952 

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 

956 

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