From 91c4978b8006878e785113665f1de33027f3accf Mon Sep 17 00:00:00 2001 From: CorgiBoyG <111257566+CorgiBoyG@users.noreply.github.com> Date: Wed, 30 Sep 2026 09:44:12 +0900 Subject: [PATCH] Python: Fix active mixed-pause Host response correlation --- .../packages/core/agent_framework/_tools.py | 99 +++- .../core/test_function_invocation_logic.py | 422 +++++++++++++++++- 2 files changed, 502 insertions(+), 19 deletions(-) diff --git a/python/packages/core/agent_framework/_tools.py b/python/packages/core/agent_framework/_tools.py index 4ab15a52f79..8d37892e27d 100644 --- a/python/packages/core/agent_framework/_tools.py +++ b/python/packages/core/agent_framework/_tools.py @@ -3270,6 +3270,8 @@ def _match_mixed_pause_responses( host_items_by_occurrence[request.id] = index matched_content_ids: set[int] = set() + preexisting_response_indexes = {index for index, item in enumerate(items) if item.get("response") is not None} + current_idless_response_indexes: set[int] = set() for match_idless_host_results in (False, True): for response in responses: is_idless_host_result = response.type == "function_result" and response.id is None @@ -3310,6 +3312,18 @@ def _match_mixed_pause_responses( if items[pending_index].get("response") is None ] if len(unanswered_indexes) == 1: + if any( + pending_index in preexisting_response_indexes + and isinstance(items[pending_index].get("response"), Mapping) + and _same_mixed_pause_response( + cast(Mapping[str, Any], items[pending_index]["response"]), + response.to_dict(), + ) + for pending_index in host_items_by_call[response.call_id] + ): + raise RuntimeError( + f"Ambiguous id-less Host response for mixed pause call_id {response.call_id!r}." + ) item_index = unanswered_indexes[0] elif not unanswered_indexes and allow_idless_host_duplicates: duplicate_indexes = [ @@ -3323,6 +3337,13 @@ def _match_mixed_pause_responses( ] if len(duplicate_indexes) == 1: item_index = duplicate_indexes[0] + elif any( + pending_index in current_idless_response_indexes + for pending_index in host_items_by_call[response.call_id] + ): + raise RuntimeError( + f"Conflicting id-less Host responses for mixed pause call_id {response.call_id!r}." + ) else: continue @@ -3337,6 +3358,8 @@ def _match_mixed_pause_responses( raise RuntimeError(f"Conflicting response for mixed pause occurrence {candidate.id!r}.") items[item_index]["response"] = candidate_state matched_content_ids.add(id(response)) + if is_idless_host_result: + current_idless_response_indexes.add(item_index) if any(item.get("response") is None for item in items): return matched_content_ids, True, [], set() @@ -3383,11 +3406,81 @@ def bind_approval_response(response: Content) -> Content | None: consume=False, ) + host_requests = [ + request + for item in items + if item.get("kind") == "host" + and item.get("response") is None + and (request := _content_from_state(item.get("request"))) is not None + and request.call_id is not None + ] + approval_request_ids = { + str(identity) + for item in items + if item.get("kind") == "approval" + and item.get("response") is None + and (request := _content_from_state(item.get("request"))) is not None + for identity in ( + request.id, + request.function_call.id if request.function_call is not None else None, + ) + if identity is not None + } + host_call_ids = {request.call_id for request in host_requests if request.call_id is not None} + host_occurrences = { + (request.call_id, request.id) + for request in host_requests + if request.call_id is not None and request.id is not None + } + approval_anchor: int | None = None + host_anchor: int | None = None + indexed_contents: list[tuple[int, Content]] = [] + content_index = -1 + for message in messages: + message_start = content_index + 1 + for content in message.contents: + content_index += 1 + indexed_contents.append((content_index, content)) + if content.type == "function_approval_response" and approval_anchor is None: + bound_approval = bind_approval_response(content) + if bound_approval is not None and approval_request_ids.intersection( + str(identity) + for identity in ( + bound_approval.additional_properties.get(_APPROVAL_REQUEST_ID_KEY), + bound_approval.id, + ) + if identity is not None + ): + approval_anchor = message_start + elif ( + content.type == "function_result" + and content.call_id is not None + and content.id is not None + and (content.call_id, content.id) in host_occurrences + and host_anchor is None + ): + host_anchor = message_start + response_start = min( + (anchor for anchor in (approval_anchor, host_anchor) if anchor is not None), + default=None, + ) + approval_is_staged = any(item.get("kind") == "approval" and item.get("response") is not None for item in items) + if response_start is None and approval_is_staged and messages: + last_message_start = len(indexed_contents) - len(messages[-1].contents) + response_start = next( + ( + index + for index, content in indexed_contents[last_message_start:] + if content.type == "function_result" and content.id is None and content.call_id in host_call_ids + ), + None, + ) responses = [ content - for message in messages - for content in message.contents - if content.type in {"function_approval_response", "function_result"} + for index, content in indexed_contents + if response_start is not None + and index >= response_start + and content.type in {"function_approval_response", "function_result"} ] matched_content_ids, incomplete, ordered_responses, host_result_ids = _match_mixed_pause_responses( items, 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..3356186d79b 100644 --- a/python/packages/core/tests/core/test_function_invocation_logic.py +++ b/python/packages/core/tests/core/test_function_invocation_logic.py @@ -4999,6 +4999,143 @@ def test_stateful_mixed_batch_assigns_idless_equal_result_to_unanswered_occurren assert host_result_ids == {id(messages[-1].contents[1]), id(messages[-1].contents[2])} +def test_stateful_mixed_batch_rejects_idless_replay_of_previously_staged_result() -> None: + """An id-less replay cannot fill a sibling occurrence after partial recovery.""" + from agent_framework._tools import ( + _stage_pending_mixed_pause_responses, + _store_pending_approval_requests, + _store_pending_mixed_pause_batch, + ) + + session = AgentSession() + approval_call = Content.from_function_call( + call_id="approval", + name="approval_func", + arguments={}, + id="approval-occurrence", + ) + approval_request = Content.from_function_approval_request( + id="approval-occurrence", + function_call=approval_call, + ) + host_requests = [ + Content.from_function_call( + call_id="shared", + name="host_func", + arguments={}, + id=f"host-occurrence-{index}", + ) + for index in range(2) + ] + for request in host_requests: + request.user_input_request = True + identified_result = Content.from_function_result(call_id="shared", result="same result") + identified_result.id = "host-occurrence-0" + approval_response = approval_request.to_function_approval_response(approved=True) + _store_pending_approval_requests(session, [approval_request]) + _store_pending_mixed_pause_batch(session, [[approval_request], *[[request] for request in host_requests]]) + + first_messages = [Message(role="user", contents=[approval_response, identified_result])] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(first_messages, session) + assert incomplete is True + assert completed is False + + replay = Content.from_function_result(call_id="shared", result="same result") + full_transcript = [Message(role="user", contents=[approval_response, identified_result, replay])] + with pytest.raises(RuntimeError, match="Ambiguous id-less Host response"): + _stage_pending_mixed_pause_responses(full_transcript, session) + + distinct_result = Content.from_function_result(call_id="shared", result="second result") + full_transcript = [Message(role="user", contents=[approval_response, identified_result, distinct_result])] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(full_transcript, session) + assert incomplete is False + assert completed is True + assert [content.result for content in full_transcript[-1].contents if content.type == "function_result"] == [ + "same result", + "second result", + ] + + +def test_stateful_mixed_batch_accepts_exact_host_results_across_messages() -> None: + """The first exact Host response anchors a split response window.""" + from agent_framework._tools import ( + _stage_pending_mixed_pause_responses, + _store_pending_approval_requests, + _store_pending_mixed_pause_batch, + ) + + session = AgentSession() + approval_call = Content.from_function_call(call_id="approval", name="approval_func", arguments={}, id="approval") + approval_request = Content.from_function_approval_request(id="approval", function_call=approval_call) + host_requests = [ + Content.from_function_call(call_id=f"host-{index}", name="host_func", arguments={}, id=f"host-{index}") + for index in range(2) + ] + for request in host_requests: + request.user_input_request = True + _store_pending_approval_requests(session, [approval_request]) + _store_pending_mixed_pause_batch(session, [[approval_request], *[[request] for request in host_requests]]) + + approval_messages = [Message(role="user", contents=[approval_request.to_function_approval_response(approved=True)])] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(approval_messages, session) + assert incomplete is True + assert completed is False + + host_messages: list[Message] = [] + for request in host_requests: + assert request.call_id is not None + assert request.id is not None + result = Content.from_function_result(call_id=request.call_id, result=request.id) + result.id = request.id + host_messages.append(Message(role="tool", contents=[result])) + incomplete, completed, _ = _stage_pending_mixed_pause_responses(host_messages, session) + assert incomplete is False + assert completed is True + assert [content.result for content in host_messages[-1].contents if content.type == "function_result"] == [ + "host-0", + "host-1", + ] + + +def test_stateful_mixed_batch_accepts_idless_host_result_after_staged_approval() -> None: + """A latest-message id-less Host result can continue a partially staged batch.""" + from agent_framework._tools import ( + _stage_pending_mixed_pause_responses, + _store_pending_approval_requests, + _store_pending_mixed_pause_batch, + ) + + session = AgentSession() + approval_call = Content.from_function_call(call_id="approval", name="approval_func", arguments={}, id="approval") + approval_request = Content.from_function_approval_request(id="approval", function_call=approval_call) + host_request = Content.from_function_call(call_id="host", name="host_func", arguments={}, id="host") + host_request.user_input_request = True + _store_pending_approval_requests(session, [approval_request]) + _store_pending_mixed_pause_batch(session, [[approval_request], [host_request]]) + + approval_messages = [Message(role="user", contents=[approval_request.to_function_approval_response(approved=True)])] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(approval_messages, session) + assert incomplete is True + assert completed is False + + historical_result = Content.from_function_result(call_id="host", result="historical") + historical_messages = [ + Message(role="tool", contents=[historical_result]), + Message(role="user", contents=["later"]), + ] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(historical_messages, session) + assert incomplete is True + assert completed is False + assert historical_messages[0].contents == [historical_result] + + current_result = Content.from_function_result(call_id="host", result="current") + current_messages = [Message(role="tool", contents=[current_result])] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(current_messages, session) + assert incomplete is False + assert completed is True + assert current_messages[-1].contents[-1].result == "current" + + def test_stateless_mixed_batch_rejects_conflicting_identified_host_results() -> None: """Conflicting results for one identified Host occurrence fail closed.""" from agent_framework._tools import _stateless_mixed_pause_batch_status @@ -5040,8 +5177,8 @@ def test_stateless_mixed_batch_rejects_conflicting_identified_host_results() -> _stateless_mixed_pause_batch_status(messages) -def test_active_mixed_pause_ignores_historical_host_requests() -> None: - """Only the session-recorded mixed batch participates in response correlation.""" +def test_active_mixed_pause_ignores_historical_idless_result_with_reused_call_id() -> None: + """A historical id-less result cannot be consumed as an active duplicate.""" from agent_framework._tools import ( _stage_pending_mixed_pause_responses, _store_pending_approval_requests, @@ -5070,14 +5207,13 @@ def test_active_mixed_pause_ignores_historical_host_requests() -> None: _store_pending_mixed_pause_batch(session, [[approval_request], [host_request]]) completed_old_host = Content.from_function_call( - call_id="old-completed", + call_id="current-host", name="old_host", arguments={}, id="old-completed-occurrence", ) completed_old_host.user_input_request = True - completed_old_result = Content.from_function_result(call_id="old-completed", result="old result") - completed_old_result.id = "old-completed-occurrence" + completed_old_result = Content.from_function_result(call_id="current-host", result="current result") abandoned_old_host = Content.from_function_call( call_id="old-abandoned", name="old_host", @@ -5087,24 +5223,20 @@ def test_active_mixed_pause_ignores_historical_host_requests() -> None: abandoned_old_host.user_input_request = True host_result = Content.from_function_result(call_id="current-host", result="current result") host_result.id = "current-host-occurrence" - messages = [ - Message(role="assistant", contents=[completed_old_host, abandoned_old_host]), - Message(role="tool", contents=[completed_old_result]), - Message( - role="user", - contents=[ - approval_request.to_function_approval_response(approved=True), - host_result, - ], - ), + current_contents = [ + approval_request.to_function_approval_response(approved=True), + host_result, ] + messages = [Message(role="assistant", contents=[completed_old_host, abandoned_old_host])] + messages.append(Message(role="tool", contents=[completed_old_result])) + messages.append(Message(role="user", contents=current_contents)) incomplete, completed, host_result_ids = _stage_pending_mixed_pause_responses(messages, session) assert incomplete is False assert completed is True assert messages[0].contents == [completed_old_host, abandoned_old_host] - assert messages[1].contents == [completed_old_result] + assert any(content is completed_old_result for message in messages[:-1] for content in message.contents) assert [(content.type, content.id) for content in messages[-1].contents] == [ ("function_approval_response", "current-approval-occurrence"), ("function_result", "current-host-occurrence"), @@ -5112,6 +5244,264 @@ def test_active_mixed_pause_ignores_historical_host_requests() -> None: assert host_result_ids == {id(messages[-1].contents[-1])} +def test_active_mixed_pause_accepts_idless_host_result_before_approval_in_same_message() -> None: + """A Host-first current reply retains its id-less result before the approval anchor.""" + from agent_framework._tools import ( + _stage_pending_mixed_pause_responses, + _store_pending_approval_requests, + _store_pending_mixed_pause_batch, + ) + + session = AgentSession() + approval_call = Content.from_function_call( + call_id="approval", + name="guarded", + arguments={}, + id="approval-occurrence", + ) + approval_request = Content.from_function_approval_request( + id="approval-occurrence", + function_call=approval_call, + ) + host_request = Content.from_function_call( + call_id="host", + name="host", + arguments={}, + id="host-occurrence", + ) + host_request.user_input_request = True + _store_pending_approval_requests(session, [approval_request]) + _store_pending_mixed_pause_batch(session, [[approval_request], [host_request]]) + + historical_result = Content.from_function_result(call_id="host", result="historical result") + host_result = Content.from_function_result(call_id="host", result="current result") + approval_response = approval_request.to_function_approval_response(approved=True) + messages = [ + Message(role="tool", contents=[historical_result]), + Message(role="user", contents=[host_result, approval_response]), + ] + + incomplete, completed, host_result_ids = _stage_pending_mixed_pause_responses(messages, session) + + assert incomplete is False + assert completed is True + assert messages[0].contents == [historical_result] + assert [(content.type, content.id) for content in messages[-1].contents] == [ + ("function_approval_response", "approval-occurrence"), + ("function_result", None), + ] + assert messages[-1].contents[-1].result == "current result" + assert host_result_ids == {id(messages[-1].contents[-1])} + + +@pytest.mark.parametrize("approval_staged", [False, True], ids=["same-resume", "partial-resume"]) +def test_active_mixed_pause_rejects_conflicting_host_first_idless_results(approval_staged: bool) -> None: + """Conflicting id-less results in the anchored message cannot be ordered by guesswork.""" + from agent_framework._tools import ( + _stage_pending_mixed_pause_responses, + _store_pending_approval_requests, + _store_pending_mixed_pause_batch, + ) + + session = AgentSession() + approval_call = Content.from_function_call( + call_id="approval", + name="guarded", + arguments={}, + id="approval-occurrence", + ) + approval_request = Content.from_function_approval_request( + id="approval-occurrence", + function_call=approval_call, + ) + host_request = Content.from_function_call( + call_id="host", + name="host", + arguments={}, + id="host-occurrence", + ) + host_request.user_input_request = True + _store_pending_approval_requests(session, [approval_request]) + _store_pending_mixed_pause_batch(session, [[approval_request], [host_request]]) + + first_result = Content.from_function_result(call_id="host", result="first") + second_result = Content.from_function_result(call_id="host", result="second") + approval_response = approval_request.to_function_approval_response(approved=True) + if approval_staged: + approval_messages = [Message(role="user", contents=[approval_response])] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(approval_messages, session) + assert incomplete is True + assert completed is False + messages = [Message(role="user", contents=[first_result, second_result])] + else: + messages = [Message(role="user", contents=[first_result, second_result, approval_response])] + + with pytest.raises(RuntimeError, match="Conflicting id-less Host responses"): + _stage_pending_mixed_pause_responses(messages, session) + + +def test_active_mixed_pause_exact_host_result_precedes_weak_idless_candidate() -> None: + """An exact occurrence result remains authoritative over an earlier id-less candidate.""" + from agent_framework._tools import ( + _stage_pending_mixed_pause_responses, + _store_pending_approval_requests, + _store_pending_mixed_pause_batch, + ) + + session = AgentSession() + approval_call = Content.from_function_call( + call_id="approval", + name="guarded", + arguments={}, + id="approval-occurrence", + ) + approval_request = Content.from_function_approval_request( + id="approval-occurrence", + function_call=approval_call, + ) + host_request = Content.from_function_call( + call_id="host", + name="host", + arguments={}, + id="host-occurrence", + ) + host_request.user_input_request = True + _store_pending_approval_requests(session, [approval_request]) + _store_pending_mixed_pause_batch(session, [[approval_request], [host_request]]) + + weak_result = Content.from_function_result(call_id="host", result="weak") + exact_result = Content.from_function_result(call_id="host", result="exact") + exact_result.id = "host-occurrence" + approval_response = approval_request.to_function_approval_response(approved=True) + messages = [Message(role="user", contents=[weak_result, exact_result, approval_response])] + + incomplete, completed, _ = _stage_pending_mixed_pause_responses(messages, session) + + assert incomplete is False + assert completed is True + assert messages[0].contents == [weak_result] + assert [content.result for content in messages[-1].contents if content.type == "function_result"] == ["exact"] + + +def test_active_mixed_pause_staged_approval_replay_does_not_anchor_full_history() -> None: + """A replayed staged approval cannot admit an intervening stale Host result.""" + from agent_framework._tools import ( + _stage_pending_mixed_pause_responses, + _store_pending_approval_requests, + _store_pending_mixed_pause_batch, + ) + + session = AgentSession() + approval_call = Content.from_function_call( + call_id="approval", + name="guarded", + arguments={}, + id="approval-occurrence", + ) + approval_request = Content.from_function_approval_request( + id="approval-occurrence", + function_call=approval_call, + ) + host_request = Content.from_function_call( + call_id="host", + name="host", + arguments={}, + id="host-occurrence", + ) + host_request.user_input_request = True + _store_pending_approval_requests(session, [approval_request]) + _store_pending_mixed_pause_batch(session, [[approval_request], [host_request]]) + + approval_response = approval_request.to_function_approval_response(approved=True) + partial_messages = [Message(role="user", contents=[approval_response])] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(partial_messages, session) + assert incomplete is True + assert completed is False + + replayed_approval = approval_request.to_function_approval_response(approved=True) + stale_result = Content.from_function_result(call_id="host", result="stale") + current_result = Content.from_function_result(call_id="host", result="current") + full_history = [ + Message(role="user", contents=[replayed_approval]), + Message(role="tool", contents=[stale_result]), + Message(role="user", contents=[current_result]), + ] + + incomplete, completed, _ = _stage_pending_mixed_pause_responses(full_history, session) + + assert incomplete is False + assert completed is True + assert full_history[0].contents == [replayed_approval] + assert full_history[1].contents == [stale_result] + assert [content.result for content in full_history[-1].contents if content.type == "function_result"] == ["current"] + + +def test_active_mixed_pause_staged_exact_host_replay_does_not_anchor_full_history() -> None: + """A replayed staged exact Host result cannot admit stale input for a sibling occurrence.""" + from agent_framework._tools import ( + _stage_pending_mixed_pause_responses, + _store_pending_approval_requests, + _store_pending_mixed_pause_batch, + ) + + session = AgentSession() + approval_call = Content.from_function_call( + call_id="approval", + name="guarded", + arguments={}, + id="approval-occurrence", + ) + approval_request = Content.from_function_approval_request( + id="approval-occurrence", + function_call=approval_call, + ) + first_host = Content.from_function_call( + call_id="shared-host", + name="host", + arguments={}, + id="host-occurrence-1", + ) + first_host.user_input_request = True + second_host = Content.from_function_call( + call_id="shared-host", + name="host", + arguments={}, + id="host-occurrence-2", + ) + second_host.user_input_request = True + _store_pending_approval_requests(session, [approval_request]) + _store_pending_mixed_pause_batch(session, [[approval_request], [first_host, second_host]]) + + approval_response = approval_request.to_function_approval_response(approved=True) + first_result = Content.from_function_result(call_id="shared-host", result="first") + first_result.id = "host-occurrence-1" + partial_messages = [Message(role="user", contents=[approval_response, first_result])] + incomplete, completed, _ = _stage_pending_mixed_pause_responses(partial_messages, session) + assert incomplete is True + assert completed is False + + replayed_first_result = Content.from_function_result(call_id="shared-host", result="first") + replayed_first_result.id = "host-occurrence-1" + stale_result = Content.from_function_result(call_id="shared-host", result="stale") + current_result = Content.from_function_result(call_id="shared-host", result="current") + full_history = [ + Message(role="tool", contents=[replayed_first_result]), + Message(role="tool", contents=[stale_result]), + Message(role="user", contents=[current_result]), + ] + + incomplete, completed, _ = _stage_pending_mixed_pause_responses(full_history, session) + + assert incomplete is False + assert completed is True + assert full_history[0].contents == [replayed_first_result] + assert full_history[1].contents == [stale_result] + assert [content.result for content in full_history[-1].contents if content.type == "function_result"] == [ + "first", + "current", + ] + + async def test_function_invocation_config_additional_tools(chat_client_base: SupportsChatGetResponse): """Test that additional_tools are available but treated as declaration_only.""" exec_counter_visible = 0