Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions docs/specs/004-python-function-calling-loop.md
Original file line number Diff line number Diff line change
Expand Up @@ -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` |
Expand Down
63 changes: 59 additions & 4 deletions python/packages/core/agent_framework/_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Copy link
Copy Markdown
Contributor

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 mode is required: after an allowed call executes, both loops clear the entire tool_choice, so a later model turn can dispatch any registered tool. That defeats the local authorization boundary for providers that do not enforce allowed_tools. Once the required-call obligation is satisfied, preserve the allowlist by transitioning to auto mode rather than dropping it.

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, ...]]:
Expand Down Expand Up @@ -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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

When a provider returns only calls excluded by required_function_name, this rejects the batch with zero executions, but both outer loops still clear required mode before the next model turn. A provider that ignores the restriction can then return the same forbidden tool again and it executes locally without the policy. Keep required mode active until at least one permitted call actually executes in both streaming and non-streaming paths.

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

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The 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? _try_execute_function_call_groups returns only the visible approval group, but this comprehension consumes one group for every allowed call, so the deferred sibling hits StopIteration and both streaming and non-streaming requests fail before returning the approval prompt. Could the reassembly key result groups by input occurrence, or preserve an explicit slot for every allowed 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,
Expand Down
149 changes: 145 additions & 4 deletions python/packages/core/tests/core/test_function_invocation_logic.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
ResponseInvalidatedException,
ResponseStream,
SupportsChatGetResponse,
ToolMode,
chat_middleware,
tool,
)
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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]
Expand All @@ -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]
Expand Down Expand Up @@ -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(
Expand All @@ -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},
)

Expand Down
Loading