-
Notifications
You must be signed in to change notification settings - Fork 2.4k
Python: Enforce tool selection before dispatch #8852
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -4395,6 +4395,45 @@ class _FunctionProcessingResult: | |
| _FunctionCallExecutor: TypeAlias = Callable[..., Awaitable[_FunctionExecutionBatch]] | ||
|
|
||
|
|
||
| def _fresh_function_calls_allowed_by_tool_choice( | ||
| function_calls: Sequence[Content], | ||
| options: dict[str, Any] | None, | ||
| ) -> tuple[list[Content], set[int]]: | ||
| """Return fresh calls permitted for local dispatch and the blocked object identities.""" | ||
| from ._types import validate_tool_mode | ||
|
|
||
| tool_mode = validate_tool_mode(options.get("tool_choice")) if options else None | ||
| if tool_mode is None: | ||
| return list(function_calls), set() | ||
|
|
||
| allowed_names: set[str] | None = None | ||
| if tool_mode.get("mode") == "none": | ||
| allowed_names = set() | ||
| elif required_name := tool_mode.get("required_function_name"): | ||
| allowed_names = {required_name} | ||
| elif (configured_names := tool_mode.get("allowed_tools")) is not None: | ||
| allowed_names = set(configured_names) | ||
|
|
||
| if allowed_names is None: | ||
| return list(function_calls), set() | ||
|
|
||
| # Provider-side selection is advisory for some protocols. Recheck fresh model | ||
| # output here so an ignored policy can never become a local side effect. | ||
| allowed_calls = [call for call in function_calls if call.name in allowed_names] | ||
| allowed_call_ids = {id(call) for call in allowed_calls} | ||
| return allowed_calls, {id(call) for call in function_calls if id(call) not in allowed_call_ids} | ||
|
|
||
|
|
||
| def _tool_choice_rejection_result(function_call: Content) -> Content: | ||
| from ._types import Content | ||
|
|
||
| return Content.from_function_result( | ||
| call_id=function_call.call_id or "", | ||
| result="Error: The requested function is not permitted by the active tool_choice policy.", | ||
| exception="FunctionInvocationPolicyError", | ||
| ) | ||
|
|
||
|
|
||
| def _messages_and_updates_for_terminal_contents( | ||
| contents: Sequence[Content], | ||
| ) -> tuple[tuple[Message, ...], tuple[ChatResponseUpdate, ...]]: | ||
|
|
@@ -4830,16 +4869,32 @@ async def _process_model_function_calls( | |
| return _FunctionProcessingResult(errors_in_a_row=errors_in_a_row, action="return") | ||
|
|
||
| # 2. Execute the batch once while preserving each call's result group. | ||
| execution = await execute_function_calls( | ||
| function_calls=function_calls, | ||
| options=options, | ||
| allowed_calls, blocked_call_ids = _fresh_function_calls_allowed_by_tool_choice(function_calls, options) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. When a provider returns only calls excluded by |
||
| execution = ( | ||
| await execute_function_calls( | ||
| function_calls=allowed_calls, | ||
| options=options, | ||
| ) | ||
| if allowed_calls | ||
| else _FunctionExecutionBatch(result_groups=[]) | ||
| ) | ||
| # Synthetic policy failures preserve call/result balance but did not run a tool | ||
| # body, so capture the real execution count before merging those result groups. | ||
| executed_call_count = execution.executed_call_count | ||
| if blocked_call_ids: | ||
| allowed_result_groups = iter(execution.result_groups) | ||
| execution.result_groups = [ | ||
| [_tool_choice_rejection_result(function_call)] | ||
| if id(function_call) in blocked_call_ids | ||
| else next(allowed_result_groups) | ||
| for function_call in function_calls | ||
|
Comment on lines
+4885
to
+4890
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. What happens when the allowed batch contains an approval-required call plus a session-deferred executable sibling and also includes a blocked call? |
||
| ] | ||
|
|
||
| # 3. Fold results into the response and translate errors or middleware termination into the next loop action. | ||
| processing_result = _handle_function_call_results( | ||
| response=response, | ||
| execution_results=execution.contents, | ||
| function_call_count=execution.executed_call_count, | ||
| function_call_count=executed_call_count, | ||
| function_call_messages=function_call_messages, | ||
| errors_in_a_row=errors_in_a_row, | ||
| had_errors=execution.had_errors, | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
This allowlist is enforced only for the current round when
modeisrequired: after an allowed call executes, both loops clear the entiretool_choice, so a later model turn can dispatch any registered tool. That defeats the local authorization boundary for providers that do not enforceallowed_tools. Once the required-call obligation is satisfied, preserve the allowlist by transitioning to auto mode rather than dropping it.