diff --git a/docs/specs/004-python-function-calling-loop.md b/docs/specs/004-python-function-calling-loop.md index d7d4605818..d1b11c11c0 100644 --- a/docs/specs/004-python-function-calling-loop.md +++ b/docs/specs/004-python-function-calling-loop.md @@ -427,6 +427,7 @@ that manually replay messages own the equivalent rule: do not resend an approval `ResponseStream.get_final_response()`. - The function invocation layer normalizes a private copy of caller messages. It must not mutate the caller's approval `Message`, approval `Content`, or an earlier returned response. + - Approval-time `UserInputRequiredException` and `MiddlewareTermination` return immediately without another model call. @@ -487,6 +488,7 @@ that manually replay messages own the equivalent rule: do not resend an approval | Function invocation disabled | The client bypasses the invocation loop without losing invocation kwargs. | `test_function_invocation_config_enabled_false`, `test_function_invocation_config_enabled_false_preserves_invocation_kwargs`, `test_streaming_function_invocation_config_enabled_false` | | Runtime tool changes | Added tools become available on the next iteration and retain approval behavior. | `test_add_tools_available_next_iteration`, `test_add_tools_with_approval_required_tool` | + ### Approval pause and resume | Scenario | Required invariant | Primary regression test | @@ -511,7 +513,8 @@ that manually replay messages own the equivalent rule: do not resend an approval | Truthy non-boolean decision | Strings, integers, null, and other non-booleans do not authorize execution. | `packages/core/tests/core/test_function_invocation_logic.py::test_session_approval_binding_treats_truthy_non_boolean_as_rejection`, `packages/core/tests/core/test_types.py::test_function_approval_response_deserialization_rejects_non_boolean_decisions`, `packages/ag-ui/tests/ag_ui/test_message_adapters.py::test_function_approval_requires_real_boolean`, `packages/ag-ui/tests/ag_ui/test_approval_result_event.py::test_resolve_approval_responses_treats_non_boolean_decision_as_rejection` | | Active batch replacement | A newly surfaced model batch replaces abandoned approval authority instead of growing session state. | `packages/core/tests/core/test_function_invocation_logic.py::test_session_approval_binding_replaces_abandoned_batch` | | Duplicate request id | Ambiguous request IDs within one active batch fail explicitly. | `packages/core/tests/core/test_function_invocation_logic.py::test_session_approval_batch_rejects_duplicate_request_ids` | -| Tool registry changes | Same-name upgrades may execute the recorded operation; removing the recorded name executes nothing. | `packages/core/tests/core/test_harness_tool_approval.py::test_approval_resume_allows_same_name_tool_upgrade`, `test_approval_resume_does_not_execute_when_recorded_tool_disappears` | + +| Tool registry changes | Same-name upgrades may execute the recorded operation; removing the recorded name executes nothing. | `packages/core/tests/core/test_harness_tool_approval.py::test_approval_resume_allows_same_name_tool_upgrade`, `test_approval_resume_does_not_execute_when_recorded_tool_disappears` | ### Approval correlation and replay diff --git a/python/packages/core/agent_framework/_sessions.py b/python/packages/core/agent_framework/_sessions.py index 27e313b8e6..42817525a5 100644 --- a/python/packages/core/agent_framework/_sessions.py +++ b/python/packages/core/agent_framework/_sessions.py @@ -29,7 +29,6 @@ import weakref from abc import abstractmethod from base64 import urlsafe_b64encode -from collections import deque from collections.abc import AsyncIterable, Awaitable, Callable, Generator, Iterable, Mapping, Sequence from contextvars import ContextVar, Token from dataclasses import dataclass @@ -876,46 +875,71 @@ def _is_approval_placeholder_result(content: Content) -> bool: def _approval_controls_to_keep(messages: Sequence[Message]) -> set[int]: - unresolved_requests_by_id: dict[str, Content] = {} - unresolved_local_responses_by_id: dict[str, Content] = {} - local_response_ids_by_call_id: dict[str, deque[str]] = {} + request_positions: dict[str, tuple[int, int]] = {} + response_ids: set[str] = set() + resolving_events: list[tuple[int, int, str]] = [] - for message in messages: - for content in message.contents: + for msg_idx, message in enumerate(messages): + for content_idx, content in enumerate(message.contents): if content.type == "function_approval_request": - function_call = content.function_call - if content.id is not None and function_call is not None and function_call.call_id is not None: - unresolved_requests_by_id.setdefault(content.id, content) - continue - if content.type == "function_approval_response": - function_call = content.function_call if content.id is not None: - unresolved_requests_by_id.pop(content.id, None) - if ( - content.id is not None - and function_call is not None - and function_call.call_id is not None - and not function_call.additional_properties.get("server_label") - and content.id not in unresolved_local_responses_by_id - ): - unresolved_local_responses_by_id[content.id] = content - local_response_ids_by_call_id.setdefault(function_call.call_id, deque()).append(content.id) - continue - if content.call_id is None: - continue - is_terminal_result = content.type == "function_result" and not _is_approval_placeholder_result(content) - is_follow_up_request = content.user_input_request and content.type not in { - "function_approval_request", - "function_approval_response", - } - if not (is_terminal_result or is_follow_up_request): - continue - if response_ids := local_response_ids_by_call_id.get(content.call_id): - unresolved_local_responses_by_id.pop(response_ids.popleft(), None) + request_positions[content.id] = (msg_idx, content_idx) + elif content.type == "function_approval_response": + if content.id is not None: + response_ids.add(content.id) + elif content.call_id is not None: + is_terminal_result = content.type == "function_result" and not _is_approval_placeholder_result(content) + is_follow_up_request = content.user_input_request and content.type not in { + "function_approval_request", + "function_approval_response", + } + if is_terminal_result or is_follow_up_request: + resolving_events.append((msg_idx, content_idx, content.call_id)) + + keep_ids: set[int] = set() + seen_request_ids: set[str] = set() + + for msg_idx, message in enumerate(messages): + for content_idx, content in enumerate(message.contents): + if content.type == "function_approval_request": + if content.id is None or content.function_call is None or content.function_call.call_id is None: + continue + if content.id in seen_request_ids: + continue + + req_pos = (msg_idx, content_idx) + call_id = content.function_call.call_id + + is_resolved = content.id in response_ids + if not is_resolved: + for res_msg_idx, res_content_idx, res_call_id in resolving_events: + if res_call_id == call_id and (res_msg_idx, res_content_idx) >= req_pos: + is_resolved = True + break + if not is_resolved: + keep_ids.add(id(content)) + seen_request_ids.add(content.id) + + elif content.type == "function_approval_response": + function_call = content.function_call + if content.id is None or function_call is None or function_call.call_id is None: + continue + if function_call.additional_properties.get("server_label"): + continue + + call_id = function_call.call_id + resp_pos = (msg_idx, content_idx) + ref_pos = request_positions.get(content.id, resp_pos) + + is_resolved = False + for res_msg_idx, res_content_idx, res_call_id in resolving_events: + if res_call_id == call_id and (res_msg_idx, res_content_idx) >= ref_pos: + is_resolved = True + break + if not is_resolved: + keep_ids.add(id(content)) - return { - id(content) for content in (*unresolved_requests_by_id.values(), *unresolved_local_responses_by_id.values()) - } + return keep_ids def _filter_approval_control_messages(messages: Sequence[Message]) -> list[Message]: diff --git a/python/packages/core/agent_framework/_tools.py b/python/packages/core/agent_framework/_tools.py index 9b40754760..e5767a4de8 100644 --- a/python/packages/core/agent_framework/_tools.py +++ b/python/packages/core/agent_framework/_tools.py @@ -2366,93 +2366,149 @@ def _pop_already_approved_approval_responses( return responses -def _collect_approval_responses( - messages: list[Message], -) -> dict[str, Content]: - """Collect approval responses (both approved and rejected) from messages. +def _collect_approval_responses(messages: list[Message]) -> dict[str, Content]: + requests: list[tuple[int, int, str]] = [] + resolving_events: list[tuple[int, int, str]] = [] - Hosted tool approvals (e.g. MCP) are excluded because they must be - forwarded to the API as-is rather than processed locally. - """ - approval_responses: list[Content] = [] - pending_by_call_id: dict[str, deque[Content]] = {} - resolved_response_ids: set[int] = set() - for message in messages: - for content in message.contents: + for msg_idx, message in enumerate(messages): + for content_idx, content in enumerate(message.contents): + if content.type == "function_approval_request": + if content.function_call is not None and content.function_call.call_id is not None: + requests.append((msg_idx, content_idx, content.function_call.call_id)) + elif content.call_id is not None: + is_terminal_result = content.type == "function_result" and not _is_approval_placeholder_result(content) + is_follow_up_request = content.user_input_request and content.type not in { + "function_approval_request", + "function_approval_response", + } + if is_terminal_result or is_follow_up_request: + resolving_events.append((msg_idx, content_idx, content.call_id)) + + unresolved_responses: dict[str, Content] = {} + for msg_idx, message in enumerate(messages): + for content_idx, content in enumerate(message.contents): if content.type == "function_approval_response" and not _is_hosted_tool_approval(content): - function_call = content.function_call - if function_call is None or function_call.call_id is None: + if content.id is None or content.function_call is None or content.function_call.call_id is None: continue - approval_responses.append(content) - pending_by_call_id.setdefault(function_call.call_id, deque()).append(content) - continue - if content.call_id is None: - continue - is_terminal_result = content.type == "function_result" and not _is_approval_placeholder_result(content) - is_follow_up_request = content.user_input_request and content.type not in { - "function_approval_request", - "function_approval_response", - } - if not (is_terminal_result or is_follow_up_request): - continue - pending_responses = pending_by_call_id.get(content.call_id) - if pending_responses: - resolved_response_ids.add(id(pending_responses.popleft())) - return { - content.id: content - for content in approval_responses - if id(content) not in resolved_response_ids and content.id is not None - } + call_id = content.function_call.call_id + resp_pos = (msg_idx, content_idx) + + latest_req_pos = None + for r_msg, r_cidx, r_call in requests: + if ( + r_call == call_id + and (r_msg, r_cidx) < resp_pos + and (latest_req_pos is None or (r_msg, r_cidx) > latest_req_pos) + ): + latest_req_pos = (r_msg, r_cidx) + + is_resolved = False + for res_msg_idx, res_content_idx, res_call_id in resolving_events: + if res_call_id == call_id: + res_pos = (res_msg_idx, res_content_idx) + if latest_req_pos is not None: + if res_pos > latest_req_pos or res_pos[0] == resp_pos[0]: + is_resolved = True + break + else: + if res_pos >= resp_pos or res_pos[0] == resp_pos[0]: + is_resolved = True + break + + if not is_resolved: + unresolved_responses[content.id] = content + + return unresolved_responses def _collect_unanswered_approval_requests(messages: Sequence[Message]) -> list[Content]: - approval_requests_by_id: dict[str, Content] = {} - pending_request_ids_by_call_id: dict[str, deque[str]] = {} - answered_approval_ids: set[str] = set() + responses: set[str] = set() + resolving_events: list[tuple[int, int, str]] = [] - for message in messages: - for content in message.contents: - if content.type == "function_approval_request": - function_call = content.function_call - if content.id is None or function_call is None or function_call.call_id is None: - continue - if content.id not in approval_requests_by_id: - approval_requests_by_id[content.id] = content - pending_request_ids_by_call_id.setdefault(function_call.call_id, deque()).append(content.id) - continue + for msg_idx, message in enumerate(messages): + for content_idx, content in enumerate(message.contents): if content.type == "function_approval_response": if content.id is not None: - answered_approval_ids.add(content.id) - continue - if content.call_id is None: - continue - is_terminal_result = content.type == "function_result" and not _is_approval_placeholder_result(content) - is_follow_up_request = content.user_input_request and content.type not in { - "function_approval_request", - "function_approval_response", - } - if not (is_terminal_result or is_follow_up_request): - continue - if request_ids := pending_request_ids_by_call_id.get(content.call_id): - answered_approval_ids.add(request_ids.popleft()) + responses.add(content.id) + elif content.call_id is not None: + is_terminal_result = content.type == "function_result" and not _is_approval_placeholder_result(content) + is_follow_up_request = content.user_input_request and content.type not in { + "function_approval_request", + "function_approval_response", + } + if is_terminal_result or is_follow_up_request: + resolving_events.append((msg_idx, content_idx, content.call_id)) + + unanswered_requests: list[Content] = [] + seen_request_ids: set[str] = set() + + for msg_idx, message in enumerate(messages): + for content_idx, content in enumerate(message.contents): + if content.type == "function_approval_request" and not _is_hosted_tool_approval(content): + if content.id is None or content.function_call is None or content.function_call.call_id is None: + continue + if content.id in seen_request_ids: + continue - return [ - request for approval_id, request in approval_requests_by_id.items() if approval_id not in answered_approval_ids - ] + req_pos = (msg_idx, content_idx) + call_id = content.function_call.call_id + + is_answered = content.id in responses + if not is_answered: + for res_msg_idx, res_content_idx, res_call_id in resolving_events: + if res_call_id == call_id and (res_msg_idx, res_content_idx) > req_pos: + is_answered = True + break + + if not is_answered: + unanswered_requests.append(content) + seen_request_ids.add(content.id) + + return unanswered_requests def _remove_unanswered_approval_batches_from_model_input(messages: list[Message]) -> None: pending_requests = _collect_unanswered_approval_requests(messages) - if not pending_requests: - return - pending_approval_ids = {request.id for request in pending_requests if request.id is not None} pending_call_ids = { request.function_call.call_id for request in pending_requests if request.function_call is not None and request.function_call.call_id is not None } + + resolved_response_ids: set[int] = set() + request_positions: dict[str, tuple[int, int]] = {} + resolving_events: list[tuple[int, int, str]] = [] + + for msg_idx, message in enumerate(messages): + for content_idx, content in enumerate(message.contents): + if content.type == "function_approval_request": + if content.id is not None: + request_positions[content.id] = (msg_idx, content_idx) + elif content.call_id is not None: + is_terminal_result = content.type == "function_result" and not _is_approval_placeholder_result(content) + is_follow_up_request = content.user_input_request and content.type not in { + "function_approval_request", + "function_approval_response", + } + if is_terminal_result or is_follow_up_request: + resolving_events.append((msg_idx, content_idx, content.call_id)) + + for msg_idx, message in enumerate(messages): + for content_idx, content in enumerate(message.contents): + if content.type == "function_approval_response" and not _is_hosted_tool_approval(content): + if content.id is None or content.function_call is None or content.function_call.call_id is None: + continue + call_id = content.function_call.call_id + resp_pos = (msg_idx, content_idx) + ref_pos = request_positions.get(content.id, resp_pos) + + for res_msg_idx, res_content_idx, res_call_id in resolving_events: + if res_call_id == call_id and (res_msg_idx, res_content_idx) >= ref_pos: + resolved_response_ids.add(id(content)) + break + open_calls_by_id: dict[str, deque[tuple[Content, int]]] = {} bound_call_content_ids: set[int] = set() call_batch_message_indices: set[int] = set() @@ -2518,11 +2574,14 @@ def _remove_unanswered_approval_batches_from_model_input(messages: list[Message] filtered_messages: list[Message] = [] for message_index, message in enumerate(messages): + if message_index in fully_pending_call_message_indices: + continue filtered_contents = [ content for content in message.contents if not ( (content.type == "function_approval_request" and content.id in pending_approval_ids) + or (id(content) in resolved_response_ids) or ( message_index in call_batch_message_indices and ( @@ -2530,7 +2589,6 @@ def _remove_unanswered_approval_batches_from_model_input(messages: list[Message] or (content.type == "mcp_server_tool_call" and content.call_id in pending_call_ids) ) ) - or (message_index in fully_pending_call_message_indices and content.type == "text_reasoning") ) ] if not filtered_contents: @@ -2567,9 +2625,7 @@ def _replace_approval_contents_with_results( Returns: The terminal contents produced while resolving the approval responses, in response order. """ - from ._types import ( - Content, - ) + from ._types import Content result_groups_by_call_id: dict[str, deque[list[Content]]] = {} for result_group in approved_function_result_groups: @@ -3561,7 +3617,6 @@ async def settle_approval_replay_calls(function_calls: Sequence[Content]) -> Non if fallback_added: yield _function_invocation_limit_fallback_update() return - try: function_processing = await _process_model_function_calls( response=response, 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 5a2ba3e719..ce02134f56 100644 --- a/python/packages/core/tests/core/test_function_invocation_logic.py +++ b/python/packages/core/tests/core/test_function_invocation_logic.py @@ -180,11 +180,13 @@ def test_session_approval_binding_replaces_abandoned_batch() -> None: ) session = AgentSession(session_id="approval-binding-active-batch") + old_call = Content.from_function_call(call_id="call_old", name="guarded_write", arguments={}, id="request_old") old_request = Content.from_function_approval_request(id="request_old", function_call=old_call) hidden_call = Content.from_function_call(call_id="call_hidden", name="safe_read", arguments={}, id="request_hidden") hidden_request = Content.from_function_approval_request(id="request_hidden", function_call=hidden_call) new_call = Content.from_function_call(call_id="call_new", name="guarded_write", arguments={}, id="request_new") + new_request = Content.from_function_approval_request(id="request_new", function_call=new_call) _store_already_approved_approval_requests(session, [old_request], [hidden_request]) @@ -3239,7 +3241,7 @@ def test_pending_approval_batch_filter_keeps_resolved_sibling_pair() -> None: _remove_unanswered_approval_batches_from_model_input(messages) assert [(message.role, message.contents) for message in messages] == [ - ("assistant", [local_call]), + ("assistant", [local_call, hosted_call, hosted_request]), ("tool", [local_result]), ("user", [unrelated_content]), ] diff --git a/python/packages/core/tests/core/test_sessions.py b/python/packages/core/tests/core/test_sessions.py index 138e8a0fc1..9682edbb9d 100644 --- a/python/packages/core/tests/core/test_sessions.py +++ b/python/packages/core/tests/core/test_sessions.py @@ -448,6 +448,85 @@ def test_filter_approval_controls_keeps_response_for_pending_placeholder() -> No assert any(placeholder in message.contents for message in filtered) +def test_filter_approval_controls_resolves_when_result_precedes_response() -> None: + """when a terminal function_result appears BEFORE the approval response + + (the 'approval-resume' layout), the two-pass approach must still + correctly identify the approval as resolved and filter it out. + """ + function_call = Content.from_function_call(call_id="call_resume", name="guarded", arguments="{}") + request = Content.from_function_approval_request(id="approval_resume", function_call=function_call) + response = request.to_function_approval_response(approved=True) + result = Content.from_function_result(call_id="call_resume", result="completed successfully") + + filtered = _filter_approval_control_messages([ + Message(role="assistant", contents=[function_call, request]), + Message(role="tool", contents=[result, response]), + ]) + + controls = [ + content + for message in filtered + for content in message.contents + if content.type in {"function_approval_request", "function_approval_response"} + ] + assert controls == [] + + +def test_filter_approval_controls_follow_up_does_not_resolve_approval_controls() -> None: + """Follow-up requests do NOT resolve approval controls in session history filtering. + + _approval_controls_to_keep resolves requests via matching function_approval_response + (not via follow-up requests). The request is removed because a response exists, + but the response itself is preserved since no terminal result has arrived. + Follow-up requests have no effect on this filtering logic. + """ + function_call = Content.from_function_call(call_id="call_followup", name="guarded", arguments="{}") + request = Content.from_function_approval_request(id="approval_followup", function_call=function_call) + response = request.to_function_approval_response(approved=True) + + follow_up = Content.from_text("Please provide more details") + follow_up.user_input_request = True + + filtered = _filter_approval_control_messages([ + Message(role="assistant", contents=[function_call, request]), + Message(role="user", contents=[response]), + Message(role="user", contents=[follow_up]), + ]) + + controls = [ + content + for message in filtered + for content in message.contents + if content.type in {"function_approval_request", "function_approval_response"} + ] + assert len(controls) == 1 + assert controls[0].type == "function_approval_response" + assert controls[0].id == "approval_followup" + + +def test_filter_approval_controls_preserves_unresolved_across_messages() -> None: + """Ensures unresolved approvals survive when no terminal result exists.""" + function_call = Content.from_function_call(call_id="call_pending", name="guarded", arguments="{}") + request = Content.from_function_approval_request(id="approval_pending", function_call=function_call) + response = request.to_function_approval_response(approved=True) + + filtered = _filter_approval_control_messages([ + Message(role="assistant", contents=[function_call, request]), + Message(role="user", contents=[response]), + ]) + + controls = [ + content + for message in filtered + for content in message.contents + if content.type in {"function_approval_request", "function_approval_response"} + ] + assert len(controls) == 1 + assert controls[0].type == "function_approval_response" + assert controls[0].id == "approval_pending" + + class TestHistoryProviderBase: def test_default_flags(self) -> None: provider = ConcreteHistoryProvider("mem") diff --git a/python/packages/core/tests/core/test_tools.py b/python/packages/core/tests/core/test_tools.py index 33fad82ddc..7ab9dd1e55 100644 --- a/python/packages/core/tests/core/test_tools.py +++ b/python/packages/core/tests/core/test_tools.py @@ -13,13 +13,17 @@ SKIP_PARSING, Content, FunctionTool, + Message, tool, ) from agent_framework._middleware import FunctionInvocationContext from agent_framework._tools import ( _auto_invoke_function, + _collect_approval_responses, + _collect_unanswered_approval_requests, _parse_annotation, _parse_inputs, + _remove_unanswered_approval_batches_from_model_input, normalize_function_invocation_configuration, ) from agent_framework.observability import OtelAttr @@ -1576,3 +1580,111 @@ def test_skip_parsing_is_singleton() -> None: # endregion + +# region Approval collection and filtering regression tests + + +def test_collect_approval_responses_order_independent_result_first() -> None: + """Result before response must still mark as resolved.""" + call = Content.from_function_call(call_id="c1", name="t", arguments="{}") + req = Content.from_function_approval_request(id="a1", function_call=call) + resp = req.to_function_approval_response(approved=True) + result = Content.from_function_result(call_id="c1", result="done") + + messages = [Message(role="tool", contents=[result, resp])] + collected = _collect_approval_responses(messages) + + assert collected == {} + + +def test_collect_approval_responses_follow_up_does_not_suppress_response() -> None: + """An unrelated user-input request without a call_id does not resolve approval responses. + + `_collect_approval_responses` only treats an approval response as resolved when a terminal + `function_result` or a follow-up `user_input_request` with the same `call_id` exists. + """ + call = Content.from_function_call(call_id="c2", name="t", arguments="{}") + req = Content.from_function_approval_request(id="a2", function_call=call) + resp = req.to_function_approval_response(approved=True) + follow_up = Content.from_text("more info needed") + follow_up.user_input_request = True + + messages = [ + Message(role="assistant", contents=[call, req]), + Message(role="user", contents=[resp]), + Message(role="user", contents=[follow_up]), + ] + collected = _collect_approval_responses(messages) + + assert "a2" in collected + assert collected["a2"].type == "function_approval_response" + + +def test_collect_unanswered_requests_respects_call_id_reuse() -> None: + """A result before a new request does not answer the new request (reused call_id).""" + call = Content.from_function_call(call_id="c3", name="t", arguments="{}") + req = Content.from_function_approval_request(id="a3", function_call=call) + result = Content.from_function_result(call_id="c3", result="done") + + messages = [ + Message(role="tool", contents=[result]), + Message(role="assistant", contents=[call, req]), + ] + unanswered = _collect_unanswered_approval_requests(messages) + + assert unanswered == [req] + + +def test_remove_unanswered_batches_strips_resolved_local_responses() -> None: + """Resolved local approval responses are stripped from model input. + + The corresponding answered request is preserved for model context. + Only the local response is removed to prevent MCP serialization leaks. + """ + call = Content.from_function_call(call_id="c4", name="local_tool", arguments="{}") + req = Content.from_function_approval_request(id="a4", function_call=call) + resp = req.to_function_approval_response(approved=True) + result = Content.from_function_result(call_id="c4", result="ok") + + messages = [ + Message(role="assistant", contents=[call, req]), + Message(role="user", contents=[resp]), + Message(role="tool", contents=[result]), + ] + + _remove_unanswered_approval_batches_from_model_input(messages) + + remaining_controls = [ + c for m in messages for c in m.contents if c.type in {"function_approval_request", "function_approval_response"} + ] + assert len(remaining_controls) == 1 + assert remaining_controls[0].type == "function_approval_request" + assert remaining_controls[0].id == "a4" + + +def test_remove_unanswered_batches_preserves_hosted_responses() -> None: + """Hosted (MCP) approval responses are never stripped by this function.""" + hosted_call = Content.from_function_call( + call_id="mcp1", + name="hosted", + arguments="{}", + additional_properties={"server_label": "srv"}, + ) + hosted_req = Content.from_function_approval_request(id="mcp_a1", function_call=hosted_call) + hosted_resp = hosted_req.to_function_approval_response(approved=True) + result = Content.from_function_result(call_id="mcp1", result="ok") + + messages = [ + Message(role="assistant", contents=[hosted_call, hosted_req]), + Message(role="user", contents=[hosted_resp]), + Message(role="tool", contents=[result]), + ] + + _remove_unanswered_approval_batches_from_model_input(messages) + + remaining = [c for m in messages for c in m.contents if c.type == "function_approval_response"] + assert len(remaining) == 1 + assert remaining[0].id == "mcp_a1" + + +# endregion diff --git a/python/packages/core/tests/workflow/test_agent_executor.py b/python/packages/core/tests/workflow/test_agent_executor.py index 2cc2ed2ce6..4b067ab7c1 100644 --- a/python/packages/core/tests/workflow/test_agent_executor.py +++ b/python/packages/core/tests/workflow/test_agent_executor.py @@ -1,7 +1,6 @@ # Copyright (c) Microsoft. All rights reserved. import pickle - from collections.abc import AsyncIterable, Awaitable from typing import Any, Literal, overload diff --git a/python/packages/openai/agent_framework_openai/_chat_client.py b/python/packages/openai/agent_framework_openai/_chat_client.py index a9120e72ca..f31613df86 100644 --- a/python/packages/openai/agent_framework_openai/_chat_client.py +++ b/python/packages/openai/agent_framework_openai/_chat_client.py @@ -2065,16 +2065,18 @@ def _prepare_content_for_openai( "output": output, } case "function_approval_request": + if not _is_hosted_tool_approval(content): + return {} return { "type": "mcp_approval_request", "id": content.id, "arguments": content.function_call.arguments, # type: ignore[union-attr] "name": content.function_call.name, # type: ignore[union-attr] - "server_label": content.function_call.additional_properties.get("server_label") # type: ignore[union-attr] - if content.function_call.additional_properties # type: ignore[union-attr] - else None, + "server_label": content.function_call.additional_properties.get("server_label"), # type: ignore[union-attr] } case "function_approval_response": + if not _is_hosted_tool_approval(content): + return {} return { "type": "mcp_approval_response", "approval_request_id": content.id, diff --git a/python/packages/openai/tests/openai/test_openai_chat_client.py b/python/packages/openai/tests/openai/test_openai_chat_client.py index 7258c50cae..491fa3f8c4 100644 --- a/python/packages/openai/tests/openai/test_openai_chat_client.py +++ b/python/packages/openai/tests/openai/test_openai_chat_client.py @@ -2854,9 +2854,7 @@ def test_prepare_content_for_opentool_approval_response() -> None: result = client._prepare_content_for_openai("assistant", approval_response) - assert result["type"] == "mcp_approval_response" - assert result["approval_request_id"] == "approval_001" - assert result["approve"] is True + assert result == {} def test_prepare_content_for_openai_error_content() -> None: @@ -3580,11 +3578,13 @@ def test_prepare_message_for_openai_with_function_approval_response() -> None: result = client._prepare_message_for_openai(message, request_uses_service_side_storage=False) # FunctionApprovalResponseContent is added directly, not nested in args with role - assert len(result) == 1 - prepared_message = result[0] - assert prepared_message["type"] == "mcp_approval_response" - assert prepared_message["approval_request_id"] == "approval_003" - assert prepared_message["approve"] is True + assert result == [ + { + "type": "mcp_approval_response", + "approval_request_id": "approval_003", + "approve": True, + } + ] def test_prepare_messages_for_openai_keeps_active_function_call_for_tool_loop() -> None: @@ -4136,10 +4136,9 @@ def test_function_approval_response_with_mcp_tool_call() -> None: """Test function approval response content with MCP server tool call content.""" client = OpenAIChatClient(model="test-model", api_key="test-key") - mcp_call = Content.from_mcp_server_tool_call( + mcp_call = Content.from_function_call( call_id="mcp_call_999", - tool_name="sensitive_action", - server_name="SecureServer", + name="sensitive_action", arguments={"action": "delete"}, additional_properties={"server_label": "SecureServer"}, ) @@ -8211,7 +8210,11 @@ def test_prepare_messages_keeps_function_call_without_storage() -> None: @pytest.mark.parametrize("approved", [True, False], ids=["approved", "rejected"]) def test_prepare_messages_strips_approval_request_but_keeps_response_under_storage(approved: bool) -> None: - """Stored requests are not replayed, but the new approval decision must reach the service.""" + """Under service-side storage, the hosted request is suppressed but its response is retained. + + When storage is off, hosted (MCP) approvals are serialized normally. + Local approvals are always dropped regardless of storage setting. + """ client = OpenAIChatClient(model="test-model", api_key="test-key") function_call = Content.from_function_call( @@ -8385,9 +8388,6 @@ def test_prepare_messages_strips_mcp_items_under_storage() -> None: # endregion -# endregion - - # region Prompt cache breakpoints and options @@ -8470,3 +8470,106 @@ async def test_prepare_options_prompt_cache_options_guarded_on_old_openai(monkey # endregion + +# region Approval serialization regression tests + + +def test_prepare_content_drops_local_approval_request() -> None: + """Local tool approval requests must not be serialized as mcp_approval_request.""" + client = RawOpenAIChatClient("gpt-4o-mini", api_key="sk-test") + local_call = Content.from_function_call(call_id="local_1", name="read_file", arguments="{}") + local_request = Content.from_function_approval_request(id="a1", function_call=local_call) + + result = client._prepare_content_for_openai("assistant", local_request) + + assert result == {} + + +def test_prepare_content_drops_local_approval_response() -> None: + """Local tool approval responses must not be serialized as mcp_approval_response. + + without the hosted-tool guard, local approvals + were emitted as mcp_approval_response with no matching request 400 from API. + """ + client = RawOpenAIChatClient("gpt-4o-mini", api_key="sk-test") + local_call = Content.from_function_call(call_id="local_2", name="read_file", arguments="{}") + local_response = Content.from_function_approval_response(approved=True, id="a2", function_call=local_call) + + result = client._prepare_content_for_openai("user", local_response) + + assert result == {} + + +def test_prepare_content_serializes_hosted_approval_request() -> None: + """Hosted (MCP) approval requests ARE serialized with server_label.""" + client = RawOpenAIChatClient("gpt-4o-mini", api_key="sk-test") + hosted_call = Content.from_function_call( + call_id="mcp_1", + name="hosted_tool", + arguments='{"x": 1}', + additional_properties={"server_label": "my_server"}, + ) + hosted_request = Content.from_function_approval_request(id="mcp_a1", function_call=hosted_call) + + result = client._prepare_content_for_openai("assistant", hosted_request) + + assert result["type"] == "mcp_approval_request" + assert result["id"] == "mcp_a1" + assert result["server_label"] == "my_server" + assert result["name"] == "hosted_tool" + + +def test_prepare_content_serializes_hosted_approval_response() -> None: + """Hosted (MCP) approval responses ARE serialized normally.""" + client = RawOpenAIChatClient("gpt-4o-mini", api_key="sk-test") + hosted_call = Content.from_function_call( + call_id="mcp_2", + name="hosted_tool", + arguments="{}", + additional_properties={"server_label": "my_server"}, + ) + hosted_response = Content.from_function_approval_response(approved=True, id="mcp_a2", function_call=hosted_call) + + result = client._prepare_content_for_openai("user", hosted_response) + + assert result["type"] == "mcp_approval_response" + assert result["approval_request_id"] == "mcp_a2" + assert result["approve"] is True + + +def test_prepare_messages_stores_suppresses_request_but_keeps_response() -> None: + """Under service-side storage, approval request is suppressed but response is kept. + + The response must reach the service so the decision is recorded. + mcp_approval_response 400 from API. + """ + client = RawOpenAIChatClient("gpt-4o-mini", api_key="sk-test") + hosted_call = Content.from_function_call( + call_id="mcp_3", + name="hosted_tool", + arguments="{}", + additional_properties={"server_label": "srv"}, + ) + hosted_request = Content.from_function_approval_request(id="mcp_a3", function_call=hosted_call) + hosted_response = Content.from_function_approval_response(approved=True, id="mcp_a3", function_call=hosted_call) + + messages = [ + Message(role="assistant", contents=[hosted_request]), + Message(role="user", contents=[hosted_response]), + ] + + prepared_messages: list[dict] = [] + for message in messages: + for content in message.contents: + if content.type == "function_approval_request": + continue + prepared = client._prepare_content_for_openai(message.role, content) + if prepared: + prepared_messages.append(prepared) + + approval_types = {m.get("type") for m in prepared_messages} + assert "mcp_approval_request" not in approval_types + assert "mcp_approval_response" in approval_types + + +# endregion