diff --git a/docs/specs/004-python-function-calling-loop.md b/docs/specs/004-python-function-calling-loop.md index bdf2fc8ac23..1f734c71d16 100644 --- a/docs/specs/004-python-function-calling-loop.md +++ b/docs/specs/004-python-function-calling-loop.md @@ -582,6 +582,7 @@ that manually replay messages own the equivalent rule: do not resend an approval | String input | Flexible string input follows the same loop behavior. | `test_base_client_with_function_calling_string_input` | | Multiple sequential rounds | Each round retains one call/result pair. | `test_base_client_with_function_calling_resets` | | Streaming call | Call chunks, one result update, and final text are emitted in order. | `test_base_client_with_streaming_function_calling` | +| Tool-selection policy | Fresh model calls execute locally only when permitted by `tool_choice`; `none`, `required_function_name`, and `allowed_tools` are enforced again before local dispatch in both response modes even when a provider does not enforce them. Rejected calls retain correlated error results without consuming the executed-function budget. Valid session-bound approval resumes retain their recorded authority. | `test_fresh_function_dispatch_enforces_tool_choice_policy`, `test_tool_choice_rejection_does_not_consume_function_call_budget`, `test_local_approval_response_executes_with_authoritative_session`, `test_streaming_approval_resume_yields_terminal_result_before_model_text` | | Function-call occurrence identity | Actionable calls gain one stable `Content.id`; a safe local empty-`call_id` fallback uses that id with a migration warning, and streaming aggregation preserves provider-assigned occurrence ids across interleaved fragments. OpenAI Chat Completions scopes fragment correlation to each request and `(choice.index, tool.index)`. | `test_actionable_function_call_gets_stable_occurrence_identity`, `test_actionable_function_call_uses_occurrence_identity_for_empty_call_id`, `test_streaming_empty_call_id_keeps_occurrence_identity_through_approval`, `test_streaming_empty_call_id_delta_reuses_opening_call_identity`, `test_streaming_interleaved_indexed_call_fragments_coalesce_by_occurrence`, `packages/core/tests/core/test_types.py::test_function_call_occurrence_id_roundtrips_without_regeneration`, `packages/openai/tests/openai/test_openai_chat_completion_client.py::test_streaming_tool_call_identity_is_request_local_and_scoped_by_choice_index` | | Reasoning-bound call | Finalized output retains reasoning, function call, function result, and final text. | `test_streaming_function_calling_response_includes_reasoning_and_tool_results` | | Calls across response messages | Every actionable call is executed once. | `test_base_client_executes_function_calls_across_multiple_response_messages` | diff --git a/python/packages/core/agent_framework/_tools.py b/python/packages/core/agent_framework/_tools.py index 4ab15a52f79..a8bc3f17995 100644 --- a/python/packages/core/agent_framework/_tools.py +++ b/python/packages/core/agent_framework/_tools.py @@ -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) + 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 + ] # 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, diff --git a/python/packages/core/tests/core/test_function_invocation_logic.py b/python/packages/core/tests/core/test_function_invocation_logic.py index 243dd0dcff0..32da61a4b44 100644 --- a/python/packages/core/tests/core/test_function_invocation_logic.py +++ b/python/packages/core/tests/core/test_function_invocation_logic.py @@ -28,6 +28,7 @@ ResponseInvalidatedException, ResponseStream, SupportsChatGetResponse, + ToolMode, chat_middleware, tool, ) @@ -1005,6 +1006,140 @@ def ai_func(arg1: str) -> str: assert response.messages[2].text == "done" +@pytest.mark.parametrize("streaming", [False, True], ids=["non_streaming", "streaming"]) +@pytest.mark.parametrize( + ("tool_choice", "expected_executions"), + [ + ({"mode": "auto", "allowed_tools": ["allowed_tool"]}, ["allowed"]), + ({"mode": "required", "required_function_name": "allowed_tool"}, ["allowed"]), + ({"mode": "none"}, []), + ], + ids=["allowed_tools", "required_function", "none"], +) +async def test_fresh_function_dispatch_enforces_tool_choice_policy( + chat_client_base: SupportsChatGetResponse, + streaming: bool, + tool_choice: ToolMode, + expected_executions: list[str], +) -> None: + """A provider response cannot dispatch local functions excluded by tool_choice.""" + executions: list[str] = [] + + @tool(name="allowed_tool", approval_mode="never_require") + def allowed_tool() -> str: + executions.append("allowed") + return "allowed" + + @tool(name="blocked_tool", approval_mode="never_require") + def blocked_tool() -> str: + executions.append("blocked") + return "blocked" + + calls = [ + Content.from_function_call(call_id="allowed", name="allowed_tool", arguments={}), + Content.from_function_call(call_id="blocked", name="blocked_tool", arguments={}), + ] + options: ChatOptions = {"tools": [allowed_tool, blocked_tool], "tool_choice": tool_choice} + if streaming: + chat_client_base.streaming_responses = [ # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + [ChatResponseUpdate(role="assistant", contents=calls)], + [ChatResponseUpdate(role="assistant", contents=[Content.from_text("done")])], + ] + stream = chat_client_base.get_response( + [Message(role="user", contents=["run tools"])], + options=options, + stream=True, + ) + response = await stream.get_final_response() + else: + chat_client_base.run_responses = [ # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + ChatResponse(messages=Message(role="assistant", contents=calls)), + ChatResponse(messages=Message(role="assistant", contents=["done"])), + ] + response = await chat_client_base.get_response( + [Message(role="user", contents=["run tools"])], + options=options, + ) + + assert executions == expected_executions + results = { + content.call_id: content + for message in response.messages + for content in message.contents + if content.type == "function_result" + } + assert results["blocked"].exception == "FunctionInvocationPolicyError" + assert results["allowed"].exception == (None if expected_executions else "FunctionInvocationPolicyError") + + +@pytest.mark.parametrize("streaming", [False, True], ids=["non_streaming", "streaming"]) +async def test_tool_choice_rejection_does_not_consume_function_call_budget( + chat_client_base: SupportsChatGetResponse, + streaming: bool, +) -> None: + """A policy-rejected call does not prevent a later allowed call from using the execution budget.""" + from agent_framework._tools import _FUNCTION_INVOCATION_BUDGET_STATE_KEY + + executions: list[str] = [] + + @tool(name="allowed_tool", approval_mode="never_require") + def allowed_tool() -> str: + executions.append("allowed") + return "allowed" + + @tool(name="blocked_tool", approval_mode="never_require") + def blocked_tool() -> str: + executions.append("blocked") + return "blocked" + + blocked_call = Content.from_function_call(call_id="blocked", name="blocked_tool", arguments={}) + allowed_call = Content.from_function_call(call_id="allowed", name="allowed_tool", arguments={}) + if streaming: + chat_client_base.streaming_responses = [ # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + [ChatResponseUpdate(role="assistant", contents=[blocked_call])], + [ChatResponseUpdate(role="assistant", contents=[allowed_call])], + [ChatResponseUpdate(role="assistant", contents=[Content.from_text("done")])], + ] + else: + chat_client_base.run_responses = [ # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + ChatResponse(messages=Message(role="assistant", contents=[blocked_call])), + ChatResponse(messages=Message(role="assistant", contents=[allowed_call])), + ChatResponse(messages=Message(role="assistant", contents=["done"])), + ] + + budget_state: dict[str, int] = {} + chat_client_base.function_invocation_configuration["max_function_calls"] = 1 # type: ignore[attr-defined] # ty: ignore[unresolved-attribute] + options: ChatOptions = { + "tools": [allowed_tool, blocked_tool], + "tool_choice": {"mode": "auto", "allowed_tools": ["allowed_tool"]}, + } + if streaming: + stream = chat_client_base.get_response( + [Message(role="user", contents=["run tools"])], + options=options, + stream=True, + client_kwargs={_FUNCTION_INVOCATION_BUDGET_STATE_KEY: budget_state}, + ) + response = await stream.get_final_response() + else: + response = await chat_client_base.get_response( + [Message(role="user", contents=["run tools"])], + options=options, + client_kwargs={_FUNCTION_INVOCATION_BUDGET_STATE_KEY: budget_state}, + ) + + assert executions == ["allowed"] + assert budget_state["total_function_calls"] == 1 + results = { + content.call_id: content + for message in response.messages + for content in message.contents + if content.type == "function_result" + } + assert results["blocked"].exception == "FunctionInvocationPolicyError" + assert results["allowed"].exception is None + + async def test_function_call_with_length_finish_reason_executes_and_continues( chat_client_base: SupportsChatGetResponse, ) -> None: @@ -8663,7 +8798,10 @@ def guarded_stream_tool() -> str: first_stream = chat_client_base.get_response( [Message(role="user", contents=["run guarded"])], stream=True, - options={"tools": [guarded_stream_tool]}, + options={ + "tools": [guarded_stream_tool], + "tool_choice": {"mode": "auto", "allowed_tools": ["guarded_stream_tool"]}, + }, client_kwargs={"session": session}, ) first_updates = [update async for update in first_stream] @@ -8677,7 +8815,10 @@ def guarded_stream_tool() -> str: resumed_stream = chat_client_base.get_response( [Message(role="user", contents=[approval_request.to_function_approval_response(approved=approved)])], stream=True, - options={"tools": [guarded_stream_tool]}, + options={ + "tools": [guarded_stream_tool], + "tool_choice": {"mode": "auto", "allowed_tools": ["guarded_stream_tool"]}, + }, client_kwargs={"session": session}, ) resumed_updates = [update async for update in resumed_stream] @@ -11029,7 +11170,7 @@ async def test_local_approval_response_executes_with_authoritative_session( paused = await chat_client_base.get_response( [Message(role="user", contents=["please continue"])], - options={"tools": [guarded_tool]}, + options={"tools": [guarded_tool], "tool_choice": {"mode": "auto", "allowed_tools": ["guarded_tool"]}}, client_kwargs={"session": session}, ) approval_request = next( @@ -11044,7 +11185,7 @@ async def test_local_approval_response_executes_with_authoritative_session( *paused.messages, Message(role="user", contents=[approval_request.to_function_approval_response(approved=True)]), ], - options={"tools": [guarded_tool]}, + options={"tools": [guarded_tool], "tool_choice": {"mode": "auto", "allowed_tools": ["guarded_tool"]}}, client_kwargs={"session": session}, )