From 30552059cc3d7684271017794d0ef5a6ade2ff8e Mon Sep 17 00:00:00 2001 From: rzx Date: Thu, 17 Sep 2026 15:02:50 +0800 Subject: [PATCH] fix(a2a): preserve execution lifecycle across reconnect and permission resume Keep request-scoped event delivery and natural-completion release synchronized with accepted requests, durable backup, permission handoff, and recovery admission. Fence late task projections and stale completion generations so a completed turn cannot overwrite or release a newer execution. Preserve terminal subscription shutdown from main, pin the SDK lifecycle contract, and make async and process tests synchronize on observable state across Python versions and Windows, including direct interpreter PIDs for the manager idle test. --- pyproject.toml | 10 +- .../run_execution_control_scenarios.py | 66 +- src/iac_code/a2a/app.py | 19 +- src/iac_code/a2a/execution_control.py | 1107 +++++- src/iac_code/a2a/executor.py | 223 +- src/iac_code/a2a/input_required.py | 32 +- src/iac_code/a2a/pipeline_executor.py | 160 +- .../a2a/request_scoped_active_task.py | 572 +++ src/iac_code/a2a/task_store.py | 486 ++- src/iac_code/a2a/transports/dispatcher.py | 665 +++- src/iac_code/services/providers/aliyun.py | 12 + tests/a2a/test_app.py | 383 +- tests/a2a/test_execution_control.py | 1189 ++++++ .../a2a/test_execution_control_regressions.py | 535 ++- tests/a2a/test_executor.py | 229 +- tests/a2a/test_input_required.py | 604 ++- .../a2a/test_permission_execution_control.py | 31 +- tests/a2a/test_pipeline_executor.py | 183 +- tests/a2a/test_resource_selector.py | 40 +- tests/a2a/test_task_store.py | 88 +- tests/a2a/test_transport_dispatcher.py | 3223 +++++++++++++++-- .../test_execution_control_scenarios.py | 20 +- tests/mcp/test_client.py | 4 +- .../test_alicloud_ros_agent_bridge.py | 113 +- tests/web/test_dynamic_mcp_commands.py | 167 +- uv.lock | 10 +- 26 files changed, 9350 insertions(+), 821 deletions(-) create mode 100644 src/iac_code/a2a/request_scoped_active_task.py diff --git a/pyproject.toml b/pyproject.toml index ad5c185a..46c604c1 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -54,18 +54,18 @@ dependencies = [ [project.optional-dependencies] http = [ - "a2a-sdk[http-server,signing]>=1.0.2,<2", + "a2a-sdk[http-server,signing]==1.1.0", "starlette>=0.39.0", "uvicorn[standard]>=0.30.0", ] a2a = [ - "a2a-sdk[http-server,signing]>=1.0.2,<2", + "a2a-sdk[http-server,signing]==1.1.0", "cryptography>=42.0", "starlette>=0.39.0", "uvicorn[standard]>=0.30.0", ] a2a-signing = [ - "a2a-sdk[signing]>=1.0.2,<2", + "a2a-sdk[signing]==1.1.0", ] a2a-grpc = [ "grpcio>=1.60.0", @@ -76,7 +76,7 @@ a2a-redis = [ ] agui = [ "ag-ui-protocol==0.1.20", - "a2a-sdk[http-server,signing]>=1.0.2,<2", + "a2a-sdk[http-server,signing]==1.1.0", "starlette>=0.39.0", "uvicorn[standard]>=0.30.0", ] @@ -99,7 +99,7 @@ dev = [ "pexpect>=4.9.0", ] desktop = [ - "a2a-sdk[http-server,signing]>=1.0.2,<2", + "a2a-sdk[http-server,signing]==1.1.0", "pyinstaller==6.21.0", "termaid>=0.1; python_version >= '3.11'", "uvicorn[standard]>=0.30.0", diff --git a/scripts/a2a/e2e/execution_control/run_execution_control_scenarios.py b/scripts/a2a/e2e/execution_control/run_execution_control_scenarios.py index 64b4d03b..bec073ff 100644 --- a/scripts/a2a/e2e/execution_control/run_execution_control_scenarios.py +++ b/scripts/a2a/e2e/execution_control/run_execution_control_scenarios.py @@ -201,6 +201,16 @@ def snapshot(self) -> list[dict[str, Any]]: with self._lock: return list(self.events) + def wait_for_text(self, text: str, *, timeout: float) -> None: + _wait_until( + lambda: self.error or self.done.is_set() or text in json.dumps(self.snapshot()), + timeout=timeout, + description="{} stream output {}".format(self.name, text), + ) + if self.error is not None: + raise RuntimeError("A2A stream failed: {}".format(self.error)) from self.error + assert text in json.dumps(self.snapshot()), "A2A stream ended before expected output" + def close(self) -> None: self._closed = True response = self._response @@ -543,12 +553,25 @@ def _assert_other_context_progress(self) -> None: other_context = _first(other.snapshot(), "contextId") assert other_context and other_context != self.context_id assert "ISOLATION_FIXTURE_FINAL" in json.dumps(other.snapshot()) - status, other_state = _http_json( - "GET", - self.server.url + "/iac-code/execution/state?" + urlencode({"contextId": other_context}), - timeout=1.0, + + def other_context_released() -> dict[str, Any] | None: + status, state = _http_json( + "GET", + self.server.url + "/iac-code/execution/state?" + urlencode({"contextId": other_context}), + timeout=1.0, + ) + assert status == 200 + if state["phase"] == "running": + return state + if state.get("terminationReason") == "natural_completion" and state.get("releaseReady") is True: + return state + return None + + _wait_until( + other_context_released, + timeout=self.timeout, + description="other context natural release", ) - assert status == 200 and other_state["phase"] == "running" assert not (self.control_dir / "release-storage").exists() lifecycle = _read_jsonl(self.run_dir / "fixture-lifecycle.jsonl") assert not any(value["event"].endswith("commit_finished") for value in lifecycle) @@ -562,6 +585,9 @@ def _assert_other_context_progress(self) -> None: "otherElapsedSeconds": time.monotonic() - started, }, ) + # Business progress above proves isolation while the primary write is held. + # Allow the normal scenario budget for backup and stream shutdown as well. + other.join(self.timeout) def _slow_termination_storage(self) -> None: _atomic_json(self.control_dir / "arm-termination-storage", {"contextId": self.context_id}) @@ -751,7 +777,11 @@ def _verify_agent_card(self) -> None: def _state(self) -> dict[str, Any]: self.server.assert_running() query = urlencode({"contextId": self.context_id}) - status, state = _http_json("GET", self.server.url + "/iac-code/execution/state?" + query) + status, state = _http_json( + "GET", + self.server.url + "/iac-code/execution/state?" + query, + timeout=self.timeout, + ) assert status == 200, state self.timeline.append({"observedAt": time.time(), **state}) _atomic_json(self.run_dir / "execution-state-timeline.json", self.timeline) @@ -939,7 +969,6 @@ def _warm_resume_paused(self) -> None: assert "DURABLE_FIXTURE_RESULT" in json.dumps(paused_recovery, ensure_ascii=False) subscription = self._subscribe() self._resume(epoch=2, request_id="resume-paused") - self._wait_state(lambda value: value["phase"] == "running", "running after paused resume") subscription.join(self.timeout) self._assert_normal_completion(subscription.snapshot()) @@ -1088,9 +1117,17 @@ def _natural_completion(self) -> None: assert transcript.count("NATURAL_FIXTURE_FINAL") == 1 assert recovery["outputText"] == ["NATURAL_FIXTURE_FINAL"] self._resume(epoch=2, request_id="resume-natural") - final = self._wait_state(lambda value: value["phase"] == "running", "natural completion hold release") + final = self._wait_state( + lambda value: ( + value["phase"] == "terminated" + and value.get("terminationReason") == "natural_completion" + and value.get("releaseReady") is True + ), + "natural completion hold release", + ) assert final.get("executionStatus") == completed.get("executionStatus") assert final.get("streamAvailable") is False + self._assert_shared_backup(final, None, require_reason=False) assert len(self._provider_calls()) == 1 def _capture_recovery(self, suffix: str) -> dict[str, Any]: @@ -1117,10 +1154,15 @@ def _assert_normal_completion(self, events: list[dict[str, Any]]) -> None: "completed normal turn", ) task = self._task() - state = self._state() - assert state["phase"] == "running" - assert state.get("terminationReason") is None - assert state.get("backup", {}).get("status") == "not_requested" + state = self._wait_state( + lambda value: ( + value["phase"] == "terminated" + and value.get("terminationReason") == "natural_completion" + and value.get("releaseReady") is True + ), + "natural completion release", + ) + self._assert_shared_backup(state, None, require_reason=False) assert state["executionId"] == self.execution_id assert state["serverInstanceId"] == self.server_instance_id assert len([value for value in self._tool_events() if value.get("event") == "tool.started"]) == 1 diff --git a/src/iac_code/a2a/app.py b/src/iac_code/a2a/app.py index 76785831..cbf24c21 100644 --- a/src/iac_code/a2a/app.py +++ b/src/iac_code/a2a/app.py @@ -653,6 +653,7 @@ async def pause_execution(request: Request) -> JSONResponse: reason=required_string(payload, "reason"), reconnect_timeout_seconds=float(timeout), ) + state = control.protocol_snapshot(state) return JSONResponse(state, status_code=200 if state["phase"] == "paused" else 202) except Exception as exc: return await execution_error_response(exc) @@ -662,14 +663,20 @@ async def get_execution_state(request: Request) -> JSONResponse: context_id = request.query_params.get("contextId") if not context_id: raise ValueError("contextId is required") - control = await execution_control_from_request(request, context_id) + service = components.execution_control_service + if service is None: + raise ExecutionControlNotFoundError("Execution control is unavailable") + state = await service.observe( + context_id=validate_protocol_id(context_id), + owner=execution_owner(request), + ) expected_execution_id = request.query_params.get("executionId") pause_id = request.query_params.get("pauseId") - if expected_execution_id is not None and expected_execution_id != control.execution_id: + if expected_execution_id is not None and expected_execution_id != state.get("executionId"): raise ExecutionControlConflictError("executionId does not identify the current execution") - if pause_id is not None and pause_id != control.pause_id: + if pause_id is not None and pause_id != state.get("pauseId"): raise ExecutionControlConflictError("pauseId does not identify the current pause") - return JSONResponse(control.snapshot()) + return JSONResponse(state) except Exception as exc: return await execution_error_response(exc) @@ -684,6 +691,7 @@ async def resume_execution(request: Request) -> JSONResponse: request_id=validate_protocol_id(required_string(payload, "requestId")), connection_epoch=required_epoch(payload), ) + state = control.protocol_snapshot(state) return JSONResponse(state, status_code=200 if state["phase"] == "running" else 202) except Exception as exc: return await execution_error_response(exc) @@ -708,6 +716,7 @@ async def terminate_execution(request: Request) -> JSONResponse: reason=reason, pause_id=pause_id, ) + state = control.protocol_snapshot(state) return JSONResponse(state, status_code=200 if state["phase"] == "terminated" else 202) except Exception as exc: return await execution_error_response(exc) @@ -752,7 +761,7 @@ async def get_session_recovery(request: Request) -> JSONResponse: "outputText": list(task_record.output_text), "task": MessageToDict(task, preserving_proto_field_name=False), "messages": [message.to_dict() for message in messages], - "executionControl": control.snapshot(), + "executionControl": control.protocol_snapshot(control.snapshot()), } return JSONResponse(project_a2a_data(recovery, public_path_roots=roots)) except Exception as exc: diff --git a/src/iac_code/a2a/execution_control.py b/src/iac_code/a2a/execution_control.py index f52e7ec8..c6e4ed03 100644 --- a/src/iac_code/a2a/execution_control.py +++ b/src/iac_code/a2a/execution_control.py @@ -9,6 +9,7 @@ import asyncio import hashlib import json +import logging import math import time import uuid @@ -20,10 +21,13 @@ from pathlib import Path from typing import Any, AsyncIterator, Literal, TypeVar, cast -from iac_code.a2a.backup import await_fenced, run_sync_fenced +from iac_code.a2a.backup import await_fenced, run_sync_fenced, run_sync_fenced_with_cancel_completion from iac_code.services.session_backup import BackupReason, BackupResult from iac_code.services.session_storage import SessionStorage -from iac_code.utils.state_io import atomic_write_json +from iac_code.utils.public_errors import sanitize_strict_text +from iac_code.utils.state_io import atomic_write_json, cross_process_file_lock + +logger = logging.getLogger(__name__) ExecutionPhase = Literal[ "running", @@ -37,6 +41,8 @@ _PAUSE_PHASES = frozenset({"pausing", "pause_committing", "paused"}) _TERMINAL_TASK_STATES = frozenset({"completed", "failed", "canceled", "input-required", "normal-turn-ended"}) +_RECOVERABLE_INPUT_ADMISSION_TTL_SECONDS = 60.0 +NATURAL_COMPLETION_FINALIZED_GENERATION = "_naturalCompletionFinalizedGeneration" _CURRENT_CONTROL: ContextVar[Any] = ContextVar("a2a_execution_control", default=None) _CURRENT_ACTIVITY_IDS: ContextVar[tuple[str, ...]] = ContextVar("a2a_execution_activity_ids", default=()) _CURRENT_PARTICIPANT_IDS: ContextVar[tuple[str, ...]] = ContextVar("a2a_execution_participant_ids", default=()) @@ -81,6 +87,369 @@ class _Activity: budget_changed: asyncio.Event = field(default_factory=asyncio.Event) +@dataclass(frozen=True) +class _RecoverableInputAdmission: + token: str + context_id: str + task_id: str + owner: str + expires_at: float + + def as_document(self) -> dict[str, Any]: + return { + "token": self.token, + "contextId": self.context_id, + "taskId": self.task_id, + "owner": self.owner, + "expiresAt": self.expires_at, + } + + +@dataclass(frozen=True) +class _RecoverableInputActivation: + previous_control: dict[str, Any] | None + execution_id: str + + +class _RecoverableInputAdmissionStore: + """One-request reservation store with a cross-process file fence when persistence is enabled.""" + + def __init__(self, persistence_root: Path | None) -> None: + self._root = persistence_root / "execution-control" if persistence_root is not None else None + self._by_token: dict[str, _RecoverableInputAdmission] = {} + self._token_by_context: dict[str, str] = {} + + def reserve(self, admission: _RecoverableInputAdmission) -> bool: + if self._root is None: + if admission.context_id in self._token_by_context: + return False + self._remember(admission) + return True + path = self._admission_path(admission.context_id) + with cross_process_file_lock(self._lock_path(admission.context_id)): + existing = self._load_document(path) + if existing is not None and not self._expired(existing): + return False + if not self._persisted_control_allows_recovery(admission): + return False + if existing is not None: + path.unlink(missing_ok=True) + atomic_write_json(path, admission.as_document()) + self._remember(admission) + return True + + def activate( + self, + admission: _RecoverableInputAdmission, + control_snapshot: dict[str, Any], + ) -> _RecoverableInputActivation | None: + """Fence one admission while publishing its running control.""" + if self._root is None: + if self._by_token.get(admission.token) != admission or self._expired(admission.as_document()): + return None + return _RecoverableInputActivation( + previous_control=None, + execution_id=str(control_snapshot["executionId"]), + ) + admission_path = self._admission_path(admission.context_id) + with cross_process_file_lock(self._lock_path(admission.context_id)): + document = self._load_document(admission_path) + if document is None or self._expired(document) or not self._matches(document, admission): + return None + # A different worker may have advanced the shared control after this + # process inspected its stale local controller. Re-check while holding + # the same fence used by every recovery reservation and activation. + if not self._persisted_control_allows_recovery(admission): + return None + revision = int(control_snapshot["revision"]) + persisted_snapshot = dict(control_snapshot) + persisted_snapshot["persistedRevision"] = revision + control_path = self._root / f"{admission.context_id}.json" + previous_control = self._load_document(control_path) + atomic_write_json(control_path, persisted_snapshot) + return _RecoverableInputActivation( + previous_control=previous_control, + execution_id=str(control_snapshot["executionId"]), + ) + + def finish_activation(self, admission: _RecoverableInputAdmission) -> bool: + if self._root is not None: + path = self._admission_path(admission.context_id) + with cross_process_file_lock(self._lock_path(admission.context_id)): + document = self._load_document(path) + if document is not None and self._matches(document, admission): + try: + path.unlink(missing_ok=True) + except OSError: + return False + self._forget(admission) + return True + + def rollback_activation( + self, + admission: _RecoverableInputAdmission, + activation: _RecoverableInputActivation, + ) -> bool: + if self._root is None: + return True + with cross_process_file_lock(self._lock_path(admission.context_id)): + control_path = self._root / f"{admission.context_id}.json" + current = self._load_document(control_path) + if current is None or current.get("executionId") != activation.execution_id: + return False + if activation.previous_control is None: + control_path.unlink(missing_ok=True) + else: + atomic_write_json(control_path, activation.previous_control) + return True + + def can_begin_without_admission( + self, + context_id: str, + local_execution_id: str | None, + local_server_instance_id: str, + local_input_handoff_ready: bool, + ) -> bool: + """Reject a new local controller when another process owns shared execution state.""" + if local_input_handoff_ready: + return False + if self._root is None: + return True + with cross_process_file_lock(self._lock_path(context_id)): + ticket = self._load_document(self._admission_path(context_id)) + if ticket is not None and not self._expired(ticket): + return False + control = self._load_document(self._root / f"{context_id}.json") + if control is None: + return not local_input_handoff_ready + if control.get("inputHandoffReady") is True: + return False + if local_execution_id is not None and ( + control.get("executionId") == local_execution_id + or control.get("serverInstanceId") == local_server_instance_id + ): + return True + return bool(control.get("phase") == "terminated" and control.get("releaseReady", False)) + + def has_active(self, context_id: str) -> bool: + if self._root is None: + return context_id in self._token_by_context + path = self._admission_path(context_id) + with cross_process_file_lock(self._lock_path(context_id)): + document = self._load_document(path) + if document is None: + return False + if self._expired(document): + path.unlink(missing_ok=True) + return False + return True + + def release(self, token: str) -> None: + admission = self._by_token.get(token) + if admission is None: + return + if self._root is not None: + path = self._admission_path(admission.context_id) + with cross_process_file_lock(self._lock_path(admission.context_id)): + document = self._load_document(path) + if document is not None and self._matches(document, admission): + path.unlink(missing_ok=True) + self._forget(admission) + + def get(self, token: str | None) -> _RecoverableInputAdmission | None: + return self._by_token.get(token or "") + + def close(self) -> None: + for token in tuple(self._by_token): + self.release(token) + + def _persisted_control_allows_recovery(self, admission: _RecoverableInputAdmission) -> bool: + assert self._root is not None + path = self._root / f"{admission.context_id}.json" + if not path.exists(): + return True + document = self._load_document(path) + if document is None or document.get("taskId") != admission.task_id: + return False + if document.get("inputHandoffReady") is True: + return True + if document.get("phase") != "terminated": + return False + backup = document.get("backup") + backup_status = backup.get("status") if isinstance(backup, dict) else None + return bool(document.get("releaseReady", False) or backup_status == "blocked") + + def _remember(self, admission: _RecoverableInputAdmission) -> None: + self._by_token[admission.token] = admission + self._token_by_context[admission.context_id] = admission.token + + def _forget(self, admission: _RecoverableInputAdmission) -> None: + self._by_token.pop(admission.token, None) + if self._token_by_context.get(admission.context_id) == admission.token: + self._token_by_context.pop(admission.context_id, None) + + def _admission_path(self, context_id: str) -> Path: + assert self._root is not None + return self._root / f".{context_id}.recoverable-input.json" + + def _lock_path(self, context_id: str) -> Path: + assert self._root is not None + return self._root / f".{context_id}.recoverable-input.lock" + + @staticmethod + def _load_document(path: Path) -> dict[str, Any] | None: + try: + value = json.loads(path.read_text(encoding="utf-8")) + except FileNotFoundError: + return None + except (OSError, ValueError): + return {"expiresAt": math.inf} + return value if isinstance(value, dict) else {"expiresAt": math.inf} + + @staticmethod + def _expired(document: dict[str, Any]) -> bool: + expires_at = document.get("expiresAt") + return isinstance(expires_at, (int, float)) and not isinstance(expires_at, bool) and expires_at <= time.time() + + @staticmethod + def _matches(document: dict[str, Any], admission: _RecoverableInputAdmission) -> bool: + return bool( + document.get("token") == admission.token + and document.get("contextId") == admission.context_id + and document.get("taskId") == admission.task_id + and document.get("owner") == admission.owner + ) + + +class RecoverableInputAdmissionLease: + """Transfer one recovery admission from the transport to the SDK producer.""" + + def __init__( + self, + token: str, + *, + acknowledge_enqueue: Callable[[str], bool], + release: Callable[[str], Awaitable[None]], + ) -> None: + self.token = token + self._acknowledge_enqueue = acknowledge_enqueue + self._release = release + self._enqueued = False + self._released = False + self._release_lock = asyncio.Lock() + self._producer_task: asyncio.Task[Any] | None = None + self._producer_done_callback: Callable[[asyncio.Task[Any]], None] | None = None + self._producer_cleanup: asyncio.Task[None] | None = None + + def acknowledge_enqueued(self, producer_task: asyncio.Task[Any] | None) -> None: + """Commit ownership only after the request queue accepted the context.""" + + if self._enqueued or not self._acknowledge_enqueue(self.token): + return + self._enqueued = True + if producer_task is None or producer_task.done(): + self._schedule_release() + return + self._producer_task = producer_task + self._producer_done_callback = lambda _task: self._schedule_release() + producer_task.add_done_callback(self._producer_done_callback) + + async def release(self) -> None: + """Release once after execution or an early producer failure.""" + + async with self._release_lock: + if self._released: + return + await self._release(self.token) + self._released = True + if self._producer_task is not None and self._producer_done_callback is not None: + self._producer_task.remove_done_callback(self._producer_done_callback) + self._producer_task = None + self._producer_done_callback = None + + def _schedule_release(self) -> None: + if self._released or (self._producer_cleanup is not None and not self._producer_cleanup.done()): + return + self._producer_cleanup = asyncio.create_task(self.release()) + + +class RecoverableInputAdmissionCarrier: + """Attach a server-issued admission lease to one queued SDK RequestContext.""" + + _ATTRIBUTE = "_iac_code_recoverable_input_admission" + + @classmethod + def attach(cls, request_context: Any, admission: str | RecoverableInputAdmissionLease | None) -> None: + setattr(request_context, cls._ATTRIBUTE, admission) + + @classmethod + def read(cls, request_context: Any) -> str | None: + admission = getattr(request_context, cls._ATTRIBUTE, None) + if isinstance(admission, RecoverableInputAdmissionLease): + return admission.token + return admission if isinstance(admission, str) and admission else None + + @classmethod + def acknowledge_enqueued(cls, request_context: Any, producer_task: asyncio.Task[Any] | None) -> None: + admission = getattr(request_context, cls._ATTRIBUTE, None) + if isinstance(admission, RecoverableInputAdmissionLease): + admission.acknowledge_enqueued(producer_task) + + @classmethod + async def release(cls, request_context: Any) -> None: + admission = getattr(request_context, cls._ATTRIBUTE, None) + if isinstance(admission, RecoverableInputAdmissionLease): + await admission.release() + + +class NaturalCompletionGenerationCarrier: + """Carry one executor turn generation to its response-stream boundary.""" + + _ATTRIBUTE = "_iac_code_natural_completion_generation" + _DELIVERED_ATTRIBUTE = "_iac_code_natural_completion_delivered" + _STATE_KEY = "iac_code_natural_completion_generation" + _DELIVERED_STATE_KEY = "iac_code_natural_completion_delivered" + + @classmethod + def prepare(cls, request_context: Any) -> None: + setattr(request_context, cls._ATTRIBUTE, None) + setattr(request_context, cls._DELIVERED_ATTRIBUTE, False) + call_context = getattr(request_context, "call_context", None) + state = getattr(call_context, "state", None) + if isinstance(state, dict): + state.pop(cls._STATE_KEY, None) + state.pop(cls._DELIVERED_STATE_KEY, None) + + @classmethod + def attach(cls, request_context: Any, generation: int | None) -> bool: + setattr(request_context, cls._ATTRIBUTE, generation) + call_context = getattr(request_context, "call_context", None) + state = getattr(call_context, "state", None) + if isinstance(state, dict): + state.pop(cls._STATE_KEY, None) + if generation is not None: + state[cls._STATE_KEY] = generation + return state.get(cls._DELIVERED_STATE_KEY) is True + return getattr(request_context, cls._DELIVERED_ATTRIBUTE, False) is True + + @classmethod + def mark_delivered(cls, context: Any) -> int | None: + setattr(context, cls._DELIVERED_ATTRIBUTE, True) + state = getattr(context, "state", None) + if isinstance(state, dict): + state[cls._DELIVERED_STATE_KEY] = True + return cls.read(context) + + @classmethod + def read(cls, context: Any) -> int | None: + generation = getattr(context, cls._ATTRIBUTE, None) + if isinstance(generation, int): + return generation + state = getattr(context, "state", None) + generation = state.get(cls._STATE_KEY) if isinstance(state, dict) else None + return generation if isinstance(generation, int) else None + + @dataclass class _Participant: participant_id: str @@ -123,14 +492,19 @@ def __init__( backup_service: Any | None, termination_cleanup: Callable[[str, str, str], Awaitable[str | None]] | None = None, on_resume: Callable[[str], None] | None = None, + input_handoff_commit: Callable[[ExecutionController, dict[str, Any]], Awaitable[None]] | None = None, execution_id: str | None = None, + execution_mode: str = "normal", ) -> None: + if execution_mode not in {"normal", "pipeline"}: + raise ValueError("execution_mode must be normal or pipeline") self.context_id = context_id self.task_id = task_id self.owner = owner self.cwd = cwd self.session_id: str | None = None self.execution_id = execution_id or "exec-" + uuid.uuid4().hex + self.execution_mode = execution_mode self.server_instance_id = server_instance_id self.phase: ExecutionPhase = "running" self.execution_status = "working" @@ -151,6 +525,7 @@ def __init__( self._condition = asyncio.Condition(self._lock) self._commit_lock = asyncio.Lock() self._operation_commit_lock = asyncio.Lock() + self._backup_lock = asyncio.Lock() self._activities: dict[str, _Activity] = {} self._participants: dict[str, _Participant] = {} self._participant_ids_by_task: dict[asyncio.Task[Any], str] = {} @@ -169,12 +544,48 @@ def __init__( self._backup_service = backup_service self._termination_cleanup = termination_cleanup self._on_resume = on_resume + self._input_handoff_commit = input_handoff_commit + self._durable_input_handoff_enabled = False self._termination_cleanup_complete = termination_cleanup is None self._termination_cleanup_inflight = False + self._turn_generation = 0 + self._turn_generation_by_task: dict[asyncio.Task[Any], int] = {} + self._natural_completion_generation: int | None = None + self._natural_completion_delivered_generation: int | None = None + self._pending_explicit_termination_reason: str | None = None + self._termination_generation = 0 def bind_session(self, session_id: str) -> None: self.session_id = session_id + def enable_durable_input_handoff(self) -> None: + self._durable_input_handoff_enabled = True + + def disable_durable_input_handoff(self) -> None: + self._durable_input_handoff_enabled = False + + def input_handoff_ready(self) -> bool: + return bool( + self._durable_input_handoff_enabled + and self.phase == "running" + and self.execution_status == "input-required" + and not self.stream_available + and not self.has_managed_work() + and not any(not task.done() for task in self._background_tasks) + ) + + def local_input_continuation_ready(self) -> bool: + """Return whether this process can continue its own drained input wait.""" + + return bool( + not self._durable_input_handoff_enabled + and self.phase == "running" + and self.execution_status == "input-required" + and not self.stream_available + and not self.has_managed_work() + and not any(not task.done() for task in self._background_tasks) + ) + def permission_suspension_allowed(self) -> bool: return self.phase == "running" and not self._connection_hold and self._resume_barrier_revision is None @@ -203,10 +614,17 @@ async def run_permission_suspension(self, operation: Callable[[], Awaitable[bool finally: await self.end_activity(activity_id) - async def attach_task(self, task: asyncio.Task[Any], *, mark_working: bool = True) -> None: + async def attach_task(self, task: asyncio.Task[Any], *, mark_working: bool = True) -> int: async with self._condition: if self.phase in {"terminating", "terminated"}: raise ExecutionControlConflictError("Execution is terminating") + generation = self._turn_generation_by_task.get(task) + if generation is None: + self._turn_generation += 1 + generation = self._turn_generation + self._turn_generation_by_task[task] = generation + self._natural_completion_generation = None + self._natural_completion_delivered_generation = None self._execution_tasks.add(task) participant_id = self._ensure_participant_locked(task, "execution") participant_ids = _CURRENT_PARTICIPANT_IDS.get() @@ -217,6 +635,7 @@ async def attach_task(self, task: asyncio.Task[Any], *, mark_working: bool = Tru self.execution_status = "working" self._invalidate_pause_commit_locked() self._condition.notify_all() + return generation def register_spawned_task(self, task: asyncio.Task[Any], *, kind: str) -> None: """Synchronously reserve a newly-created Task before it can be scheduled.""" @@ -265,21 +684,108 @@ async def rollover_normal_execution(self, *, task_id: str, cwd: str) -> None: self._pending_staged_backup = None self._termination_cleanup_complete = self._termination_cleanup is None self._termination_cleanup_inflight = False + self._natural_completion_generation = None + self._natural_completion_delivered_generation = None + self._pending_explicit_termination_reason = None self.revision += 1 snapshot = self.snapshot() await self._persist_snapshot(snapshot) - async def detach_task(self, task: asyncio.Task[Any], *, execution_status: str) -> None: + async def detach_task( + self, + task: asyncio.Task[Any], + *, + execution_status: str, + natural_completion: bool = False, + ) -> int | None: + handoff_snapshot: dict[str, Any] | None = None + handoff_commit = self._input_handoff_commit async with self._condition: + generation = self._turn_generation_by_task.pop(task, None) self._execution_tasks.discard(task) self._remove_participant_locked(task) if not self._execution_tasks: self.stream_available = False if execution_status != "working" or self.execution_status == "working": self.execution_status = execution_status + if ( + natural_completion + and execution_status in _TERMINAL_TASK_STATES + and generation is not None + and generation == self._turn_generation + ): + self._natural_completion_generation = generation + self._natural_completion_delivered_generation = None self._condition.notify_all() self._schedule_pause_commit_locked() self._maybe_mark_release_ready_locked() + if self.input_handoff_ready() and handoff_commit is not None: + self.revision += 1 + handoff_snapshot = self.snapshot() + if handoff_snapshot is not None and handoff_commit is not None: + await handoff_commit(self, handoff_snapshot) + return generation if natural_completion else None + + async def observe_state(self) -> dict[str, Any]: + """Observe state and retry a durable finalization that previously stopped short.""" + + claimed_snapshot: dict[str, Any] | None = None + termination_generation: int | None = None + async with self._condition: + if self._can_claim_natural_completion_locked(): + termination_generation = self._claim_natural_completion_locked() + claimed_snapshot = self.snapshot() + snapshot = self.snapshot() + self._retry_termination_if_needed_locked() + if claimed_snapshot is not None: + await self._persist_natural_completion_claim(claimed_snapshot) + assert termination_generation is not None + self._spawn( + self._finish_natural_completion(termination_generation), + "natural-completion", + ) + return snapshot + + async def finalize_natural_completion( + self, + *, + task_id: str, + completion_generation: int, + ) -> dict[str, Any]: + """Finalize a drained business turn after its response stream has been delivered.""" + + claimed_snapshot: dict[str, Any] | None = None + termination_generation: int | None = None + async with self._condition: + if task_id != self.task_id: + raise ExecutionControlConflictError("Execution identity does not match the current execution") + if completion_generation != self._natural_completion_generation: + snapshot = self.snapshot() + snapshot[NATURAL_COMPLETION_FINALIZED_GENERATION] = None + return snapshot + self._natural_completion_delivered_generation = completion_generation + if self._can_claim_natural_completion_locked(): + termination_generation = self._claim_natural_completion_locked() + claimed_snapshot = self.snapshot() + elif self.phase != "terminating" or self.termination_reason != "natural_completion": + snapshot = self.snapshot() + snapshot[NATURAL_COMPLETION_FINALIZED_GENERATION] = None + return snapshot + if claimed_snapshot is not None: + await self._persist_natural_completion_claim(claimed_snapshot) + assert termination_generation is not None + await await_fenced(self._finish_natural_completion(termination_generation)) + async with self._condition: + snapshot = self.snapshot() + snapshot[NATURAL_COMPLETION_FINALIZED_GENERATION] = completion_generation + return snapshot + + async def _persist_natural_completion_claim(self, snapshot: dict[str, Any]) -> None: + try: + await self._persist_snapshot(snapshot) + except Exception: + async with self._condition: + self._commit_error = "state_commit_failed" async def pause( self, @@ -365,6 +871,7 @@ async def resume( } fingerprint = _fingerprint("resume", payload) should_terminate = False + termination_generation: int | None = None async with self._condition: self._validate_target_locked(task_id=self.task_id, execution_id=execution_id) duplicate = self._idempotent_locked(request_id, "resume", fingerprint) @@ -390,7 +897,7 @@ async def resume( raise ExecutionControlConflictError("pauseId does not identify the active pause") loop = asyncio.get_running_loop() if self._expires_monotonic is not None and loop.time() >= self._expires_monotonic: - self._claim_termination_locked("disconnect_timeout") + termination_generation = self._claim_termination_locked("disconnect_timeout") should_terminate = True snapshot = self.snapshot() else: @@ -409,7 +916,8 @@ async def resume( ) self._condition.notify_all() if should_terminate: - self._spawn(self._finish_termination(), "deadline-termination") + assert termination_generation is not None + self._spawn(self._finish_termination(termination_generation), "deadline-termination") raise ExecutionControlConflictError("Reconnect deadline has expired; termination was claimed") return snapshot @@ -453,13 +961,30 @@ async def terminate( self._request_ids[request_id] = ("terminate", fingerprint, pause_id) self.connection_epoch = connection_epoch if self.phase == "terminated": + if self.termination_reason == "natural_completion": + # Natural release ends a response boundary, not the Task. + # A later explicit StopChat is a new terminal mutation and + # must publish a fresh canceled backup before release. + termination_generation = self._claim_termination_locked(reason) + snapshot = self.snapshot() + self._spawn( + self._finish_termination(termination_generation), + "post-natural-explicit-termination", + ) + return snapshot self._retry_termination_if_needed_locked() return self.snapshot() if self.phase == "terminating": + if self.termination_reason == "natural_completion": + # Serialize an explicit terminal mutation behind the + # in-flight non-canceling cleanup. The natural finalizer + # hands ownership over before publishing its backup. + if self._pending_explicit_termination_reason is None: + self._pending_explicit_termination_reason = reason return self.snapshot() - self._claim_termination_locked(reason) + termination_generation = self._claim_termination_locked(reason) snapshot = self.snapshot() - self._spawn(self._finish_termination(), "explicit-termination") + self._spawn(self._finish_termination(termination_generation), "explicit-termination") return snapshot async def checkpoint(self) -> None: @@ -625,6 +1150,7 @@ def snapshot(self) -> dict[str, Any]: "contextId": self.context_id, "taskId": self.task_id, "executionId": self.execution_id, + "owner": self.owner, "serverInstanceId": self.server_instance_id, "pauseId": self.pause_id, "pauseReason": self.pause_reason, @@ -642,8 +1168,20 @@ def snapshot(self) -> dict[str, Any]: "backup": dict(self.backup), "externalOperations": [dict(operation) for operation in self.external_operations], "releaseReady": self.release_ready, + "inputHandoffReady": self.input_handoff_ready(), + "localInputContinuationReady": self.local_input_continuation_ready(), } + @staticmethod + def protocol_snapshot(snapshot: dict[str, Any]) -> dict[str, Any]: + """Return the stable execution-control wire shape without coordination-only fields.""" + + public_snapshot = dict(snapshot) + public_snapshot.pop("owner", None) + public_snapshot.pop("inputHandoffReady", None) + public_snapshot.pop("localInputContinuationReady", None) + return public_snapshot + def has_managed_work(self) -> bool: return bool( any(not task.done() for task in self._execution_tasks) @@ -651,6 +1189,43 @@ def has_managed_work(self) -> bool: or any(not participant.task.done() for participant in self._participants.values()) ) + def can_admit_recoverable_input_continuation(self, task_id: str) -> bool: + """Return whether a sidecar-proven continuation can safely claim this control.""" + return bool( + self.task_id == task_id + and not self.has_managed_work() + and not any(not task.done() for task in self._background_tasks) + and ( + (self.phase == "terminated" and (self.release_ready or self.backup.get("status") == "blocked")) + or self.input_handoff_ready() + ) + ) + + def can_replace_with_recoverable_input_continuation(self, task_id: str) -> bool: + """Allow an admitted continuation to supersede a blocked terminal backup.""" + return bool( + self.can_admit_recoverable_input_continuation(task_id) + and (self.input_handoff_ready() or (not self.release_ready and self.backup.get("status") == "blocked")) + ) + + async def wait_until_recoverable_input_continuation( + self, + task_id: str, + *, + timeout: float, + ) -> None: + """Wait for the current execution's durable handoff or terminal release.""" + + async def wait() -> None: + async with self._condition: + await self._condition.wait_for( + lambda: self.task_id != task_id or self.can_admit_recoverable_input_continuation(task_id) + ) + if self.task_id != task_id: + raise ExecutionControlConflictError("Execution task changed before recovery") + + await asyncio.wait_for(wait(), timeout=timeout) + async def close(self) -> None: current = asyncio.current_task() managed_tasks = tuple( @@ -664,6 +1239,14 @@ async def close(self) -> None: if task is not current and not task.done() ) ) + if managed_tasks: + logger.warning( + "A2A execution control close canceling tasks context_id=%s execution_id=%s phase=%s task_names=%s", + sanitize_strict_text(self.context_id), + sanitize_strict_text(self.execution_id), + sanitize_strict_text(self.phase), + ",".join(sanitize_strict_text(task.get_name()) for task in managed_tasks), + ) for task in managed_tasks: task.cancel() if managed_tasks: @@ -906,10 +1489,14 @@ async def _deadline(self, generation: int, pause_id: str, delay: float) -> None: or not self._connection_hold ): return - self._claim_termination_locked("disconnect_timeout") - await self._finish_termination() - - def _claim_termination_locked(self, reason: str) -> None: + termination_generation = self._claim_termination_locked("disconnect_timeout") + await self._finish_termination(termination_generation) + + def _claim_termination_locked(self, reason: str) -> int: + self._natural_completion_generation = None + self._natural_completion_delivered_generation = None + self._pending_explicit_termination_reason = None + self._termination_generation += 1 self._termination_pause_id = self.pause_id if reason == "disconnect_timeout" else None self._pause_generation += 1 self._connection_hold = False @@ -921,8 +1508,57 @@ def _claim_termination_locked(self, reason: str) -> None: self._backup_state_committed = False self._termination_cleanup_complete = self._termination_cleanup is None self._termination_cleanup_inflight = False + self._pending_staged_backup = None self.release_ready = False self._condition.notify_all() + return self._termination_generation + + def _can_claim_natural_completion_locked(self) -> bool: + return bool( + self._natural_completion_generation is not None + and self._natural_completion_delivered_generation == self._natural_completion_generation + and self.phase == "running" + and self._resume_barrier_revision is None + and self.execution_status in _TERMINAL_TASK_STATES + and not self.has_managed_work() + ) + + def _claim_natural_completion_locked(self) -> int: + """Atomically claim non-canceling finalization for a drained business turn.""" + + self._natural_completion_generation = None + self._natural_completion_delivered_generation = None + self._pending_explicit_termination_reason = None + self._termination_generation += 1 + self._termination_pause_id = None + self._pause_generation += 1 + self._connection_hold = False + self._resume_barrier_revision = None + self.pause_id = None + self.pause_reason = None + self.expires_at = None + self._expires_monotonic = None + self.phase = "terminating" + self.revision += 1 + self.termination_reason = "natural_completion" + self.backup = {"status": "pending"} + self._backup_state_committed = False + self._termination_cleanup_complete = self._termination_cleanup is None + self._termination_cleanup_inflight = False + self._pending_staged_backup = None + self.release_ready = False + logger.info( + "A2A execution natural completion claimed context_id=%s task_id=%s execution_id=%s status=%s", + sanitize_strict_text(self.context_id), + sanitize_strict_text(self.task_id), + sanitize_strict_text(self.execution_id), + sanitize_strict_text(self.execution_status), + ) + self._condition.notify_all() + return self._termination_generation + + def _owns_termination_locked(self, generation: int) -> bool: + return self._termination_generation == generation and self.phase in {"terminating", "terminated"} def _retry_termination_if_needed_locked(self) -> None: if self.phase != "terminated" or self.release_ready: @@ -932,22 +1568,31 @@ def _retry_termination_if_needed_locked(self) -> None: self._termination_cleanup_inflight = True self.backup = {"status": "pending"} self._backup_state_committed = False - self._spawn(self._retry_termination_cleanup(), "retry-termination-cleanup") + self._spawn( + self._retry_termination_cleanup(self._termination_generation), + "retry-termination-cleanup", + ) return if self.backup.get("status") == "blocked": self.backup = {"status": "pending"} self._backup_state_committed = False - self._spawn(self._retry_backup(), "retry-termination-backup") + self._spawn(self._retry_backup(self._termination_generation), "retry-termination-backup") elif self._commit_error is not None: self._commit_error = None if self._backup_state_committed: self._maybe_mark_release_ready_locked() else: - self._spawn(self._retry_backup_state_commit(), "retry-backup-state-commit") + self._spawn( + self._retry_backup_state_commit(self._termination_generation), + "retry-backup-state-commit", + ) - async def _finish_termination(self) -> None: + async def _finish_termination(self, generation: int) -> None: current = asyncio.current_task() async with self._condition: + if not self._owns_termination_locked(generation): + return + termination_reason = self.termination_reason or "terminated" cleanup = None if not self._termination_cleanup_complete and not self._termination_cleanup_inflight: self._termination_cleanup_inflight = True @@ -963,8 +1608,24 @@ async def _finish_termination(self) -> None: for participant in self._participants.values() if participant.task is not current and not participant.task.done() ) - cleanup_error = await self._run_termination_cleanup(cleanup, mark_complete=False) + cleanup_error = await self._run_termination_cleanup( + cleanup, + generation=generation, + termination_reason=termination_reason, + mark_complete=False, + ) + async with self._condition: + if not self._owns_termination_locked(generation): + return tasks = tuple(dict.fromkeys((*execution_tasks, *activity_tasks, *participant_tasks))) + if tasks: + logger.warning( + "A2A execution termination canceling tasks context_id=%s execution_id=%s reason=%s task_names=%s", + sanitize_strict_text(self.context_id), + sanitize_strict_text(self.execution_id), + sanitize_strict_text(self.termination_reason or "none"), + ",".join(sanitize_strict_text(task.get_name()) for task in tasks), + ) for task in tasks: task.cancel() # A transport may reuse its producer Task after execute() returns (the @@ -978,8 +1639,15 @@ async def _finish_termination(self) -> None: if owned_tasks: await asyncio.gather(*owned_tasks, return_exceptions=True) if cleanup_error is None: - cleanup_error = await self._run_termination_cleanup(cleanup, mark_complete=True) + cleanup_error = await self._run_termination_cleanup( + cleanup, + generation=generation, + termination_reason=termination_reason, + mark_complete=True, + ) async with self._condition: + if not self._owns_termination_locked(generation): + return for task in tuple(self._participant_ids_by_task): if task.done(): self._remove_participant_locked(task) @@ -998,24 +1666,100 @@ async def _finish_termination(self) -> None: await self._commit_blocked_backup_state( "execution termination cleanup failed", cleanup_error, + generation=generation, + clear_cleanup_inflight=True, + ) + return + await self._perform_backup(generation) + + async def _finish_natural_completion(self, generation: int) -> None: + """Finalize a drained execution without canceling any business task.""" + + async with self._condition: + if ( + not self._owns_termination_locked(generation) + or self.phase != "terminating" + or self.termination_reason != "natural_completion" + ): + return + cleanup = None + if not self._termination_cleanup_complete and not self._termination_cleanup_inflight: + self._termination_cleanup_inflight = True + cleanup = self._termination_cleanup + cleanup_error = await self._run_termination_cleanup( + cleanup, + generation=generation, + termination_reason="natural_completion", + mark_complete=True, + ) + async with self._condition: + if ( + not self._owns_termination_locked(generation) + or self.phase != "terminating" + or self.termination_reason != "natural_completion" + ): + return + pending_explicit_reason = self._pending_explicit_termination_reason + if pending_explicit_reason is not None: + explicit_generation = self._claim_termination_locked(pending_explicit_reason) + self._spawn( + self._finish_termination(explicit_generation), + "post-natural-explicit-termination", + ) + return + if self.has_managed_work(): + cleanup_error = cleanup_error or RuntimeError("Natural completion acquired new managed work") + self.stream_available = False + self.phase = "terminated" + self.revision += 1 + terminated = self.snapshot() + try: + await self._persist_snapshot(terminated) + except Exception: + async with self._condition: + self._commit_error = "state_commit_failed" + if cleanup_error is not None: + await self._commit_blocked_backup_state( + "execution natural completion cleanup failed", + cleanup_error, + generation=generation, clear_cleanup_inflight=True, ) return - await self._perform_backup() + await self._perform_backup(generation) + async with self._condition: + await self._condition.wait_for( + lambda: self.release_ready or self.backup.get("status") == "blocked" or self._commit_error is not None + ) + logger.info( + "A2A execution natural completion settled context_id=%s task_id=%s execution_id=%s " + "status=%s backup_status=%s release_ready=%s commit_error=%s", + sanitize_strict_text(self.context_id), + sanitize_strict_text(self.task_id), + sanitize_strict_text(self.execution_id), + sanitize_strict_text(self.execution_status), + sanitize_strict_text(str(self.backup.get("status") or "none")), + self.release_ready, + sanitize_strict_text(self._commit_error or "none"), + ) async def _run_termination_cleanup( self, cleanup: Callable[[str, str, str], Awaitable[str | None]] | None, *, + generation: int, + termination_reason: str, mark_complete: bool = True, ) -> BaseException | None: if cleanup is None: return None try: - execution_status = await cleanup(self.context_id, self.task_id, self.termination_reason or "terminated") + execution_status = await cleanup(self.context_id, self.task_id, termination_reason) except Exception as exc: return exc async with self._condition: + if not self._owns_termination_locked(generation): + return None if execution_status is not None: self.execution_status = execution_status if mark_complete: @@ -1023,28 +1767,54 @@ async def _run_termination_cleanup( self._termination_cleanup_inflight = False return None - async def _retry_termination_cleanup(self) -> None: - error = await self._run_termination_cleanup(self._termination_cleanup) + async def _retry_termination_cleanup(self, generation: int) -> None: + async with self._condition: + if not self._owns_termination_locked(generation): + return + termination_reason = self.termination_reason or "terminated" + error = await self._run_termination_cleanup( + self._termination_cleanup, + generation=generation, + termination_reason=termination_reason, + ) if error is not None: await self._commit_blocked_backup_state( "execution termination cleanup failed", error, + generation=generation, clear_cleanup_inflight=True, ) return - await self._perform_backup() + await self._perform_backup(generation) + + async def _retry_backup(self, generation: int) -> None: + await self._perform_backup(generation) - async def _retry_backup(self) -> None: - await self._perform_backup() + async def _perform_backup(self, generation: int) -> None: + async with self._backup_lock: + async with self._condition: + if not self._owns_termination_locked(generation): + return + await self._perform_backup_serialized(generation) - async def _perform_backup(self) -> None: + async def _perform_backup_serialized(self, generation: int) -> None: try: await self._persist_external_operations() except Exception as exc: - await self._commit_blocked_backup_state("external operation state could not be persisted", exc) + await self._commit_blocked_backup_state( + "external operation state could not be persisted", + exc, + generation=generation, + ) return + async with self._condition: + if not self._owns_termination_locked(generation): + return + termination_reason = self.termination_reason if self._backup_service is None or self.session_id is None: async with self._condition: + if not self._owns_termination_locked(generation): + return self.backup = {"status": "disabled"} self._backup_state_committed = False self.release_ready = False @@ -1055,13 +1825,17 @@ async def _perform_backup(self) -> None: await self._persist_snapshot(snapshot) except Exception: async with self._condition: - self._commit_error = "state_commit_failed" + if self._owns_termination_locked(generation): + self._commit_error = "state_commit_failed" return async with self._condition: + if not self._owns_termination_locked(generation): + return self._commit_error = None self._backup_state_committed = True self._maybe_mark_release_ready_locked() return + result: BackupResult | None = None try: result = self._pending_staged_backup if result is None: @@ -1071,7 +1845,7 @@ async def _perform_backup(self) -> None: self.session_id, reason=( BackupReason.DISCONNECT_TIMEOUT - if self.termination_reason == "disconnect_timeout" + if termination_reason == "disconnect_timeout" else BackupReason.TERMINAL ), critical=True, @@ -1091,7 +1865,6 @@ async def _perform_backup(self) -> None: ), ) succeeded = (not result.enabled) or (result.succeeded and result.shared_committed) - self._pending_staged_backup = None if succeeded else result if result.staged_committed else None backup = { "status": "disabled" if not result.enabled else "shared_committed" if succeeded else "blocked", "generation": result.generation, @@ -1102,6 +1875,11 @@ async def _perform_backup(self) -> None: succeeded = False backup = {"status": "blocked", "error": str(exc)} async with self._condition: + if not self._owns_termination_locked(generation): + return + self._pending_staged_backup = ( + None if succeeded or result is None else result if result.staged_committed else None + ) self.backup = backup self._backup_state_committed = False self.release_ready = False @@ -1112,15 +1890,20 @@ async def _perform_backup(self) -> None: await self._persist_snapshot(snapshot) except Exception: async with self._condition: - self._commit_error = "state_commit_failed" + if self._owns_termination_locked(generation): + self._commit_error = "state_commit_failed" return if succeeded: async with self._condition: + if not self._owns_termination_locked(generation): + return self._commit_error = None self._backup_state_committed = True self._maybe_mark_release_ready_locked() else: async with self._condition: + if not self._owns_termination_locked(generation): + return self._commit_error = None self._backup_state_committed = True @@ -1129,9 +1912,12 @@ async def _commit_blocked_backup_state( message: str, exc: BaseException, *, + generation: int, clear_cleanup_inflight: bool = False, ) -> None: async with self._condition: + if not self._owns_termination_locked(generation): + return self.backup = {"status": "blocked", "error": f"{message}: {type(exc).__name__}"} if clear_cleanup_inflight: # A retry that observes the blocked state must also be able to @@ -1148,11 +1934,19 @@ async def _commit_blocked_backup_state( await self._persist_snapshot(snapshot) except Exception: async with self._condition: - if self.revision == revision and self.backup.get("status") == "blocked": + if ( + self._owns_termination_locked(generation) + and self.revision == revision + and self.backup.get("status") == "blocked" + ): self._commit_error = "state_commit_failed" return async with self._condition: - if self.revision == revision and self.backup.get("status") == "blocked": + if ( + self._owns_termination_locked(generation) + and self.revision == revision + and self.backup.get("status") == "blocked" + ): self._commit_error = None self._backup_state_committed = True @@ -1170,8 +1964,10 @@ async def _persist_external_operations(self) -> None: path = SessionStorage().session_dir(self.cwd, self.session_id) / "a2a" / "external-operations.json" await run_sync_fenced(atomic_write_json, path, document) - async def _retry_backup_state_commit(self) -> None: + async def _retry_backup_state_commit(self, generation: int) -> None: async with self._condition: + if not self._owns_termination_locked(generation): + return snapshot = self.snapshot() snapshot["commitError"] = None backup_complete = self.backup.get("status") in {"disabled", "shared_committed"} @@ -1179,9 +1975,12 @@ async def _retry_backup_state_commit(self) -> None: await self._persist_snapshot(snapshot) except Exception: async with self._condition: - self._commit_error = "state_commit_failed" + if self._owns_termination_locked(generation): + self._commit_error = "state_commit_failed" return async with self._condition: + if not self._owns_termination_locked(generation): + return self._commit_error = None self._backup_state_committed = True if backup_complete: @@ -1227,10 +2026,18 @@ def completed(done: asyncio.Task[Any]) -> None: if not done.cancelled(): with suppress(BaseException): done.exception() + try: + done.get_loop().create_task(self._notify_recovery_waiters()) + except RuntimeError: + pass task.add_done_callback(completed) return task + async def _notify_recovery_waiters(self) -> None: + async with self._condition: + self._condition.notify_all() + class ExecutionControlService: def __init__(self, *, persistence_root: Path | None, backup_service: Any | None) -> None: @@ -1239,6 +2046,7 @@ def __init__(self, *, persistence_root: Path | None, backup_service: Any | None) self._backup_service = backup_service self._controls: dict[str, ExecutionController] = {} self._context_start_locks: dict[str, asyncio.Lock] = {} + self._recoverable_input_admissions = _RecoverableInputAdmissionStore(persistence_root) self._termination_cleanup: Callable[[str, str, str], Awaitable[str | None]] | None = None self._on_resume: Callable[[str], None] | None = None @@ -1255,6 +2063,21 @@ def set_termination_cleanup( for control in self._controls.values(): control._termination_cleanup = cleanup + async def _commit_input_handoff( + self, + control: ExecutionController, + snapshot: dict[str, Any], + ) -> None: + """Publish a drained recovered input wait before another request can start locally.""" + async with self._context_start_locks.setdefault(control.context_id, asyncio.Lock()): + if ( + self._controls.get(control.context_id) is not control + or control.revision != snapshot.get("revision") + or not control.input_handoff_ready() + ): + return + await control._persist_snapshot(snapshot) + async def begin_execution( self, *, @@ -1262,7 +2085,9 @@ async def begin_execution( task_id: str, owner: str, cwd: str, + execution_mode: str = "normal", continue_input_required: bool = False, + recoverable_input_admission: str | None = None, ) -> ExecutionController: task = asyncio.current_task() if task is None: @@ -1272,18 +2097,37 @@ async def begin_execution( control = self._controls.get(context_id) if control is not None and control.owner != owner: raise ExecutionControlNotFoundError("Execution was not found") + admission = self._recoverable_input_admissions.get(recoverable_input_admission) + admitted_recovery = bool( + admission is not None + and admission.context_id == context_id + and admission.task_id == task_id + and admission.owner == owner + ) + if recoverable_input_admission is not None and not admitted_recovery: + raise ExecutionControlConflictError("Recoverable input continuation admission is stale") + replace_blocked_input_wait = bool( + control is not None + and admitted_recovery + and control.can_replace_with_recoverable_input_continuation(task_id) + ) if control is not None and ( ( control.task_id != task_id and control.phase in {"pausing", "pause_committing", "paused", "resuming", "terminating"} ) - or (control.phase in {"terminating", "terminated"} and not control.release_ready) + or ( + control.phase in {"terminating", "terminated"} + and not control.release_ready + and not replace_blocked_input_wait + ) ): raise ExecutionControlConflictError("Current execution must finish recovery before a new task starts") reuse = bool( control is not None and control.task_id == task_id and control.phase not in {"terminating", "terminated"} + and not control.input_handoff_ready() and ( control.stream_available or (continue_input_required and control.execution_status == "input-required") @@ -1299,6 +2143,16 @@ async def begin_execution( ): await control.rollover_normal_execution(task_id=task_id, cwd=cwd) reuse = True + if not reuse and admission is None: + shared_start_allowed = await run_sync_fenced( + self._recoverable_input_admissions.can_begin_without_admission, + context_id, + control.execution_id if control is not None else None, + self.server_instance_id, + control.input_handoff_ready() if control is not None else False, + ) + if not shared_start_allowed: + raise ExecutionControlConflictError("Current execution is active in another process") if not reuse: retired_control = control path = None @@ -1314,14 +2168,163 @@ async def begin_execution( backup_service=self._backup_service, termination_cleanup=self._termination_cleanup, on_resume=self._on_resume, + input_handoff_commit=self._commit_input_handoff, + execution_mode=execution_mode, ) - self._controls[context_id] = control assert control is not None + if admission is not None and not reuse: + control.enable_durable_input_handoff() + control.revision += 1 + activation: _RecoverableInputActivation | None = None + + async def cleanup_cancelled_activation( + result: _RecoverableInputActivation | None, + error: BaseException | None, + ) -> None: + if result is not None and error is None: + await run_sync_fenced( + self._recoverable_input_admissions.rollback_activation, + admission, + result, + ) + control.disable_durable_input_handoff() + await control.detach_task(task, execution_status="input-required") + await control.close() + + await control.attach_task(task, mark_working=False) + try: + activation = await run_sync_fenced_with_cancel_completion( + self._recoverable_input_admissions.activate, + cleanup_cancelled_activation, + admission, + control.snapshot(), + ) + if activation is None: + raise ExecutionControlConflictError("Recoverable input continuation admission is stale") + control.persisted_revision = control.revision + self._controls[context_id] = control + await run_sync_fenced(self._recoverable_input_admissions.finish_activation, admission) + except asyncio.CancelledError: + if activation is not None: + await run_sync_fenced( + self._recoverable_input_admissions.rollback_activation, + admission, + activation, + ) + if retired_control is None: + self._controls.pop(context_id, None) + else: + self._controls[context_id] = retired_control + control.disable_durable_input_handoff() + await control.detach_task(task, execution_status="input-required") + await control.close() + raise + except BaseException: + if activation is not None: + await run_sync_fenced( + self._recoverable_input_admissions.rollback_activation, + admission, + activation, + ) + if retired_control is None: + self._controls.pop(context_id, None) + else: + self._controls[context_id] = retired_control + control.disable_durable_input_handoff() + await control.detach_task(task, execution_status="input-required") + await control.close() + raise + self._controls[context_id] = control if retired_control is not None: await retired_control.close() - await control.attach_task(task, mark_working=False) + if admission is None or reuse: + await control.attach_task(task, mark_working=False) return control + async def reserve_recoverable_input_continuation( + self, + *, + context_id: str, + task_id: str, + owner: str, + ) -> str | None: + """Reserve one request-scoped continuation after the caller proves the sidecar wait.""" + async with self._context_start_locks.setdefault(context_id, asyncio.Lock()): + control = self._controls.get(context_id) + if control is not None and ( + control.owner != owner or not control.can_admit_recoverable_input_continuation(task_id) + ): + return None + if control is not None and control.input_handoff_ready(): + # A previous detach may have completed locally while its durable + # handoff write failed. Re-publish the same revision before a + # request can reserve and replace this controller. + await control._persist_snapshot(control.snapshot()) + token = "recovery-" + uuid.uuid4().hex + admission = _RecoverableInputAdmission( + token=token, + context_id=context_id, + task_id=task_id, + owner=owner, + expires_at=time.time() + _RECOVERABLE_INPUT_ADMISSION_TTL_SECONDS, + ) + reserved = await run_sync_fenced( + self._recoverable_input_admissions.reserve, + admission, + ) + return token if reserved else None + + async def wait_until_recoverable_input_continuation( + self, + *, + context_id: str, + task_id: str, + timeout: float, + ) -> None: + """Wait until the local controller can safely yield to one recovery request.""" + + control = self._controls.get(context_id) + if control is None: + return + await control.wait_until_recoverable_input_continuation(task_id, timeout=timeout) + + async def finalize_natural_completion( + self, + *, + context_id: str, + task_id: str, + owner: str, + completion_generation: int, + finalized_cleanup: Callable[[], Awaitable[None]] | None = None, + ) -> dict[str, Any] | None: + """Settle a naturally exhausted response without invoking termination.""" + + async with self._context_start_locks.setdefault(context_id, asyncio.Lock()): + control = self._controls.get(context_id) + if control is None: + return None + if control.owner != owner: + raise ExecutionControlNotFoundError("Execution was not found") + state = await control.finalize_natural_completion( + task_id=task_id, + completion_generation=completion_generation, + ) + if ( + finalized_cleanup is not None + and state.get("terminationReason") == "natural_completion" + and state.get(NATURAL_COMPLETION_FINALIZED_GENERATION) == completion_generation + ): + await finalized_cleanup() + return state + + async def release_recoverable_input_continuation(self, token: str) -> None: + """Release an unused recovery reservation; consumed reservations are a no-op.""" + admission = self._recoverable_input_admissions.get(token) + if admission is None: + return + async with self._context_start_locks.setdefault(admission.context_id, asyncio.Lock()): + await run_sync_fenced(self._recoverable_input_admissions.release, token) + def get_for_context(self, context_id: str) -> ExecutionController | None: return self._controls.get(context_id) @@ -1331,9 +2334,34 @@ async def require(self, *, context_id: str, owner: str) -> ExecutionController: raise ExecutionControlNotFoundError("Execution was not found") return control + async def observe(self, *, context_id: str, owner: str) -> dict[str, Any]: + """Observe an active controller or its read-only persisted terminal state.""" + + control = self._controls.get(context_id) + if control is not None: + if control.owner != owner: + raise ExecutionControlNotFoundError("Execution was not found") + return control.protocol_snapshot(await control.observe_state()) + snapshot = await run_sync_fenced(self._load_persisted_terminal_snapshot, context_id) + if snapshot is None or snapshot.get("owner") != owner: + raise ExecutionControlNotFoundError("Execution was not found") + return ExecutionController.protocol_snapshot(snapshot) + def snapshot_for_context(self, context_id: str) -> dict[str, Any] | None: control = self._controls.get(context_id) - return control.snapshot() if control is not None else None + return control.snapshot() if control is not None else self._load_persisted_terminal_snapshot(context_id) + + def _load_persisted_terminal_snapshot(self, context_id: str) -> dict[str, Any] | None: + if self._persistence_root is None: + return None + path = self._persistence_root / "execution-control" / f"{context_id}.json" + try: + value = json.loads(path.read_text(encoding="utf-8")) + except (FileNotFoundError, OSError, ValueError): + return None + if not isinstance(value, dict) or value.get("contextId") != context_id or value.get("phase") != "terminated": + return None + return value def has_active_work(self) -> bool: active_phases = {"pausing", "pause_committing", "paused", "resuming", "terminating"} @@ -1346,6 +2374,7 @@ def has_active_work(self) -> bool: async def close(self) -> None: await asyncio.gather(*(control.close() for control in tuple(self._controls.values())), return_exceptions=True) + await run_sync_fenced(self._recoverable_input_admissions.close) def bind_execution_control(control: ExecutionController | None) -> Token[ExecutionController | None]: diff --git a/src/iac_code/a2a/executor.py b/src/iac_code/a2a/executor.py index 574e0a99..463e528b 100644 --- a/src/iac_code/a2a/executor.py +++ b/src/iac_code/a2a/executor.py @@ -31,6 +31,8 @@ ) from iac_code.a2a.execution_control import ( ExecutionControlService, + NaturalCompletionGenerationCarrier, + RecoverableInputAdmissionCarrier, bind_execution_control, clear_execution_participants, current_execution_control, @@ -64,6 +66,7 @@ WaitingInputCancelResult, cancel_waiting_input_task_from_sidecar, recoverable_task_id_from_sidecar, + sandbox_release_recoverable_task_id_from_sidecar, terminal_task_state_from_sidecar, ) from iac_code.a2a.pipeline_journal import A2APipelineJournal @@ -76,6 +79,10 @@ from iac_code.a2a.pipeline_stream import BACKUP_COMMITTED_EVENT_TYPE, PipelineA2AEventPublisher from iac_code.a2a.projection import a2a_safe_mode_enabled from iac_code.a2a.request_mode import resolve_request_run_mode +from iac_code.a2a.request_scoped_active_task import ( + DirectPipelineRouteGateCarrier, + PipelineLifecycleEventQueueCarrier, +) from iac_code.a2a.resource_selector import ( PendingResourceSelection, ResourceSelectionCheckpointStore, @@ -97,6 +104,7 @@ from iac_code.a2a.thinking_metadata import A2AThinkingMetadata from iac_code.a2a.types import ( TASK_STATE_CANCELED, + TASK_STATE_COMPLETED, TASK_STATE_FAILED, TASK_STATE_INPUT_REQUIRED, TASK_STATE_WORKING, @@ -1242,6 +1250,10 @@ def _string_value(value: Any) -> str: class IacCodeA2AExecutor(AgentExecutor): + _RECOVERABLE_RELEASE_REASONS = frozenset( + {"disconnect_timeout", "client_disconnect_control_timeout", "natural_completion"} + ) + def __init__( self, *, @@ -1281,6 +1293,44 @@ def __init__( execution_control_service.set_termination_cleanup(self._terminate_detached_execution) execution_control_service.set_resume_callback(self._task_store.touch_context) + async def wait_until_recoverable_pipeline_input(self, *, context_id: str, task_id: str) -> None: + if self._execution_control_service is None: + return + try: + await self._execution_control_service.wait_until_recoverable_input_continuation( + context_id=context_id, + task_id=task_id, + timeout=30, + ) + except TimeoutError as exc: + raise InvalidParamsError("Pipeline continuation is still finalizing; retry the same request.") from exc + + async def finalize_natural_execution( + self, + *, + context_id: str, + task_id: str, + owner: str, + completion_generation: int, + ) -> None: + """Finalize a delivered response without canceling its business task.""" + + if self._execution_control_service is None: + return + + async def release_reversible_permission_closes() -> None: + closing_tokens = await self._permission_input_registry.reversible_closing_tokens(task_id) + for closing_token in closing_tokens: + await self._permission_input_registry.reopen_task(closing_token) + + await self._execution_control_service.finalize_natural_completion( + context_id=context_id, + task_id=task_id, + owner=owner, + completion_generation=completion_generation, + finalized_cleanup=release_reversible_permission_closes, + ) + async def resolve_sideband_permission( self, response: PermissionResponse, *, metadata: Any = None ) -> Message | None: @@ -1324,7 +1374,8 @@ async def execute(self, context: RequestContext, event_queue: EventQueue) -> Non owner = self._task_store.owner_for_context(getattr(context, "call_context", None)) current_task = asyncio.current_task() if existing is not None and existing.owner == owner and current_task is not None: - await existing.attach_task(current_task, mark_working=False) + await existing.attach_task(current_task, mark_working=True) + await existing.mark_execution_started() bind_execution_control(existing) requested_llm_headers = resolve_a2a_llm_headers(metadata) with contextlib.ExitStack() as request_scope: @@ -1376,16 +1427,38 @@ async def activate_bound_llm_headers() -> None: control = current_execution_control() if control is not None: execution_status = "unknown" + natural_completion = False try: record = await self._task_store.get_task_record(control.task_id) execution_status = record.state + natural_completion = execution_status in { + TASK_STATE_INPUT_REQUIRED, + TASK_STATE_COMPLETED, + TASK_STATE_FAILED, + } and not await self._permission_input_registry.has_pending_task(control.task_id) except ValueError: pass current_task = asyncio.current_task() if current_task is not None: - await control.detach_task(current_task, execution_status=execution_status) + completion_generation = await control.detach_task( + current_task, + execution_status=execution_status, + natural_completion=natural_completion, + ) + response_already_delivered = NaturalCompletionGenerationCarrier.attach( + context, + completion_generation, + ) + if completion_generation is not None and response_already_delivered: + await self.finalize_natural_execution( + context_id=context_id, + task_id=control.task_id, + owner=self._task_store.owner_for_context(getattr(context, "call_context", None)), + completion_generation=completion_generation, + ) reset_execution_control(execution_scope) reset_execution_participants(participant_scope) + await RecoverableInputAdmissionCarrier.release(context) async def _execute( self, @@ -1426,6 +1499,17 @@ async def activate_llm_headers() -> None: return permission_response = parse_permission_response(getattr(context, "message", None)) if permission_response is not None: + if ( + PipelineLifecycleEventQueueCarrier.read(context) + and not PipelineLifecycleEventQueueCarrier.is_bound(context) + ): + await self._publish_status( + event_queue, + task_id=permission_response.task_id, + context_id=permission_response.context_id, + state=TaskState.TASK_STATE_WORKING, + ) + PipelineLifecycleEventQueueCarrier.mark_bound(context) response_metadata = getattr(context, "metadata", None) or getattr( getattr(context, "message", None), "metadata", None ) @@ -1443,6 +1527,7 @@ async def activate_llm_headers() -> None: with a2a_request_context(aliyun_credential=response_credential): approved = await self._permission_input_registry.answer( permission_response, + aliyun_credential=response_credential, before_delivery=commit_llm_headers, ) await activate_bound_llm_headers() @@ -1648,16 +1733,49 @@ async def release_context_execution() -> None: restore_interrupted=not pipeline_mode, ) if self._execution_control_service is not None: + recoverable_input_admission = RecoverableInputAdmissionCarrier.read(context) control = await self._execution_control_service.begin_execution( context_id=context_id, task_id=task.task_id, owner=owner, cwd=cwd, + execution_mode="pipeline" if pipeline_mode and not route_pipeline_handoff_to_normal else "normal", continue_input_required=pipeline_mode and not route_pipeline_handoff_to_normal, + recoverable_input_admission=recoverable_input_admission, ) bind_execution_control(control) await control.checkpoint() await control.mark_execution_started() + current_task = getattr(context, "current_task", None) + needs_lifecycle_binding_frame = ( + pipeline_mode + and not route_pipeline_handoff_to_normal + and PipelineLifecycleEventQueueCarrier.read(context) + and ( + recoverable_input_admission is not None + or ( + isinstance(current_task, Task) + and current_task.status is not None + and current_task.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + ) + ) + ) + if needs_lifecycle_binding_frame: + # A cold recovered request reuses the SDK's existing + # INPUT_REQUIRED Task, so there is no automatic initial + # Task frame. Bind ROS to the newly started execution + # before Pipeline context/sidecar restoration can finish + # without producing a public event. Same-process reuse + # (localInputContinuationReady) legitimately reenters + # without a dispatcher-signed admission but hits the + # same empty-first-frame window, so publish here too. + await self._publish_status( + event_queue, + task_id=task_id, + context_id=context_id, + state=TaskState.TASK_STATE_WORKING, + ) + PipelineLifecycleEventQueueCarrier.mark_bound(context) await publish_initial_task_if_missing() await self._task_store.ensure_task_not_expired(task.task_id) except InvalidParamsError: @@ -1742,16 +1860,30 @@ async def release_context_execution() -> None: resource_selector_enabled=resource_selector_enabled, ) try: - pipeline_result = await pipeline_executor.execute( - context=context, - event_queue=event_queue, - task=task, - task_id=task_id, - context_id=context_id, - cwd=cwd, - pipeline_input=pipeline_input, - active_followup_only=active_pipeline_owner is not None, - ) + direct_route_gate = DirectPipelineRouteGateCarrier.read(context) + if direct_route_gate is None: + pipeline_result = await pipeline_executor.execute( + context=context, + event_queue=event_queue, + task=task, + task_id=task_id, + context_id=context_id, + cwd=cwd, + pipeline_input=pipeline_input, + active_followup_only=active_pipeline_owner is not None, + ) + else: + pipeline_result = await pipeline_executor.execute( + context=context, + event_queue=event_queue, + task=task, + task_id=task_id, + context_id=context_id, + cwd=cwd, + pipeline_input=pipeline_input, + active_followup_only=active_pipeline_owner is not None, + direct_route_gate=direct_route_gate, + ) if active_pipeline_owner is not None and pipeline_result is False: owner_finished = active_pipeline_owner.done() if owner_finished: @@ -2904,6 +3036,8 @@ async def _terminate_detached_execution(self, context_id: str, task_id: str, rea had_pending_permission = await self._permission_input_registry.has_pending_task(task_id) had_pending_resource_selection = await self._resource_selection_registry.has_pending_task(task_id) if had_pending_permission: + if reason == "natural_completion": + raise RuntimeError("Natural completion cannot release an active permission wait") await self._permission_input_registry.cancel_task(task_id) await self._task_store.set_pending_permissions(task_id, []) await self._task_store.discard_context_runtime(context_id, persist_context=False) @@ -2933,9 +3067,64 @@ async def _terminate_detached_execution(self, context_id: str, task_id: str, rea canceled=task_record.state == TASK_STATE_CANCELED, ) if task_record.state not in {TASK_STATE_INPUT_REQUIRED, TASK_STATE_CANCELED}: - if not await self._task_store.commit_inactive_execution_task(task_id=task_id, context_id=context_id): + if not await self._task_store.commit_inactive_execution_task( + task_id=task_id, + context_id=context_id, + expected_state=task_record.state, + ): raise RuntimeError("Execution task did not reach a terminal state") + await self._task_store.discard_context_runtime(context_id, persist_context=False) return task_record.state + if reason == "natural_completion" and task_record.state == TASK_STATE_CANCELED: + if not await self._task_store.commit_inactive_execution_task( + task_id=task_id, + context_id=context_id, + expected_state=TASK_STATE_CANCELED, + ): + raise RuntimeError("Canceled execution changed during natural completion") + await self._task_store.discard_context_runtime(context_id, persist_context=False) + return TASK_STATE_CANCELED + if ( + reason in self._RECOVERABLE_RELEASE_REASONS + and task_record.state == TASK_STATE_INPUT_REQUIRED + and not had_pending_permission + ): + if control is not None and control.execution_mode == "normal": + committed = await self._task_store.commit_inactive_execution_task( + task_id=task_id, + context_id=context_id, + expected_state=TASK_STATE_INPUT_REQUIRED, + ) + if committed: + await self._task_store.discard_context_runtime(context_id, persist_context=False) + committed_task = await self._task_store.get_task_record(task_id) + if committed_task.state == TASK_STATE_INPUT_REQUIRED: + return TASK_STATE_INPUT_REQUIRED + waiting_task_id = await run_sync_fenced( + sandbox_release_recoverable_task_id_from_sidecar, + cwd=context_record.cwd, + session_id=context_record.session_id, + context_id=context_id, + ) + if waiting_task_id == task_id: + committed = await self._task_store.commit_inactive_execution_task( + task_id=task_id, + context_id=context_id, + expected_state=TASK_STATE_INPUT_REQUIRED, + ) + if committed: + await self._task_store.discard_context_runtime(context_id, persist_context=False) + committed_task = await self._task_store.get_task_record(task_id) + committed_waiting_task_id = await run_sync_fenced( + sandbox_release_recoverable_task_id_from_sidecar, + cwd=context_record.cwd, + session_id=context_record.session_id, + context_id=context_id, + ) + if committed_task.state == TASK_STATE_INPUT_REQUIRED and committed_waiting_task_id == task_id: + return TASK_STATE_INPUT_REQUIRED + if reason == "natural_completion": + raise RuntimeError("Natural input completion lacks a recoverable handoff proof") cancel_result = await run_sync_fenced( cancel_waiting_input_task_from_sidecar, cwd=context_record.cwd, @@ -3506,9 +3695,7 @@ def audit_claim(value: str) -> bool: event_queue, task_id=response.task_id, context_id=response.context_id, - state=( - TaskState.TASK_STATE_INPUT_REQUIRED if normal_turn_finished else TaskState.TASK_STATE_CANCELED - ), + state=(TaskState.TASK_STATE_INPUT_REQUIRED if normal_turn_finished else TaskState.TASK_STATE_CANCELED), session_id=context_record.session_id, ) await self._notify_terminal_task( @@ -3837,9 +4024,7 @@ def _read(name: str) -> str | None: return None if re.fullmatch(r"[a-z0-9][a-z0-9-]{0,62}", region_id) is None: language = self._resolve_preferred_language(metadata) or "en" - raise InvalidParamsError( - translate_message("Unsupported Alibaba Cloud region ID.", language=language) - ) + raise InvalidParamsError(translate_message("Unsupported Alibaba Cloud region ID.", language=language)) configured = AliyunCredentials.load() if configured is None: return None diff --git a/src/iac_code/a2a/input_required.py b/src/iac_code/a2a/input_required.py index 3a9e23d6..2f217071 100644 --- a/src/iac_code/a2a/input_required.py +++ b/src/iac_code/a2a/input_required.py @@ -30,6 +30,7 @@ emit_permission_boundary_audit, sanitize_prompt_text, ) +from iac_code.services.providers.aliyun import AliyunCredential, current_aliyun_credential_override from iac_code.services.providers.aliyun_identity import AliyunCallerIdentityUnavailableError from iac_code.types.stream_events import PermissionRequestEvent @@ -116,6 +117,7 @@ class PendingPermission: context_id: str input_id: str request: PermissionRequestEvent + aliyun_credential: AliyunCredential | None = field(default=None, repr=False) language: str = field(default_factory=lambda: get_a2a_preferred_language() or "en") resolution_owner: PermissionResolutionOwner | None = None scope: str = "pipeline" @@ -1158,6 +1160,7 @@ async def register( context_id=context_id, input_id=input_id, request=request, + aliyun_credential=current_aliyun_credential_override(), resolution_owner=resolution_owner, scope=scope, coordinates=dict(coordinates) if coordinates is not None else None, @@ -1171,16 +1174,22 @@ async def answer( self, response: PermissionResponse, *, + aliyun_credential: AliyunCredential | None = None, before_delivery: Callable[[], Awaitable[None]] | None = None, ) -> bool: pending = await self._lookup(response) + + async def prepare_delivery() -> None: + if aliyun_credential is not None and pending.aliyun_credential is not None: + pending.aliyun_credential.refresh_from(aliyun_credential) + if before_delivery is not None: + await before_delivery() + if pending.resolution_owner is not None: - if before_delivery is None: - return await pending.resolution_owner.resolve_permission(pending, response) return await pending.resolution_owner.resolve_permission( pending, response, - before_delivery=before_delivery, + before_delivery=prepare_delivery, ) coordinator = self._permission_wait_coordinator @@ -1205,7 +1214,7 @@ def audit_new_claim(value: str) -> bool: source="user", on_new_claim=audit_new_claim, before_delivery=lambda record: self._backup_claim_before_delivery(pending, record), - before_release=before_delivery, + before_release=prepare_delivery, ) except (LookupError, ValueError) as exc: raise InvalidParamsError(f"permission_resume_invalid: {exc}") from exc @@ -1231,8 +1240,7 @@ def audit_new_claim(value: str) -> bool: ) if approved and not audit_ok: approved = False - if before_delivery is not None: - await before_delivery() + await prepare_delivery() future.set_result(approved) return approved @@ -1393,6 +1401,18 @@ async def reopen_task(self, token: PermissionTaskClosingToken | None) -> None: self._closing_tasks.pop(token.task_id, None) self._condition.notify_all() + async def reversible_closing_tokens(self, task_id: str) -> tuple[PermissionTaskClosingToken, ...]: + """Snapshot reversible closes that belong to the current task generation.""" + + async with self._condition: + closing = self._closing_tasks.get(task_id) + if closing is None or closing.permanent: + return () + return tuple( + PermissionTaskClosingToken(task_id=task_id, token_id=token_id) + for token_id in closing.reversible_tokens + ) + async def fail(self, pending: PendingPermission) -> None: if pending.resolution_owner is not None: await pending.resolution_owner.fail_permission(pending) diff --git a/src/iac_code/a2a/pipeline_executor.py b/src/iac_code/a2a/pipeline_executor.py index 01d3e73a..54e3a74d 100644 --- a/src/iac_code/a2a/pipeline_executor.py +++ b/src/iac_code/a2a/pipeline_executor.py @@ -47,6 +47,10 @@ pending_backup_publication_envelope, ) from iac_code.a2a.pipeline_transport_delivery import PipelineTransportDeliveryClosedError +from iac_code.a2a.request_scoped_active_task import ( + DirectPipelineRouteGate, + PipelineLifecycleEventQueueCarrier, +) from iac_code.a2a.resource_selector import ( PendingResourceSelection, ResourceSelectionCheckpointStore, @@ -97,6 +101,7 @@ TextDeltaEvent, ) from iac_code.utils.path_locks import PathLockRegistry +from iac_code.utils.public_errors import sanitize_strict_text logger = logging.getLogger(__name__) _CONTEXT_LOCK_ACQUIRE_TIMEOUT_SECONDS = 1 @@ -106,6 +111,9 @@ _TERMINAL_A2A_STATUSES = {"completed", "failed", "canceled"} _WAITING_A2A_STATUSES = {"waiting_input", "input_required"} _RUNNING_A2A_STATUSES = {"working"} +_SANDBOX_RELEASE_RECOVERABLE_INPUT_KINDS = frozenset( + {"ask_user_question", "candidate_selection", "deployment_confirmation", "pipeline_pause_confirmation"} +) _PENDING_BACKUP_VISIBILITY = "pending_backup" _COMMITTED_BACKUP_VISIBILITY = "committed" _WAITING_INPUT_CANCEL_LOCKS = PathLockRegistry() @@ -215,6 +223,12 @@ class A2APipelineRuntime: restart_requested: asyncio.Event = field(default_factory=asyncio.Event) interrupt_settled: asyncio.Event = field(default_factory=_new_set_asyncio_event) + def bind_publisher_event_queue(self, event_queue: Any) -> None: + """Route resumed Pipeline publications through the current SDK lifecycle.""" + + if self.publisher is not None: + self.publisher.event_queue = event_queue + @dataclass(frozen=True) class _StreamConsumeResult: @@ -634,6 +648,7 @@ async def execute( pipeline_input: PipelineUserInput | str | None = None, prompt: str | None = None, active_followup_only: bool = False, + direct_route_gate: DirectPipelineRouteGate | None = None, permission_checkpoint: dict[str, Any] | None = None, resource_selection_checkpoint: dict[str, Any] | None = None, ) -> bool | None: @@ -680,6 +695,20 @@ def runtime_factory(session_id: str) -> Any: except asyncio.CancelledError: # Context creation drains and closes an unfinished runtime before # cancellation reaches here, even if no Pipeline exists yet. + control = current_execution_control() + current_task = asyncio.current_task() + logger.warning( + "A2A Pipeline runtime setup canceled task_id=%s context_id=%s " + "asyncio_task=%s cancelling=%s control_execution_id=%s control_phase=%s " + "control_reason=%s", + sanitize_strict_text(task_id), + sanitize_strict_text(context_id), + sanitize_strict_text(current_task.get_name() if current_task is not None else "none"), + getattr(current_task, "cancelling", lambda: 0)() if current_task is not None else 0, + sanitize_strict_text(control.execution_id if control is not None else "none"), + sanitize_strict_text(control.phase if control is not None else "none"), + sanitize_strict_text(control.termination_reason if control is not None else "none"), + ) task.active_task = None task.state = TASK_STATE_CANCELED self._task_store.mirror_task(task) @@ -743,6 +772,8 @@ def runtime_factory(session_id: str) -> Any: cwd=cwd, pipeline_input=pipeline_input, preserve_task_record=preserve_active_task, + direct_route_gate=direct_route_gate, + bind_publisher_event_queue=PipelineLifecycleEventQueueCarrier.read(context), ) if routed: return True @@ -771,10 +802,33 @@ def runtime_factory(session_id: str) -> Any: try: await asyncio.wait_for(lock.acquire(), timeout=_CONTEXT_LOCK_ACQUIRE_TIMEOUT_SECONDS) except TimeoutError: + if direct_route_gate is not None: + direct_route_gate.require_recovery() await self._fail_already_active(event_queue, task=task, task_id=task_id, context_id=context_id) return try: + if direct_route_gate is not None: + # The SDK lifecycle can outlive the Pipeline's internal owner at + # an input boundary. Fence the new continuation before it + # publishes into that still-active lifecycle. + await direct_route_gate.activate(event_queue) + elif ( + task.state == TASK_STATE_INPUT_REQUIRED + and PipelineLifecycleEventQueueCarrier.read(context) + and not PipelineLifecycleEventQueueCarrier.is_bound(context) + ): + # A recovered SDK lifecycle reuses the existing Task projection, + # so the SDK does not emit another initial Task. Publish a + # request-scoped frame before restoring the finite sidecar stream; + # otherwise a cold resume that has no public Pipeline event can + # complete as a successful but empty response. + await self._publish_status( + event_queue, + task_id=task_id, + context_id=context_id, + state=TaskState.TASK_STATE_WORKING, + ) owner_task = asyncio.current_task() task_persistence_started = False @@ -1127,8 +1181,20 @@ async def resume_detached_resource_selection( except asyncio.CancelledError: try: task.state = TASK_STATE_CANCELED + control = current_execution_control() execution_termination_reason = current_execution_termination_reason() cancel_source = "execution_control" if execution_termination_reason is not None else "executor" + logger.warning( + "A2A pipeline execution canceled task_id=%s context_id=%s " + "cancel_source=%s execution_id=%s control_phase=%s " + "termination_reason=%s", + task_id, + context_id, + cancel_source, + getattr(control, "execution_id", None), + getattr(control, "phase", None), + sanitize_strict_text(execution_termination_reason or ""), + ) cancel_reason = execution_termination_reason or _("Task canceled.") cancel_data = {"source": cancel_source, "reason": cancel_reason} cancel_handoff_data = {"canceled": True, "reason": _("Task canceled.")} @@ -1379,14 +1445,46 @@ async def _route_active_pipeline_interrupt( cwd: str, pipeline_input: PipelineUserInput, preserve_task_record: bool, + direct_route_gate: DirectPipelineRouteGate | None = None, + bind_publisher_event_queue: bool = False, ) -> bool: runtime = ctx.runtime if getattr(runtime, "pipeline", None) is None: return False - if not await _register_active_interrupt(runtime): + if not await _register_active_interrupt( + runtime, + event_queue=event_queue, + direct_route_gate=direct_route_gate, + bind_publisher_event_queue=bind_publisher_event_queue, + ): logger.info("Ignoring A2A pipeline interrupt after terminal publication started") + # Direct routing has a recovery gate that makes the dispatcher + # recreate the request lifecycle. A base SDK lifecycle needs an + # explicit non-terminal frame so the caller does not mistake an + # empty stream (or a terminal retry) for this request's result. + if direct_route_gate is None and bind_publisher_event_queue: + await self._publish_status( + event_queue, + task_id=task_id, + context_id=context_id, + state=TaskState.TASK_STATE_INPUT_REQUIRED, + text=_retry_text(), + ) return True + # The SDK lifecycle that recovered an input-required task may not + # replay its existing Task projection. Publish a request-scoped + # non-terminal frame before routing so the upstream caller can bind + # execution control even when the accepted interrupt itself completes + # without producing another public event. + if direct_route_gate is None and bind_publisher_event_queue: + await self._publish_status( + event_queue, + task_id=task_id, + context_id=context_id, + state=TaskState.TASK_STATE_WORKING, + ) + interrupt_registered = True async def settle_interrupt() -> None: @@ -4598,6 +4696,37 @@ def waiting_input_task_id_from_sidecar(*, cwd: str, session_id: str, context_id: ) +def sandbox_release_recoverable_task_id_from_sidecar(*, cwd: str, session_id: str, context_id: str) -> str | None: + task_id = waiting_input_task_id_from_sidecar(cwd=cwd, session_id=session_id, context_id=context_id) + if task_id is None: + return None + pipeline_dir = existing_a2a_pipeline_dir_for_session(cwd=cwd, session_id=session_id) + snapshot_store = A2APipelineSnapshotStore(pipeline_dir) + journal = A2APipelineJournal(pipeline_dir) + pending_input = _pending_input_from_snapshot( + _authoritative_snapshot_for_task( + snapshot_store=snapshot_store, + journal=journal, + task_id=task_id, + context_id=context_id, + ), + task_id=task_id, + context_id=context_id, + ) + if pending_input is None: + return None + kind = pending_input.get("kind") + if kind not in _SANDBOX_RELEASE_RECOVERABLE_INPUT_KINDS: + return None + if kind == "ask_user_question" and not _string_value( + pending_input.get("toolUseId") or pending_input.get("tool_use_id") + ): + return None + if kind == "pipeline_pause_confirmation" and pending_input.get("paused") is not True: + return None + return task_id + + def cancel_waiting_input_task_from_sidecar( *, cwd: str, @@ -5368,6 +5497,13 @@ async def _drive_stream_events( try: event = await anext(stream) except asyncio.CancelledError: + control = current_execution_control() + logger.warning( + "A2A pipeline source canceled execution_id=%s control_phase=%s termination_reason=%s", + getattr(control, "execution_id", None), + getattr(control, "phase", None), + sanitize_strict_text(current_execution_termination_reason() or ""), + ) completion.cancel() raise except BaseException as exc: @@ -5434,12 +5570,32 @@ async def settle() -> None: raise cancellation -async def _register_active_interrupt(runtime: Any) -> bool: +async def _register_active_interrupt( + runtime: Any, + *, + event_queue: Any | None = None, + direct_route_gate: DirectPipelineRouteGate | None = None, + bind_publisher_event_queue: bool = False, +) -> bool: async with _outbound_lock(runtime): if bool(getattr(runtime, "terminal_publication_started", False)): + if direct_route_gate is not None: + direct_route_gate.require_recovery() return False runtime.active_interrupt_count = _active_interrupt_count(runtime) + 1 _interrupt_settled_event(runtime).clear() + try: + if direct_route_gate is not None: + if event_queue is None: + raise RuntimeError("Direct Pipeline route gate requires an event queue") + await direct_route_gate.activate(event_queue) + if event_queue is not None and (direct_route_gate is not None or bind_publisher_event_queue): + runtime.bind_publisher_event_queue(event_queue) + except BaseException: + runtime.active_interrupt_count = max(0, _active_interrupt_count(runtime) - 1) + if runtime.active_interrupt_count == 0: + _interrupt_settled_event(runtime).set() + raise return True diff --git a/src/iac_code/a2a/request_scoped_active_task.py b/src/iac_code/a2a/request_scoped_active_task.py new file mode 100644 index 00000000..a0b1f847 --- /dev/null +++ b/src/iac_code/a2a/request_scoped_active_task.py @@ -0,0 +1,572 @@ +"""Request-scoped event subscriptions for the A2A SDK ActiveTask runtime.""" + +from __future__ import annotations + +import asyncio +import logging +from collections.abc import Awaitable, Callable +from enum import Enum +from typing import TYPE_CHECKING, Any, cast + +from a2a.server.agent_execution.active_task import ( + TERMINAL_TASK_STATES, + ActiveTask, + _RequestCompleted, + _RequestStarted, +) +from a2a.server.agent_execution.active_task_registry import ActiveTaskRegistry +from a2a.server.events.event_queue_v2 import QueueShutDown +from a2a.server.tasks.task_manager import TaskManager +from a2a.types import Task, TaskState, TaskStatusUpdateEvent +from a2a.utils.errors import InvalidParamsError + +from iac_code.a2a.backup import await_fenced +from iac_code.a2a.execution_control import RecoverableInputAdmissionCarrier +from iac_code.utils.public_errors import sanitize_strict_text + +if TYPE_CHECKING: + from collections.abc import AsyncGenerator + + from a2a.server.agent_execution import RequestContext + from a2a.server.context import ServerCallContext + from a2a.server.events import Event + + +logger = logging.getLogger(__name__) + + +class RequestScopedActiveTask(ActiveTask): + """Hide events from earlier requests until this request actually starts.""" + + def __init__( + self, + *args: Any, + recovery_admission: str | None = None, + **kwargs: Any, + ) -> None: + super().__init__(*args, **kwargs) + self._direct_message_lock = asyncio.Lock() + self._recovery_admission = recovery_admission + self._recovery_replacement_pending = False + + @property + def direct_message_lock(self) -> asyncio.Lock: + """Serialize direct messages injected into the running executor.""" + + return self._direct_message_lock + + async def retire_for_recovery(self) -> None: + """Stop an idle SDK lifecycle before reopening its durable task.""" + + async with self._lock: + self._is_finished.set() + self._request_queue.shutdown(immediate=True) + lifecycle_tasks = self._retirement_tasks() + + await self._finish_retirement(lifecycle_tasks) + + def _retirement_tasks(self) -> tuple[asyncio.Task[Any], ...]: + current = asyncio.current_task() + return tuple( + task + for task in (self._producer_task, self._consumer_task) + if task is not None and task is not current + ) + + async def _finish_retirement(self, lifecycle_tasks: tuple[asyncio.Task[Any], ...]) -> None: + await await_fenced(self._finish_retirement_owned(lifecycle_tasks)) + + async def _finish_retirement_owned(self, lifecycle_tasks: tuple[asyncio.Task[Any], ...]) -> None: + if lifecycle_tasks: + logger.warning( + "A2A SDK lifecycle retirement canceling tasks task_id=%s task_names=%s", + sanitize_strict_text(self._task_id), + ",".join(sanitize_strict_text(task.get_name()) for task in lifecycle_tasks), + ) + for task in lifecycle_tasks: + if not task.done(): + task.cancel() + if lifecycle_tasks: + await asyncio.gather(*lifecycle_tasks, return_exceptions=True) + await self._event_queue_agent.close(immediate=True) + await self._event_queue_subscribers.close(immediate=True) + + async def enqueue_request(self, request_context: RequestContext): + admission = RecoverableInputAdmissionCarrier.read(request_context) + async with self._lock: + if self._recovery_replacement_pending: + raise InvalidParamsError(f"Task {self._task_id} recovery replacement is pending.") + if self._recovery_admission is not None and admission != self._recovery_admission: + raise InvalidParamsError(f"Task {self._task_id} recovery continuation is reserved.") + request_id = await super().enqueue_request(request_context) + if self._recovery_admission is not None: + self._recovery_admission = None + RecoverableInputAdmissionCarrier.acknowledge_enqueued(request_context, self._producer_task) + return request_id + + def has_unfinished_requests(self) -> bool: + """Report SDK requests accepted by this lifecycle but not yet completed.""" + + return self._queue_has_unfinished_tasks(self._request_queue) + + def has_unsettled_requests(self) -> bool: + """Report accepted requests whose SDK projection or fan-out is not settled.""" + + return bool( + self.has_unfinished_requests() + or self._request_lock.locked() + or self._queue_has_unfinished_tasks(getattr(self._event_queue_agent, "_incoming_queue", None)) + or self._queue_has_unfinished_tasks(getattr(self._event_queue_agent, "queue", None)) + ) + + @staticmethod + def _queue_has_unfinished_tasks(queue: Any) -> bool: + if queue is None: + return False + unfinished = getattr(queue, "unfinished_tasks", None) + if isinstance(unfinished, int): + return unfinished > 0 + return bool(getattr(queue, "_unfinished_tasks", 0)) + + async def wait_for_accepted_requests(self) -> None: + """Wait through executor completion and the SDK consumer's durable projection.""" + + await self._request_queue.join() + projection = asyncio.create_task( + self._wait_for_accepted_request_projection(), + name=f"accepted-request-projection:{self._task_id}", + ) + consumer = self._consumer_task + if consumer is None: + projection.cancel() + await asyncio.gather(projection, return_exceptions=True) + raise InvalidParamsError(f"Task {self._task_id} lifecycle ended before the accepted request settled.") + try: + done, _ = await asyncio.wait((projection, consumer), return_when=asyncio.FIRST_COMPLETED) + if consumer in done and self._request_lock.locked(): + raise InvalidParamsError(f"Task {self._task_id} lifecycle ended before the accepted request settled.") + if projection in done: + await projection + if self._request_lock.locked(): + raise InvalidParamsError( + f"Task {self._task_id} lifecycle ended before the accepted request settled." + ) + return + await asyncio.sleep(0) + if projection.done(): + await projection + if self._request_lock.locked(): + raise InvalidParamsError( + f"Task {self._task_id} lifecycle ended before the accepted request settled." + ) + return + raise InvalidParamsError(f"Task {self._task_id} lifecycle ended before the accepted request settled.") + finally: + if not projection.done(): + projection.cancel() + await asyncio.gather(projection, return_exceptions=True) + + async def _wait_for_accepted_request_projection(self) -> None: + join_incoming = getattr(self._event_queue_agent, "test_only_join_incoming_queue", None) + if callable(join_incoming): + await join_incoming() + agent_queue = getattr(self._event_queue_agent, "queue", None) + join_agent_queue = getattr(agent_queue, "join", None) + if callable(join_agent_queue): + await join_agent_queue() + + async def subscribe( + self, + *, + request: RequestContext | None = None, + include_initial_task: bool = False, + replace_status_update_with_task: bool = False, + ) -> AsyncGenerator[Event, None]: + async with self._lock: + if self._is_finished.is_set(): + raise InvalidParamsError(f"Task {self._task_id} is already completed.") + self._reference_count += 1 + + tapped_queue = None + request_id = None + request_started = request is None + pre_start_error: BaseException | None = None + + try: + tapped_queue = await self._event_queue_subscribers.tap() + if request is not None: + request_id = await self.enqueue_request(request) + + if include_initial_task: + yield await self.get_task() + + while True: + try: + event, updated_task = cast("Any", await tapped_queue.dequeue_event()) + except QueueShutDown: + if request is not None and not request_started: + if pre_start_error is not None: + raise pre_start_error + raise InvalidParamsError(f"Task {self._task_id} ended before the request started.") + break + except asyncio.CancelledError: + break + + try: + if isinstance(event, DirectMessageRequestStarted): + continue + if isinstance(event, _RequestStarted): + if request_id is not None and event.request_id == request_id: + request_started = True + continue + if not request_started: + if isinstance(event, BaseException): + pre_start_error = event + continue + if isinstance(event, BaseException): + raise event + if isinstance(event, _RequestCompleted): + if request_id is not None and event.request_id == request_id: + return + continue + if self.is_stale_terminal_projection(event, updated_task): + continue + if replace_status_update_with_task and isinstance(event, TaskStatusUpdateEvent): + event = updated_task + yield cast("Event", event) + finally: + tapped_queue.task_done() + finally: + if tapped_queue is not None: + await tapped_queue.close(immediate=True) + async with self._lock: + self._reference_count -= 1 + await self._maybe_cleanup() + + @staticmethod + def is_stale_terminal_projection(event: Any, updated_task: Any) -> bool: + """Drop a terminal event that the canonical task store rejected as stale.""" + + return bool( + isinstance(event, (Task, TaskStatusUpdateEvent)) + and event.status.state in TERMINAL_TASK_STATES + and isinstance(updated_task, Task) + and updated_task.status.state not in TERMINAL_TASK_STATES + ) + + +class DirectMessageRequestStarted: + """Internal fan-out fence separating a direct request from older events.""" + + +class DirectPipelineRouteOutcome(Enum): + PENDING = "pending" + ACTIVE = "active" + RECOVERY_REQUIRED = "recovery_required" + + +class DirectPipelineRouteGate: + """Bind a direct stream boundary to the running pipeline accepting its input.""" + + def __init__(self) -> None: + self.marker = DirectMessageRequestStarted() + self.outcome = DirectPipelineRouteOutcome.PENDING + + async def activate(self, event_queue: Any) -> None: + if self.outcome is not DirectPipelineRouteOutcome.PENDING: + return + await event_queue.enqueue_event(self.marker) + self.outcome = DirectPipelineRouteOutcome.ACTIVE + + def require_recovery(self) -> None: + if self.outcome is DirectPipelineRouteOutcome.PENDING: + self.outcome = DirectPipelineRouteOutcome.RECOVERY_REQUIRED + + +class DirectPipelineRouteGateCarrier: + """Attach the transport's route gate to one direct RequestContext.""" + + _ATTRIBUTE = "_iac_code_direct_pipeline_route_gate" + + @classmethod + def attach(cls, request_context: Any, gate: DirectPipelineRouteGate) -> None: + setattr(request_context, cls._ATTRIBUTE, gate) + + @classmethod + def read(cls, request_context: Any) -> DirectPipelineRouteGate | None: + gate = getattr(request_context, cls._ATTRIBUTE, None) + return gate if isinstance(gate, DirectPipelineRouteGate) else None + + +class PipelineLifecycleEventQueueCarrier: + """Mark a RequestContext whose SDK lifecycle owns its event queue.""" + + _ATTRIBUTE = "_iac_code_pipeline_lifecycle_event_queue" + _BOUND_ATTRIBUTE = "_iac_code_pipeline_lifecycle_event_queue_bound" + + @classmethod + def attach(cls, request_context: Any) -> None: + setattr(request_context, cls._ATTRIBUTE, True) + + @classmethod + def read(cls, request_context: Any) -> bool: + return getattr(request_context, cls._ATTRIBUTE, False) is True + + @classmethod + def mark_bound(cls, request_context: Any) -> None: + setattr(request_context, cls._BOUND_ATTRIBUTE, True) + + @classmethod + def is_bound(cls, request_context: Any) -> bool: + return getattr(request_context, cls._BOUND_ATTRIBUTE, False) is True + + +class DirectPipelineRecoveryRequiredError(RuntimeError): + """The old owner won terminal publication, so this request must recover.""" + + +class RequestScopedActiveTaskRegistry(ActiveTaskRegistry): + """Create request-scoped ActiveTask instances while retaining SDK lifecycle rules.""" + + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self._recoveries_in_progress: set[str] = set() + + async def reconcile_and_replace_for_recovery( + self, + task_id: str, + *, + call_context: ServerCallContext, + context_id: str, + acquire_admission: Callable[[], Awaitable[str | None]], + release_admission: Callable[[str], Awaitable[None]] | None = None, + ) -> str | None: + """Claim durable recovery and replace the old lifecycle in one local critical section.""" + + scoped = await self._reserve_recovery_replacement(task_id) + admission = None + replacement = None + retirement_started = False + published = False + try: + if scoped is not None: + await self._drain_accepted_recovery_predecessors(scoped, task_id, call_context) + admission = await acquire_admission() + if admission is None: + return None + + lifecycle_tasks = await self._begin_recovery_retirement(task_id, scoped) + retirement_started = scoped is not None + if scoped is not None: + await scoped._finish_retirement(lifecycle_tasks) + replacement = self._create_active_task( + task_id=task_id, + call_context=call_context, + context_id=context_id, + recovery_admission=admission, + ) + await replacement.start(call_context=call_context, create_task_if_missing=True) + await self._publish_recovery_replacement(task_id, scoped, replacement) + published = True + return admission + except BaseException: + await await_fenced( + self._abort_recovery_replacement( + replacement=replacement, + admission=admission, + release_admission=release_admission, + ) + ) + raise + finally: + await await_fenced( + self._clear_recovery_replacement( + task_id, + scoped, + remove_finished=retirement_started and not published, + ) + ) + + async def _reserve_recovery_replacement(self, task_id: str) -> RequestScopedActiveTask | None: + async with self._lock: + if task_id in self._recoveries_in_progress: + raise InvalidParamsError(f"Task {task_id} recovery replacement is pending.") + self._recoveries_in_progress.add(task_id) + existing = self._active_tasks.get(task_id) + scoped = cast("RequestScopedActiveTask | None", existing) + if scoped is not None and scoped._recovery_replacement_pending: + self._recoveries_in_progress.discard(task_id) + raise InvalidParamsError(f"Task {task_id} recovery replacement is pending.") + if scoped is not None: + scoped._recovery_replacement_pending = True + + try: + if scoped is not None: + # An enqueue that crossed the fence already owns predecessor + # status. Wait for its short critical section without holding + # the registry-wide mapping lock. + async with scoped._lock: + pass + return scoped + except BaseException: + await await_fenced(self._clear_recovery_replacement(task_id, scoped, remove_finished=False)) + raise + + async def _begin_recovery_retirement( + self, + task_id: str, + scoped: RequestScopedActiveTask | None, + ) -> tuple[asyncio.Task[Any], ...]: + async with self._lock: + if self._active_tasks.get(task_id) is not scoped: + raise InvalidParamsError(f"Task {task_id} lifecycle changed during recovery.") + if scoped is None: + return () + scoped._is_finished.set() + scoped._request_queue.shutdown(immediate=True) + return scoped._retirement_tasks() + + async def _publish_recovery_replacement( + self, + task_id: str, + scoped: RequestScopedActiveTask | None, + replacement: RequestScopedActiveTask, + ) -> None: + async with self._lock: + if self._active_tasks.get(task_id) is not scoped: + raise InvalidParamsError(f"Task {task_id} lifecycle changed during recovery.") + self._active_tasks[task_id] = replacement + + @staticmethod + async def _abort_recovery_replacement( + *, + replacement: RequestScopedActiveTask | None, + admission: str | None, + release_admission: Callable[[str], Awaitable[None]] | None, + ) -> None: + if replacement is not None: + await replacement.retire_for_recovery() + if admission is not None and release_admission is not None: + await release_admission(admission) + + async def _clear_recovery_replacement( + self, + task_id: str, + scoped: RequestScopedActiveTask | None, + *, + remove_finished: bool, + ) -> None: + async with self._lock: + if scoped is not None: + scoped._recovery_replacement_pending = False + if ( + self._active_tasks.get(task_id) is scoped + and (remove_finished or scoped._is_finished.is_set()) + ): + self._active_tasks.pop(task_id, None) + self._recoveries_in_progress.discard(task_id) + + async def _drain_accepted_recovery_predecessors( + self, + scoped: RequestScopedActiveTask, + task_id: str, + call_context: ServerCallContext, + ) -> None: + """Preserve accepted old-lifecycle requests before a terminal recovery decision.""" + + if not scoped.has_unsettled_requests(): + return + task = await self._task_store.get(task_id, call_context) + if task is None or ( + task.status.state not in TERMINAL_TASK_STATES + and task.status.state != TaskState.TASK_STATE_INPUT_REQUIRED + ): + return + await scoped.wait_for_accepted_requests() + + async def cancel_recovery_reservation(self, task_id: str, admission: str) -> None: + """Remove an unconsumed reservation and its unpublished replacement lifecycle.""" + + active_task = None + async with self._lock: + candidate = self._active_tasks.get(task_id) + if isinstance(candidate, RequestScopedActiveTask) and candidate._recovery_admission == admission: + active_task = self._active_tasks.pop(task_id) + if active_task is not None: + await cast("RequestScopedActiveTask", active_task).retire_for_recovery() + + async def retire_for_recovery(self, task_id: str) -> None: + """Remove and stop a lifecycle during tests or shutdown.""" + + async with self._lock: + active_task = self._active_tasks.pop(task_id, None) + if active_task is not None: + await cast("RequestScopedActiveTask", active_task).retire_for_recovery() + + async def get_or_create( + self, + task_id: str, + call_context: ServerCallContext, + context_id: str | None = None, + create_task_if_missing: bool = False, + ) -> RequestScopedActiveTask: + async with self._lock: + if task_id in self._recoveries_in_progress: + raise InvalidParamsError(f"Task {task_id} recovery replacement is pending.") + existing = self._active_tasks.get(task_id) + if existing is not None and not existing._is_finished.is_set(): + return cast("RequestScopedActiveTask", existing) + if existing is not None: + self._active_tasks.pop(task_id, None) + + active_task = self._create_active_task( + task_id=task_id, + call_context=call_context, + context_id=context_id, + ) + self._active_tasks[task_id] = active_task + + await active_task.start( + call_context=call_context, + create_task_if_missing=create_task_if_missing, + ) + return active_task + + def _create_active_task( + self, + *, + task_id: str, + call_context: ServerCallContext, + context_id: str | None, + recovery_admission: str | None = None, + ) -> RequestScopedActiveTask: + task_manager = TaskManager( + task_id=task_id, + context_id=context_id, + task_store=self._task_store, + initial_message=None, + context=call_context, + ) + return RequestScopedActiveTask( + agent_executor=self._agent_executor, + task_id=task_id, + task_manager=task_manager, + push_sender=self._push_sender, + on_cleanup=self._on_active_task_cleanup, + recovery_admission=recovery_admission, + ) + + def _on_active_task_cleanup(self, active_task: ActiveTask) -> None: + cleanup = asyncio.create_task( + self._remove_task_if_same(active_task), + name=f"remove-finished-active-task:{active_task.task_id}", + ) + self._cleanup_tasks.add(cleanup) + cleanup.add_done_callback(self._cleanup_tasks.discard) + + async def _remove_task_if_same(self, active_task: ActiveTask) -> None: + async with self._lock: + if active_task.task_id in self._recoveries_in_progress: + return + if self._active_tasks.get(active_task.task_id) is active_task: + self._active_tasks.pop(active_task.task_id, None) diff --git a/src/iac_code/a2a/task_store.py b/src/iac_code/a2a/task_store.py index ccedee52..b433437d 100644 --- a/src/iac_code/a2a/task_store.py +++ b/src/iac_code/a2a/task_store.py @@ -22,7 +22,11 @@ from iac_code.a2a.backup import await_fenced, run_sync_fenced from iac_code.a2a.events import with_iac_code_session_metadata -from iac_code.a2a.execution_control import current_execution_control, current_execution_termination_reason +from iac_code.a2a.execution_control import ( + ExecutionController, + current_execution_control, + current_execution_termination_reason, +) from iac_code.a2a.metadata_redaction import strip_llm_headers_from_metadata from iac_code.a2a.metrics import A2AMetrics, NoOpA2AMetrics from iac_code.a2a.persistence import A2AContextSnapshot, A2APersistenceStore, A2ATaskSnapshot @@ -43,9 +47,11 @@ from iac_code.services.session_storage import SessionStorage from iac_code.services.telemetry.attributes import normalize_telemetry_channel from iac_code.utils.file_security import atomic_write_text +from iac_code.utils.public_errors import sanitize_strict_text logger = logging.getLogger(__name__) A2ATaskSnapshotList: TypeAlias = list[A2ATaskSnapshot] +_RUNTIME_CLOSE_TIMEOUT_SECONDS = 2.0 class A2ATaskStore(TaskStore): @@ -89,6 +95,10 @@ def __init__( self._permission_wait_active_probe: Callable[[], bool] | None = None self._execution_control_snapshot_provider: Callable[[str], dict[str, Any] | None] | None = None self._execution_control_active_probe: Callable[[], bool] | None = None + self._recoverable_input_admission_reserver: ( + Callable[..., Awaitable[str | None]] | None + ) = None + self._recoverable_input_admission_releaser: Callable[[str], Awaitable[None]] | None = None def set_permission_wait_active_probe(self, probe: Callable[[], bool] | None) -> None: self._permission_wait_active_probe = probe @@ -97,9 +107,13 @@ def set_execution_control_provider( self, snapshot_provider: Callable[[str], dict[str, Any] | None] | None, active_probe: Callable[[], bool] | None, + recoverable_input_admission_reserver: Callable[..., Awaitable[str | None]] | None = None, + recoverable_input_admission_releaser: Callable[[str], Awaitable[None]] | None = None, ) -> None: self._execution_control_snapshot_provider = snapshot_provider self._execution_control_active_probe = active_probe + self._recoverable_input_admission_reserver = recoverable_input_admission_reserver + self._recoverable_input_admission_releaser = recoverable_input_admission_releaser def touch_context(self, context_id: str) -> None: """Restart the idle interval after a connection hold, without storage I/O.""" @@ -145,80 +159,316 @@ async def save(self, task: Task, context: ServerCallContext | None = None) -> No task_id = validate_protocol_id(task.id) _strip_task_llm_headers(task) async with self._mutation_lock: - record = self._tasks.get(task_id) - next_state = _task_state_from_sdk_task(task) - incoming_updated_at = _task_updated_at_from_sdk_task(task) - preserve_terminal = bool( - record is not None - and record.state - in {TASK_STATE_CANCELED, TASK_STATE_COMPLETED, TASK_STATE_FAILED, TASK_STATE_INPUT_REQUIRED} - and self._execution_is_terminating(task.context_id) - ) - if preserve_terminal: - # Queued SDK working events must not undo the executor's final - # state while its immutable termination snapshot is committed. - assert record is not None - task = _copy_task(task) - task.status.state = TaskState.Value("TASK_STATE_" + record.state.upper().replace("-", "_")) - active_finalization = bool( - record is not None - and record.state - in {TASK_STATE_CANCELED, TASK_STATE_COMPLETED, TASK_STATE_FAILED, TASK_STATE_INPUT_REQUIRED} - and next_state in {TASK_STATE_SUBMITTED, TASK_STATE_WORKING} - and record.active_task is not None - and not record.active_task.done() + self._save_locked(task, owner=owner, task_id=task_id) + + async def reconcile_recoverable_input_required_task( + self, + task: Task, + context_record: A2AContextRecord, + context: ServerCallContext | None = None, + ) -> str | None: + """Restore a sidecar-proven input wait and reserve its request-scoped continuation.""" + owner = self._owner(context) + task_id = validate_protocol_id(task.id) + if task.status.state != TaskState.TASK_STATE_INPUT_REQUIRED: + raise ValueError("Recoverable task must be input-required") + if task.context_id != context_record.context_id: + raise ValueError("Recoverable task context does not match context record") + _strip_task_llm_headers(task) + context_id = validate_protocol_id(task.context_id) + async with self._termination_commit_locks.setdefault(context_id, asyncio.Lock()): + async with self._mutation_lock: + control = self._execution_snapshot(context_id) + if control is not None: + if control.get("taskId") != task_id: + return None + phase = control.get("phase") + backup = control.get("backup") + backup_status = backup.get("status") if isinstance(backup, Mapping) else None + terminated_recovery = phase == "terminated" and ( + control.get("releaseReady", False) or backup_status == "blocked" + ) + if not terminated_recovery and control.get("inputHandoffReady") is not True: + return None + + reserve_admission = self._recoverable_input_admission_reserver + if reserve_admission is None: + return None + admission = await reserve_admission( + context_id=context_id, + task_id=task_id, + owner=owner, + ) + if admission is None: + return None + + committed = False + try: + record = self._tasks.get(task_id) + if record is not None and ( + record.context_id != context_id + or (record.active_task is not None and not record.active_task.done()) + ): + return None + persisted_task = self._load_task_snapshot(task_id) if record is None else None + if persisted_task is not None and persisted_task.context_id != context_id: + return None + + live_context = self._contexts.get(context_id) + context_source = live_context or context_record + if context_source.active_task_id not in {None, task_id}: + return None + + current_state = ( + record.state if record is not None else persisted_task.state if persisted_task else None + ) + if current_state == TASK_STATE_INPUT_REQUIRED and context_source.active_task_id is None: + # A prior request may have committed Task/Context and then + # failed before the SDK producer consumed its admission. + # Re-admit from the same sidecar proof without requiring + # the optional session-local mirror to exist after restart. + committed = True + return admission + + updated_at = _task_updated_at_from_sdk_task(task) + output_text = ( + record.output_text + if record is not None + else (persisted_task.output_text if persisted_task else []) + ) + expected_permission_backup_generation = ( + record.expected_permission_backup_generation + if record is not None + else (persisted_task.expected_permission_backup_generation if persisted_task else None) + ) + task_snapshot = A2ATaskSnapshot( + task_id=task_id, + context_id=context_id, + state=TASK_STATE_INPUT_REQUIRED, + owner=owner, + output_text=list(output_text), + updated_at=updated_at, + expected_permission_backup_generation=expected_permission_backup_generation, + ) + context_snapshot = A2AContextSnapshot( + context_id=context_id, + session_id=context_source.session_id, + cwd=context_source.cwd, + telemetry_channel=context_source.telemetry_channel, + active_task_id=None, + ) + + # Recovery is a cross-process handoff boundary. Use the same + # strict write-and-readback contract as termination so a failed + # task write cannot be followed by a released Context. + self._persist_terminated_task_snapshots_strict(task_snapshot, context_snapshot) + + if record is None: + record = _record_from_snapshot(task_snapshot) + self._tasks[task_id] = record + self._metrics.record_task_created() + else: + record.state = TASK_STATE_INPUT_REQUIRED + record.owner = owner + record.updated_at = updated_at + record.active_task = None + record.touch() + self._task_persistence_dirty.discard(task_id) + + if live_context is None: + live_context = context_record + if live_context.lock is None: + live_context.lock = asyncio.Lock() + self._contexts[context_id] = live_context + live_context.active_task_id = None + live_context.touch() + + self._attach_context_metadata(task) + self._attach_pending_permissions(task) + owner_tasks = self._sdk_tasks.setdefault(owner, {}) + previous = owner_tasks.get(task_id) + if previous is not None: + self._remove_sdk_task_from_index(owner, task_id, previous.context_id) + owner_tasks[task_id] = _copy_task(task) + self._sdk_tasks_by_context.setdefault(owner, {}).setdefault(context_id, set()).add(task_id) + committed = True + return admission + finally: + if not committed: + await self.release_recoverable_input_admission(admission) + + def can_continue_local_input_required_task(self, *, context_id: str, task_id: str) -> bool: + """Return whether the owning process can reuse its drained execution directly.""" + + control = self._execution_snapshot(context_id) + return bool( + control + and control.get("taskId") == task_id + and control.get("localInputContinuationReady") is True + ) + + async def release_recoverable_input_admission(self, admission: str | None) -> None: + if admission is None or self._recoverable_input_admission_releaser is None: + return + await self._recoverable_input_admission_releaser(admission) + + def _save_locked( + self, + task: Task, + *, + owner: str, + task_id: str, + ) -> None: + record = self._tasks.get(task_id) + owner_tasks = self._sdk_tasks.setdefault(owner, {}) + previous = owner_tasks.get(task_id) + next_state = _task_state_from_sdk_task(task) + incoming_updated_at = _task_updated_at_from_sdk_task(task) + preserve_terminal = bool( + record is not None + and record.state + in {TASK_STATE_CANCELED, TASK_STATE_COMPLETED, TASK_STATE_FAILED, TASK_STATE_INPUT_REQUIRED} + and self._execution_is_terminating(task.context_id) + ) + if preserve_terminal: + # Queued SDK events must not undo the executor's final state while + # its immutable termination snapshot is being committed or retained. + assert record is not None + task = _copy_task(task) + task.status.state = TaskState.Value("TASK_STATE_" + record.state.upper().replace("-", "_")) + execution = self._execution_snapshot(task.context_id) + # Executor cleanup clears active_task before delayed SDK events finish projecting. + # The control snapshot keeps the same turn's business terminal state until + # mark_execution_started changes executionStatus for a legitimate next turn. + active_finalization = bool( + record is not None + and record.state + in {TASK_STATE_CANCELED, TASK_STATE_COMPLETED, TASK_STATE_FAILED, TASK_STATE_INPUT_REQUIRED} + and next_state in {TASK_STATE_SUBMITTED, TASK_STATE_WORKING} + and ( + (record.active_task is not None and not record.active_task.done()) + or ( + execution is not None + and execution.get("phase") == "running" + and execution.get("taskId") == task_id + and execution.get("executionStatus") == record.state + ) ) - stale_state_projection = bool( - record is not None - and record.state + ) + stale_recovered_working_projection = bool( + record is not None + and record.state == TASK_STATE_WORKING + and previous is not None + and _task_state_from_sdk_task(previous) + in {TASK_STATE_CANCELED, TASK_STATE_COMPLETED, TASK_STATE_FAILED} + ) + stale_state_projection = bool( + record is not None + and ( + record.state in {TASK_STATE_CANCELED, TASK_STATE_COMPLETED, TASK_STATE_FAILED, TASK_STATE_INPUT_REQUIRED} - and next_state in {TASK_STATE_SUBMITTED, TASK_STATE_WORKING} - and not active_finalization - and incoming_updated_at < record.updated_at + or stale_recovered_working_projection ) - if stale_state_projection: - # The SDK consumes executor events asynchronously. An event - # older than a detached final state must not roll either task - # projection back. - return - self._attach_context_metadata(task) - self._attach_pending_permissions(task) - owner_tasks = self._sdk_tasks.setdefault(owner, {}) - previous = owner_tasks.get(task_id) - if previous is not None: - self._remove_sdk_task_from_index(owner, task_id, previous.context_id) - owner_tasks[task_id] = _copy_task(task) - self._sdk_tasks_by_context.setdefault(owner, {}).setdefault(task.context_id, set()).add(task_id) - if preserve_terminal or active_finalization: - # During active finalization the SDK still needs the delayed - # nonterminal frame for stream ordering, but the durable record - # already describes the backup boundary and must stay final. - return - # The SDK saves the full Task before yielding every streaming frame. Executors already - # mirror task records explicitly at output/state durability boundaries, so repeated SDK - # saves with the same projected state only need the session-local mirror below. - persist_shared_snapshot = ( - record is None - or task_id in self._task_persistence_dirty - or record.state != next_state - or record.owner != owner - ) - if record is None: - record = A2ATaskRecord( + and next_state != record.state + and not active_finalization + and incoming_updated_at < record.updated_at + ) + if stale_state_projection: + # The SDK consumes executor events asynchronously. An event older + # than a detached or recovered final state must not roll either + # projection back, including a queued terminal event from the old + # producer that arrives after a recovered execution starts. + # Preserve a delayed nonterminal event on TaskManager's mutable + # object so its status message can be folded into the following + # INPUT_REQUIRED history. Only terminal or recovered-lifecycle + # projections may close/corrupt the current SDK lifecycle and + # therefore need to be repaired in place. + if next_state in { + TASK_STATE_CANCELED, + TASK_STATE_COMPLETED, + TASK_STATE_FAILED, + } or stale_recovered_working_projection: + self._restore_rejected_sdk_projection( + task, + previous=previous, + record=record, + owner_tasks=owner_tasks, task_id=task_id, - context_id=task.context_id, - state=next_state, - owner=owner, - updated_at=incoming_updated_at, ) - self._tasks[task_id] = record - self._metrics.record_task_created() - else: - record.state = next_state - record.owner = owner - record.updated_at = incoming_updated_at - record.touch() - self._mirror_task(record, persist_shared_snapshot=persist_shared_snapshot) + return + self._attach_context_metadata(task) + self._attach_pending_permissions(task) + if previous is not None: + self._remove_sdk_task_from_index(owner, task_id, previous.context_id) + owner_tasks[task_id] = _copy_task(task) + self._sdk_tasks_by_context.setdefault(owner, {}).setdefault(task.context_id, set()).add(task_id) + if preserve_terminal or active_finalization: + # During active finalization the SDK still needs the delayed + # nonterminal frame for stream ordering, but the durable record + # already describes the backup boundary and must stay final. + return + # The SDK saves the full Task before yielding every streaming frame. Executors already + # mirror task records explicitly at output/state durability boundaries, so repeated SDK + # saves with the same projected state only need the session-local mirror below. + persist_shared_snapshot = ( + record is None + or task_id in self._task_persistence_dirty + or record.state != next_state + or record.owner != owner + ) + if record is None: + record = A2ATaskRecord( + task_id=task_id, + context_id=task.context_id, + state=next_state, + owner=owner, + updated_at=incoming_updated_at, + ) + self._tasks[task_id] = record + self._metrics.record_task_created() + else: + record.state = next_state + record.owner = owner + record.updated_at = incoming_updated_at + record.touch() + self._mirror_task(record, persist_shared_snapshot=persist_shared_snapshot) + + def _restore_rejected_sdk_projection( + self, + task: Task, + *, + previous: Task | None, + record: A2ATaskRecord, + owner_tasks: dict[str, Task], + task_id: str, + ) -> None: + """Repair the SDK TaskManager object after rejecting an out-of-date event.""" + + previous_state = _task_state_from_sdk_task(previous) if previous is not None else None + live_active_projection = bool( + record.state in {TASK_STATE_CANCELED, TASK_STATE_COMPLETED, TASK_STATE_FAILED, TASK_STATE_INPUT_REQUIRED} + and previous_state in {TASK_STATE_SUBMITTED, TASK_STATE_WORKING} + and record.active_task is not None + and not record.active_task.done() + ) + if previous is not None and previous.context_id == record.context_id and ( + previous_state == record.state or live_active_projection + ): + task.CopyFrom(previous) + return + restored = Task(id=task_id, context_id=record.context_id) + if previous is not None and previous.context_id == record.context_id: + restored.CopyFrom(previous) + restored.status.CopyFrom( + TaskStatus( + state=TaskState.Value("TASK_STATE_" + record.state.upper().replace("-", "_")), + timestamp=_timestamp_from_epoch(record.updated_at), + ) + ) + self._attach_context_metadata(restored) + self._attach_pending_permissions(restored) + owner_tasks[task_id] = _copy_task(restored) + task.CopyFrom(restored) def _attach_context_metadata(self, task: Task) -> None: context = self._contexts.get(task.context_id) @@ -233,7 +483,7 @@ def _attach_context_metadata(self, task: Task) -> None: if control is not None: iac_code = metadata.setdefault("iac_code", {}) if isinstance(iac_code, dict): - iac_code["executionControl"] = control + iac_code["executionControl"] = ExecutionController.protocol_snapshot(control) ParseDict(metadata, task.metadata) def _attach_pending_permissions(self, task: Task) -> None: @@ -1025,6 +1275,11 @@ async def cancel_task(self, task_id: str) -> bool: record = self._tasks.get(validate_protocol_id(task_id)) if record is None or record.active_task is None or record.active_task.done(): return False + logger.warning( + "A2A task store canceling active task task_id=%s asyncio_task=%s source=cancel_task", + sanitize_strict_text(task_id), + sanitize_strict_text(record.active_task.get_name()), + ) record.active_task.cancel() return True @@ -1033,7 +1288,12 @@ async def cancel_inactive_input_required_task(self, *, task_id: str, context_id: return await self.commit_inactive_execution_task(task_id=task_id, context_id=context_id, cancel_wait=True) async def commit_inactive_execution_task( - self, *, task_id: str, context_id: str, cancel_wait: bool = False + self, + *, + task_id: str, + context_id: str, + cancel_wait: bool = False, + expected_state: str | None = None, ) -> bool: """Strictly commit an execution's final task/context snapshots off the event loop.""" @@ -1050,6 +1310,7 @@ async def commit_inactive_execution_task( or record.context_id != context_id or (record.active_task is not None and not record.active_task.done()) or record.state not in allowed_states + or (expected_state is not None and record.state != expected_state) ): return False if cancel_wait: @@ -1161,6 +1422,12 @@ async def cancel_task_and_wait(self, task_id: str, *, timeout: float | None = No if record is None or record.active_task is None or record.active_task.done(): return False active_task = record.active_task + logger.warning( + "A2A task store canceling active task task_id=%s asyncio_task=%s " + "source=cancel_task_and_wait", + sanitize_strict_text(task_id), + sanitize_strict_text(active_task.get_name()), + ) active_task.cancel() if active_task is asyncio.current_task(): @@ -1208,6 +1475,29 @@ async def is_task_active(self, task_id: str) -> bool: async with self._mutation_lock: return self._task_is_active_locked(task_id) + async def wait_until_task_inactive(self, task_id: str, *, timeout: float) -> None: + """Wait for the current in-process domain owner without canceling it.""" + + task_id = validate_protocol_id(task_id) + loop = asyncio.get_running_loop() + deadline = loop.time() + timeout + while True: + async with self._mutation_lock: + record = self._tasks.get(task_id) + active_task = record.active_task if record is not None else None + if active_task is None or active_task.done(): + return + remaining = deadline - loop.time() + if remaining <= 0: + raise TimeoutError("Timed out waiting for A2A task owner to finish") + try: + await asyncio.wait_for(asyncio.shield(active_task), timeout=remaining) + except asyncio.CancelledError: + if not active_task.cancelled(): + raise + except Exception: + pass + async def has_active_work(self) -> bool: """Return whether shutting down would interrupt in-process A2A work.""" async with self._mutation_lock: @@ -1561,27 +1851,47 @@ def _write_session_snapshot(session_dir: Path, path: Path, data: dict[str, Any]) async def _close_runtime(runtime: Any | None) -> None: + if runtime is None: + return + close_task = asyncio.create_task(_close_runtime_unbounded(runtime)) + try: + done, _pending = await asyncio.wait({close_task}, timeout=_RUNTIME_CLOSE_TIMEOUT_SECONDS) + except asyncio.CancelledError: + close_task.cancel() + close_task.add_done_callback(_consume_runtime_close_result) + raise + if not done: + logger.warning("Timed out closing A2A runtime") + close_task.cancel() + close_task.add_done_callback(_consume_runtime_close_result) + return + _consume_runtime_close_result(close_task) + + +def _consume_runtime_close_result(task: asyncio.Task[Any]) -> None: + try: + task.result() + except asyncio.CancelledError: + pass + except Exception: + logger.exception("Failed to close A2A runtime") + + +async def _close_runtime_unbounded(runtime: Any | None) -> None: if runtime is None: return close = getattr(runtime, "aclose", None) if callable(close): - try: - result = close() - if asyncio.iscoroutine(result): - await result - return - except Exception: - logger.exception("Failed to close A2A runtime") - return + result = close() + if asyncio.iscoroutine(result): + await result + return manager = getattr(runtime, "mcp_manager", None) if manager is not None: - try: - await manager.disconnect_all() - except Exception: - logger.exception("Failed to disconnect A2A MCP manager") + await manager.disconnect_all() agent_runtime = getattr(runtime, "agent_runtime", None) if agent_runtime is not None and agent_runtime is not runtime: - await _close_runtime(agent_runtime) + await _close_runtime_unbounded(agent_runtime) def _runtime_path_directories(runtime: Any | None) -> tuple[list[str], list[str], list[str]]: @@ -1644,7 +1954,11 @@ def close_result(done: asyncio.Task[Any]) -> None: if loop.is_closed(): discard_marker(done) return - loop.create_task(_close_runtime(runtime)) + # This callback is already detached from request/execution progress. + # Start the close body directly so waiters observe the same cleanup + # ordering without adding another scheduling hop. + runtime_close_task = loop.create_task(_close_runtime_unbounded(runtime)) + runtime_close_task.add_done_callback(_consume_runtime_close_result) if discarded_task_waiters is None or discarded_task_waiters.get(done, 0) <= 0: discard_marker(done) diff --git a/src/iac_code/a2a/transports/dispatcher.py b/src/iac_code/a2a/transports/dispatcher.py index 808e0acb..67b9808d 100644 --- a/src/iac_code/a2a/transports/dispatcher.py +++ b/src/iac_code/a2a/transports/dispatcher.py @@ -11,6 +11,7 @@ import httpx from a2a.server.agent_execution.active_task import INTERRUPTED_TASK_STATES, TERMINAL_TASK_STATES +from a2a.server.context import ServerCallContext from a2a.server.events.event_queue import EventQueue from a2a.server.events.event_queue_v2 import QueueShutDown from a2a.server.request_handlers import DefaultRequestHandler @@ -48,6 +49,11 @@ from iac_code.a2a.app import normalize_v03_jsonrpc_version from iac_code.a2a.artifacts import A2AArtifactStore from iac_code.a2a.events import make_text_part +from iac_code.a2a.execution_control import ( + NaturalCompletionGenerationCarrier, + RecoverableInputAdmissionCarrier, + RecoverableInputAdmissionLease, +) from iac_code.a2a.executor import IacCodeA2AExecutor from iac_code.a2a.exposure import normalize_a2a_exposure_types from iac_code.a2a.input_required import parse_permission_response @@ -88,6 +94,16 @@ from iac_code.a2a.push_secrets import A2APushSecretKeyring from iac_code.a2a.push_worker import A2APushDeliveryWorker from iac_code.a2a.request_mode import resolve_request_run_mode +from iac_code.a2a.request_scoped_active_task import ( + DirectMessageRequestStarted, + DirectPipelineRecoveryRequiredError, + DirectPipelineRouteGate, + DirectPipelineRouteGateCarrier, + DirectPipelineRouteOutcome, + PipelineLifecycleEventQueueCarrier, + RequestScopedActiveTask, + RequestScopedActiveTaskRegistry, +) from iac_code.a2a.runtime_registry import A2ARuntimeOwner, A2ARuntimeRegistration, register_runtime_owner from iac_code.a2a.task_store import A2ATaskStore from iac_code.i18n import _ @@ -104,6 +120,7 @@ logger = logging.getLogger(__name__) _ACTIVE_MESSAGE_STREAM_COMPLETED = object() _ASGI_STREAM_END = object() +_RECOVERABLE_INPUT_ADMISSION_STATE_KEY = "iac_code.recoverable_input_admission" @dataclass @@ -398,6 +415,8 @@ def create_runtime_components( set_execution_control_provider( execution_control_service.snapshot_for_context, execution_control_service.has_active_work, + execution_control_service.reserve_recoverable_input_continuation, + execution_control_service.release_recoverable_input_continuation, ) if push_notifications: assert persistence is not None @@ -504,6 +523,11 @@ def __init__( **kwargs: Any, ) -> None: super().__init__(*args, **kwargs) + self._active_task_registry = RequestScopedActiveTaskRegistry( + agent_executor=self.agent_executor, + task_store=self.task_store, + push_sender=self._push_sender, + ) self._backup_service = backup_service or SessionBackupService() self._metrics = metrics or NoOpA2AMetrics() self._detached_message_producers: set[asyncio.Task[Any]] = set() @@ -516,6 +540,158 @@ async def on_list_tasks(self, params: ListTasksRequest, context): self._validate_extensions(context) return await super().on_list_tasks(params, context) + async def _setup_active_task(self, params: SendMessageRequest, call_context): + active_task, request_context = await super()._setup_active_task(params, call_context) + PipelineLifecycleEventQueueCarrier.attach(request_context) + NaturalCompletionGenerationCarrier.prepare(request_context) + admission = self._peek_recoverable_input_admission(call_context) + lease = None + if admission is not None and isinstance(self.task_store, A2ATaskStore): + lease = RecoverableInputAdmissionLease( + admission, + acknowledge_enqueue=lambda token: self._acknowledge_recoverable_input_enqueue(call_context, token), + release=self.task_store.release_recoverable_input_admission, + ) + RecoverableInputAdmissionCarrier.attach(request_context, lease) + return active_task, request_context + + @staticmethod + def _peek_recoverable_input_admission(call_context) -> str | None: + state = getattr(call_context, "state", None) + admission = state.get(_RECOVERABLE_INPUT_ADMISSION_STATE_KEY) if isinstance(state, dict) else None + return admission if isinstance(admission, str) and admission else None + + @staticmethod + def _acknowledge_recoverable_input_enqueue(call_context, admission: str) -> bool: + state = getattr(call_context, "state", None) + if not isinstance(state, dict) or state.get(_RECOVERABLE_INPUT_ADMISSION_STATE_KEY) != admission: + return False + state.pop(_RECOVERABLE_INPUT_ADMISSION_STATE_KEY, None) + return True + + @staticmethod + def _stage_recoverable_input_admission(call_context, admission: str | None) -> None: + state = getattr(call_context, "state", None) + if not isinstance(state, dict): + if admission is None: + return + raise RuntimeError("Server call context cannot carry a recovery admission") + state.pop(_RECOVERABLE_INPUT_ADMISSION_STATE_KEY, None) + if admission is not None: + state[_RECOVERABLE_INPUT_ADMISSION_STATE_KEY] = admission + + @staticmethod + def _take_recoverable_input_admission(call_context) -> str | None: + state = getattr(call_context, "state", None) + if isinstance(state, dict): + admission = state.pop(_RECOVERABLE_INPUT_ADMISSION_STATE_KEY, None) + return admission if isinstance(admission, str) and admission else None + return None + + async def _reconcile_and_replace_recovered_sdk_task( + self, + params: SendMessageRequest, + context: Any, + ) -> str | None: + task_id = params.message.task_id + if not task_id: + return None + context_id = params.message.context_id + if not context_id: + return None + registry = cast("RequestScopedActiveTaskRegistry | None", getattr(self, "_active_task_registry", None)) + if registry is None: + return await self._reconcile_recoverable_pipeline_task(params, context) + reconcile_and_replace = getattr(registry, "reconcile_and_replace_for_recovery", None) + if not callable(reconcile_and_replace): + return await self._reconcile_recoverable_pipeline_task(params, context) + return await reconcile_and_replace( + task_id, + call_context=context, + context_id=context_id, + acquire_admission=lambda: self._reconcile_recoverable_pipeline_task(params, context), + release_admission=( + self.task_store.release_recoverable_input_admission + if isinstance(self.task_store, A2ATaskStore) + else None + ), + ) + + async def _release_untransferred_recovery( + self, + params: SendMessageRequest, + context: Any, + ) -> None: + admission = self._take_recoverable_input_admission(context) + if admission is None: + return + task_id = params.message.task_id + if task_id: + registry = cast("RequestScopedActiveTaskRegistry", self._active_task_registry) + await registry.cancel_recovery_reservation(task_id, admission) + if isinstance(self.task_store, A2ATaskStore): + await self.task_store.release_recoverable_input_admission(admission) + + async def _finalize_natural_execution(self, params: SendMessageRequest, context: Any) -> None: + """Settle the execution only after the request's response stream is exhausted.""" + + await self._finalize_natural_execution_boundary( + task_id=getattr(params.message, "task_id", None), + context_id=getattr(params.message, "context_id", None), + context=context, + ) + + async def _finalize_natural_execution_boundary(self, *, task_id: Any, context_id: Any, context: Any) -> None: + """Settle a pending natural completion once its response stream is delivered. + + The generation is carried on the delivering request's ``ServerCallContext``; + when it is absent (e.g. a bare observer that did not deliver the executor + turn) ``mark_delivered`` returns ``None`` and this is a no-op, so the + controller never claims natural completion from a non-delivering path. + """ + + completion_generation = NaturalCompletionGenerationCarrier.mark_delivered(context) + finalize = getattr(getattr(self, "agent_executor", None), "finalize_natural_execution", None) + if ( + not callable(finalize) + or not task_id + or not context_id + or completion_generation is None + or not isinstance(self.task_store, A2ATaskStore) + ): + return + await finalize( + context_id=context_id, + task_id=task_id, + owner=self.task_store.owner_for_context(context), + completion_generation=completion_generation, + ) + + async def _settle_nonstream_input_required_result( + self, + result: Message | Task, + params: SendMessageRequest, + context: Any, + ) -> Message | Task: + """Wait for the accepted request before finalizing an input boundary.""" + + if ( + not isinstance(result, Task) + or result.status.state != TaskState.TASK_STATE_INPUT_REQUIRED + or bool(getattr(getattr(params, "configuration", None), "return_immediately", False)) + ): + return result + get_active_task = getattr(self._active_task_registry, "get", None) + if not callable(get_active_task): + return result + active_task = await get_active_task(result.id) + if isinstance(active_task, RequestScopedActiveTask): + await active_task.wait_for_accepted_requests() + refreshed = await self.task_store.get(result.id, context) + if refreshed is not None: + return apply_history_length(refreshed, params.configuration) + return result + async def on_message_send(self, params: SendMessageRequest, context): self._validate_extensions(context) self._validate_pipeline_message_request(params) @@ -539,11 +715,23 @@ async def on_message_send(self, params: SendMessageRequest, context): task=task, ): pass + await self._finalize_natural_execution(params, context) refreshed = await self.task_store.get(permission_response.task_id, context) return refreshed or task await self._hydrate_recoverable_pipeline_task_id(params) - await self._reconcile_recoverable_pipeline_task(params, context) - return await super().on_message_send(params, context) + admission = await self._reconcile_and_replace_recovered_sdk_task(params, context) + self._stage_recoverable_input_admission(context, admission) + try: + result = await super().on_message_send(params, context) + result = await self._settle_nonstream_input_required_result( + result, + params, + context, + ) + await self._finalize_natural_execution(params, context) + return result + finally: + await self._release_untransferred_recovery(params, context) async def on_message_send_stream(self, params: SendMessageRequest, context): self._validate_extensions(context) @@ -552,64 +740,105 @@ async def on_message_send_stream(self, params: SendMessageRequest, context): if permission_response is not None: if not params.message.task_id: params.message.task_id = permission_response.task_id - resolve = getattr(getattr(self, "agent_executor", None), "resolve_sideband_permission", None) - if callable(resolve): - ack = await resolve(permission_response, metadata=params.message) - if ack is not None: - yield ack - return + task_active = False + task_id_for_check = params.message.task_id + task_store = getattr(self, "task_store", None) + if task_id_for_check and isinstance(task_store, A2ATaskStore): + task_active = await task_store.is_task_active(task_id_for_check) + if not task_active: + resolve = getattr(getattr(self, "agent_executor", None), "resolve_sideband_permission", None) + if callable(resolve): + ack = await resolve(permission_response, metadata=params.message) + if ack is not None: + yield ack + return if permission_response is None: await self._hydrate_recoverable_pipeline_task_id(params) - await self._reconcile_recoverable_pipeline_task(params, context) - task_id = params.message.task_id or None - if task_id and isinstance(self.task_store, A2ATaskStore) and await self.task_store.is_task_active(task_id): - task = await self.task_store.get(task_id, context) - active_task = await self._active_task_registry.get(task_id) - if ( - task is not None - and active_task is not None - and task.status.state not in TERMINAL_TASK_STATES - and (task.status.state not in INTERRUPTED_TASK_STATES or permission_response is not None) - ): - active_stream = self._on_active_message_send_stream( - params, - context, - task=task, - active_task=active_task, + admission = await self._reconcile_and_replace_recovered_sdk_task(params, context) + else: + admission = None + self._stage_recoverable_input_admission(context, admission) + try: + task_id = params.message.task_id or None + if task_id and isinstance(self.task_store, A2ATaskStore) and await self.task_store.is_task_active(task_id): + task = await self.task_store.get(task_id, context) + active_task = await self._active_task_registry.get(task_id) + if ( + task is not None + and active_task is not None + and task.status.state not in TERMINAL_TASK_STATES + and (task.status.state not in INTERRUPTED_TASK_STATES or permission_response is not None) + ): + route_gate = ( + DirectPipelineRouteGate() + if permission_response is None and resolve_request_run_mode(params.message) is RunMode.PIPELINE + else None + ) + active_stream = self._on_active_message_send_stream( + params, + context, + task=task, + active_task=active_task, + route_gate=route_gate, + ) + recovery_required = False + try: + try: + async for event in active_stream: + yield event + except DirectPipelineRecoveryRequiredError: + recovery_required = True + finally: + await active_stream.aclose() + if not recovery_required: + await self._finalize_natural_execution(params, context) + return + admission = await self._wait_for_direct_pipeline_recovery(params, context) + self._stage_recoverable_input_admission(context, admission) + if permission_response is not None and isinstance(self.task_store, A2ATaskStore): + task = await self.task_store.get(permission_response.task_id, context) + if task is not None: + direct_stream = self._on_inactive_permission_send_stream(params, context, task=task) + try: + async for event in direct_stream: + yield event + finally: + await direct_stream.aclose() + await self._finalize_natural_execution(params, context) + return + base_stream = super().on_message_send_stream(params, context) + pipeline_mode = resolve_request_run_mode(params.message) is RunMode.PIPELINE + tracked_stream = ( + base_stream + if pipeline_mode + else _iterate_with_pipeline_transport_tracking( + base_stream, + task_id=getattr(params.message, "task_id", None) or None, + context_id=getattr(params.message, "context_id", None) or None, ) - try: - async for event in active_stream: - yield event - finally: - await active_stream.aclose() - return - if permission_response is not None and isinstance(self.task_store, A2ATaskStore): - task = await self.task_store.get(permission_response.task_id, context) - if task is not None: - direct_stream = self._on_inactive_permission_send_stream(params, context, task=task) - try: - async for event in direct_stream: - yield event - finally: - await direct_stream.aclose() - return - base_stream = super().on_message_send_stream(params, context) - tracked_stream = ( - base_stream - if resolve_request_run_mode(params.message) is RunMode.PIPELINE - else _iterate_with_pipeline_transport_tracking( - base_stream, - task_id=getattr(params.message, "task_id", None) or None, - context_id=getattr(params.message, "context_id", None) or None, ) - ) - try: - async for event in tracked_stream: - mark_pipeline_transport_delivery_dequeued(event) - yield event - acknowledge_pipeline_transport_delivery(event) + natural_boundary_delivered = False + stream_exhausted = False + close_after_natural_boundary = False + try: + async for event in tracked_stream: + mark_pipeline_transport_delivery_dequeued(event) + event_state = _task_event_state(event) + natural_boundary_delivered = not pipeline_mode and ( + event_state in TERMINAL_TASK_STATES or event_state == TaskState.TASK_STATE_INPUT_REQUIRED + ) + yield event + acknowledge_pipeline_transport_delivery(event) + stream_exhausted = True + except GeneratorExit: + close_after_natural_boundary = natural_boundary_delivered + raise + finally: + await tracked_stream.aclose() + if stream_exhausted or close_after_natural_boundary: + await self._finalize_natural_execution(params, context) finally: - await tracked_stream.aclose() + await self._release_untransferred_recovery(params, context) async def _on_inactive_permission_send_stream(self, params: SendMessageRequest, context, *, task: Task): """Resume an existing input boundary without asking the SDK to recreate its task.""" @@ -652,7 +881,13 @@ async def run_permission_response() -> None: # finish exactly once even if the response transport disappears. handed_off = True drain_task = asyncio.create_task( - self._drain_inactive_permission_response(queue, producer, completed), + self._drain_inactive_permission_response( + queue, + producer, + completed, + params=params, + context=context, + ), name=f"a2a-detached-permission-producer-{task.id}", ) self._detached_message_producers.add(drain_task) @@ -664,11 +899,14 @@ async def run_permission_response() -> None: with suppress(asyncio.CancelledError): await producer - @staticmethod async def _drain_inactive_permission_response( + self, queue: asyncio.Queue[Any], producer: asyncio.Task[None], completed: object, + *, + params: SendMessageRequest, + context: Any, ) -> None: try: while True: @@ -676,6 +914,7 @@ async def _drain_inactive_permission_response( if value is completed: break await producer + await self._finalize_natural_execution(params, context) except asyncio.CancelledError: if not producer.done(): producer.cancel() @@ -684,7 +923,15 @@ async def _drain_inactive_permission_response( except BaseException: logger.debug("Detached permission response continuation failed", exc_info=True) - async def _on_active_message_send_stream(self, params: SendMessageRequest, context, *, task: Task, active_task): + async def _on_active_message_send_stream( + self, + params: SendMessageRequest, + context, + *, + task: Task, + active_task, + route_gate: DirectPipelineRouteGate | None = None, + ): request_context = await self._request_context_builder.build( params=params, task_id=task.id, @@ -692,42 +939,70 @@ async def _on_active_message_send_stream(self, params: SendMessageRequest, conte task=task, context=context, ) - async with active_task._lock: - if active_task._is_finished.is_set(): - raise InvalidParamsError(_("Task {task_id} is already completed.").format(task_id=active_task.task_id)) - active_task._reference_count += 1 - tapped_queue = await active_task._event_queue_subscribers.tap() + if route_gate is not None: + DirectPipelineRouteGateCarrier.attach(request_context, route_gate) + direct_message_lock = active_task.direct_message_lock + await direct_message_lock.acquire() + reference_registered = False + tapped_queue = None + producer_task = None + producer_owns_lock = False + delivery_tracker = None - async def run_active_message() -> None: - try: - await self.agent_executor.execute(request_context, active_task._event_queue_agent) - finally: - await self._wait_for_active_message_events(active_task) - with suppress(QueueShutDown): - await tapped_queue._put_internal((_ACTIVE_MESSAGE_STREAM_COMPLETED, None)) + try: + async with active_task._lock: + if active_task._is_finished.is_set(): + raise InvalidParamsError( + _("Task {task_id} is already completed.").format(task_id=active_task.task_id) + ) + active_task._reference_count += 1 + reference_registered = True + tapped_queue = await active_task._event_queue_subscribers.tap() + request_started = route_gate.marker if route_gate is not None else DirectMessageRequestStarted() + if route_gate is None: + await active_task._event_queue_agent.enqueue_event(request_started) + + async def run_active_message() -> None: + try: + await self.agent_executor.execute(request_context, active_task._event_queue_agent) + finally: + try: + await self._wait_for_active_message_events(active_task) + with suppress(QueueShutDown): + await tapped_queue._put_internal((_ACTIVE_MESSAGE_STREAM_COMPLETED, None)) + finally: + direct_message_lock.release() - delivery_tracker = create_pipeline_transport_delivery_tracker() - with bind_pipeline_transport_delivery_route( - delivery_tracker, - task_id=task.id, - context_id=getattr(task, "context_id", None) or params.message.context_id, - ): - with bind_pipeline_transport_delivery_tracker(delivery_tracker): - producer_task = asyncio.create_task(run_active_message()) + delivery_tracker = create_pipeline_transport_delivery_tracker() + with bind_pipeline_transport_delivery_route( + delivery_tracker, + task_id=task.id, + context_id=getattr(task, "context_id", None) or params.message.context_id, + ): + with bind_pipeline_transport_delivery_tracker(delivery_tracker): + producer_task = asyncio.create_task(run_active_message()) + producer_owns_lock = True - try: + boundary_reached = False while True: try: dequeued = await tapped_queue.dequeue_event() except QueueShutDown: break - event, _updated_task = cast(Any, dequeued) - if event is _ACTIVE_MESSAGE_STREAM_COMPLETED: - tapped_queue.task_done() - break - if isinstance(event, BaseException): - raise event + event, updated_task = cast(Any, dequeued) try: + if event is _ACTIVE_MESSAGE_STREAM_COMPLETED: + break + if isinstance(event, DirectMessageRequestStarted): + if event is request_started: + boundary_reached = True + continue + if not boundary_reached: + continue + if isinstance(event, BaseException): + raise event + if RequestScopedActiveTask.is_stale_terminal_projection(event, updated_task): + continue mark_pipeline_transport_delivery_dequeued(event) if isinstance(event, Task): self._validate_task_id_match(task.id, event.id) @@ -737,27 +1012,127 @@ async def run_active_message() -> None: finally: tapped_queue.task_done() acknowledge_pipeline_transport_delivery(event) - except (asyncio.CancelledError, GeneratorExit): - raise - finally: + if route_gate is not None: + assert producer_task is not None + await producer_task + producer_task = None + if route_gate.outcome is DirectPipelineRouteOutcome.RECOVERY_REQUIRED: + raise DirectPipelineRecoveryRequiredError + if route_gate.outcome is DirectPipelineRouteOutcome.PENDING: + raise RuntimeError("Active Pipeline follow-up did not resolve its route gate") + except (asyncio.CancelledError, GeneratorExit): + raise + finally: + if delivery_tracker is not None: close_pipeline_transport_delivery_tracker(delivery_tracker) + if tapped_queue is not None: await tapped_queue.close(immediate=True) + if reference_registered: async with active_task._lock: active_task._reference_count -= 1 await active_task._maybe_cleanup() - if producer_task.done(): + if not producer_owns_lock: + direct_message_lock.release() + if producer_task is not None: + if route_gate is not None: + cleanup_task = asyncio.create_task( + self._finish_detached_active_pipeline_request( + producer_task, + params=params, + context=context, + route_gate=route_gate, + ), + name=f"a2a-detached-pipeline-recovery-{task.id}", + ) + self._track_detached_message_producer(cleanup_task) + elif producer_task.done(): await self._cleanup_active_message_producer(producer_task, task.id) else: cleanup_task = asyncio.create_task( self._cleanup_active_message_producer(producer_task, task.id), name=f"a2a-detached-message-producer-{task.id}", ) - detached_producers = getattr(self, "_detached_message_producers", None) - if detached_producers is None: - detached_producers = set() - self._detached_message_producers = detached_producers - detached_producers.add(cleanup_task) - cleanup_task.add_done_callback(detached_producers.discard) + self._track_detached_message_producer(cleanup_task) + + def _track_detached_message_producer(self, task: asyncio.Task[Any]) -> None: + detached_producers = getattr(self, "_detached_message_producers", None) + if detached_producers is None: + detached_producers = set() + self._detached_message_producers = detached_producers + detached_producers.add(task) + task.add_done_callback(detached_producers.discard) + + async def _finish_detached_active_pipeline_request( + self, + producer_task: asyncio.Task[Any], + *, + params: SendMessageRequest, + context: Any, + route_gate: DirectPipelineRouteGate, + ) -> None: + task_id = params.message.task_id or "unknown" + detached_context = self._copy_detached_recovery_context(context) + base_stream = None + try: + await producer_task + if route_gate.outcome is DirectPipelineRouteOutcome.ACTIVE: + return + if route_gate.outcome is DirectPipelineRouteOutcome.PENDING: + raise RuntimeError("Active Pipeline follow-up did not resolve its route gate") + + admission = await self._wait_for_direct_pipeline_recovery(params, detached_context) + self._stage_recoverable_input_admission(detached_context, admission) + base_stream = DefaultRequestHandler.on_message_send_stream(self, params, detached_context) + async for _event in base_stream: + pass + except asyncio.CancelledError: + if not producer_task.done(): + producer_task.cancel() + await asyncio.gather(producer_task, return_exceptions=True) + raise + except Exception as exc: + logger.error( + "Detached active Pipeline recovery task_id=%s failed: %s", + sanitize_strict_text(task_id), + sanitize_strict_text(str(exc)), + ) + finally: + if base_stream is not None: + await base_stream.aclose() + await self._release_untransferred_recovery(params, detached_context) + + @staticmethod + def _copy_detached_recovery_context(context: ServerCallContext) -> ServerCallContext: + state = dict(context.state) + state.pop(_RECOVERABLE_INPUT_ADMISSION_STATE_KEY, None) + return ServerCallContext( + state=state, + user=context.user, + tenant=context.tenant, + requested_extensions=set(context.requested_extensions), + ) + + async def _wait_for_direct_pipeline_recovery( + self, + params: SendMessageRequest, + context: Any, + ) -> str: + task_id = params.message.task_id + context_id = params.message.context_id + if not task_id or not context_id: + raise InvalidParamsError("Pipeline recovery requires task and context ids") + wait = getattr(self.agent_executor, "wait_until_recoverable_pipeline_input", None) + if callable(wait): + await wait(context_id=context_id, task_id=task_id) + if isinstance(self.task_store, A2ATaskStore): + try: + await self.task_store.wait_until_task_inactive(task_id, timeout=30) + except TimeoutError as exc: + raise InvalidParamsError("Pipeline continuation is still finalizing; retry the same request.") from exc + admission = await self._reconcile_and_replace_recovered_sdk_task(params, context) + if admission is None: + raise InvalidParamsError("Pipeline continuation is not ready; retry the same request.") + return admission async def _hydrate_recoverable_pipeline_task_id(self, params: SendMessageRequest) -> None: if resolve_request_run_mode(params.message) is not RunMode.PIPELINE or not isinstance( @@ -789,48 +1164,53 @@ async def _hydrate_recoverable_pipeline_task_id(self, params: SendMessageRequest if task_id: message.task_id = task_id - async def _reconcile_recoverable_pipeline_task(self, params: SendMessageRequest, context) -> None: + async def _reconcile_recoverable_pipeline_task(self, params: SendMessageRequest, context) -> str | None: if resolve_request_run_mode(params.message) is not RunMode.PIPELINE or not isinstance( self.task_store, A2ATaskStore ): - return + return None message = getattr(params, "message", None) if message is None: - return + return None task_id = getattr(message, "task_id", None) if not isinstance(task_id, str) or not task_id: - return + return None try: task = await self.task_store.get(task_id, context) except Exception: logger.debug("Failed to load A2A task %s before pipeline reconciliation", task_id, exc_info=True) - return - if task is None or task.status.state not in TERMINAL_TASK_STATES: - return + return None + if task is None or ( + task.status.state not in TERMINAL_TASK_STATES and task.status.state != TaskState.TASK_STATE_INPUT_REQUIRED + ): + return None context_id = getattr(message, "context_id", None) or getattr(task, "context_id", None) if not isinstance(context_id, str) or not context_id: - return + return None try: context_record = await self.task_store.get_context_record(context_id) recoverable_task_id = recoverable_task_id_from_sidecar( cwd=context_record.cwd, session_id=context_record.session_id, context_id=context_id, + include_running=False, ) except Exception: logger.debug("Failed to inspect recoverable A2A pipeline task %s", task_id, exc_info=True) - return + return None if recoverable_task_id != task_id: - return + return None - task.status.CopyFrom(TaskStatus(state=TaskState.Name(TaskState.TASK_STATE_INPUT_REQUIRED))) - task.status.timestamp.GetCurrentTime() - await self.task_store.save(task, context) - if context_record.active_task_id == task_id: - context_record.active_task_id = None - context_record.touch() - self.task_store.mirror_context(context_record) + if task.status.state in TERMINAL_TASK_STATES: + task.status.CopyFrom(TaskStatus(state=TaskState.Name(TaskState.TASK_STATE_INPUT_REQUIRED))) + task.status.timestamp.GetCurrentTime() + if self.task_store.can_continue_local_input_required_task(context_id=context_id, task_id=task_id): + return None + admission = await self.task_store.reconcile_recoverable_input_required_task(task, context_record, context) + if admission is None: + raise InvalidParamsError("Pipeline continuation is already being recovered; retry the same request.") + return admission async def _wait_for_active_message_events(self, active_task) -> None: event_queue_agent = getattr(active_task, "_event_queue_agent", None) @@ -870,9 +1250,10 @@ async def on_cancel_task(self, params: CancelTaskRequest, context) -> Task | Non if durable_cancel == "lost": return await self._reconcile_inactive_pipeline_input_required_task(task, context) if durable_cancel == "normal": - canceled = await self._reconcile_inactive_terminal_task(task, context, "canceled") - await self.task_store.discard_context_runtime(task.context_id) - return canceled + canceled = await self._cancel_inactive_normal_input_required_task(task, context) + if canceled is not None: + return canceled + raise TaskNotCancelableError canceled_task = await self._cancel_inactive_pipeline_waiting_input_task(task, context) if canceled_task is not None: return canceled_task @@ -985,7 +1366,17 @@ async def _cancel_inactive_normal_input_required_task(self, task: Task, context) return None if waiting_task_id is not None: return None - canceled = await self._reconcile_inactive_terminal_task(task, context, "canceled") + committed = await self.task_store.cancel_inactive_input_required_task( + task_id=task.id, + context_id=task.context_id, + ) + if not committed: + return None + canceled = await self._reconcile_inactive_terminal_task( + task, + context, + "canceled", + ) await self.task_store.discard_context_runtime(task.context_id) return canceled @@ -1043,20 +1434,42 @@ async def on_subscribe_to_task(self, params: SubscribeToTaskRequest, context): raise TaskNotFoundError(f"Task {params.id} is not active") active_task_registry = getattr(self, "_active_task_registry", None) active_task = await active_task_registry.get(params.id) if active_task_registry is not None else None + terminal_state_seen = False + interrupted = False async for event in super().on_subscribe_to_task(params, context): event_state = _task_event_state(event) + terminal_state_seen = terminal_state_seen or event_state in TERMINAL_TASK_STATES yield event - if event_state in TERMINAL_TASK_STATES or event_state in INTERRUPTED_TASK_STATES: - return - - # a2a-sdk 1.1 can close its subscriber queue after the producer finishes but - # before the consumer publishes the last update. Recover the persisted terminal - # snapshot so subscribers do not observe a silent, non-terminal stream ending. + if terminal_state_seen: + break + if event_state in INTERRUPTED_TASK_STATES: + interrupted = True + break + if not interrupted and not terminal_state_seen: + # a2a-sdk 1.1 can close its subscriber queue after the producer finishes but + # before the consumer publishes the last update. Recover the persisted terminal + # snapshot so subscribers do not observe a silent, non-terminal stream ending. + if active_task is not None: + await active_task._is_finished.wait() + final_task = await self.task_store.get(params.id, context) + if final_task is not None and final_task.status.state in TERMINAL_TASK_STATES: + yield final_task + # Settle any natural completion the delivered executor turn registered. Draining + # the accepted request first ensures the executor turn's finally has registered + # its natural-completion generation, so the finalize claims and settles the + # release synchronously before this subscriber stream closes instead of leaving + # the controller pinned until a later GET retry. No-op unless this delivering + # request carried a natural-completion generation. if active_task is not None: - await active_task._is_finished.wait() - final_task = await self.task_store.get(params.id, context) - if final_task is not None and final_task.status.state in TERMINAL_TASK_STATES: - yield final_task + try: + await active_task.wait_for_accepted_requests() + except InvalidParamsError: + pass + await self._finalize_natural_execution_boundary( + task_id=params.id, + context_id=task.context_id, + context=context, + ) async def on_create_task_push_notification_config( self, params: TaskPushNotificationConfig, context diff --git a/src/iac_code/services/providers/aliyun.py b/src/iac_code/services/providers/aliyun.py index a66f6d59..16cc615e 100644 --- a/src/iac_code/services/providers/aliyun.py +++ b/src/iac_code/services/providers/aliyun.py @@ -111,6 +111,12 @@ class AliyunCredential: credential_source: str = field(default="", repr=False, compare=False) credential_source_path: str = field(default="", repr=False, compare=False) + def refresh_from(self, credential: "AliyunCredential") -> None: + """Replace this request-scoped credential without changing its object identity.""" + + for credential_field in fields(self): + setattr(self, credential_field.name, getattr(credential, credential_field.name)) + _aliyun_credential_override: contextvars.ContextVar[AliyunCredential | None] = contextvars.ContextVar( "iac_code_aliyun_credential_override", default=None @@ -176,6 +182,12 @@ def use_aliyun_credential(credential: AliyunCredential) -> Iterator[None]: _aliyun_credential_override.reset(token) +def current_aliyun_credential_override() -> AliyunCredential | None: + """Return only the request-scoped override, without consulting global credential sources.""" + + return _aliyun_credential_override.get() + + def mask_sensitive(value: str) -> str: """Mask a sensitive value with '*' characters of the same length.""" if not value: diff --git a/tests/a2a/test_app.py b/tests/a2a/test_app.py index bb8edd9b..fb9e4b7b 100644 --- a/tests/a2a/test_app.py +++ b/tests/a2a/test_app.py @@ -43,7 +43,7 @@ from iac_code.a2a.pipeline_executor import recoverable_task_id_from_sidecar from iac_code.a2a.pipeline_journal import A2APipelineJournal from iac_code.a2a.pipeline_snapshot import A2APipelineSnapshotStore, reduce_pipeline_events -from iac_code.a2a.task_store import A2ATaskStore +from iac_code.a2a.task_store import A2ATaskStore, _task_updated_at_from_sdk_task from iac_code.a2a.transports.dispatcher import create_runtime_components from iac_code.mcp.errors import MCPNeedsAuthError from iac_code.pipeline.engine.events import PipelineEvent, PipelineEventType @@ -62,6 +62,7 @@ from iac_code.services.session_metadata import SESSION_LAYOUT_VERSION_V2, SessionMetadata, write_session_metadata from iac_code.services.session_storage import SessionStorage from iac_code.types.stream_events import TextDeltaEvent, ToolResultEvent +from iac_code.utils.state_io import atomic_write_json from .fakes import FakeAgentLoop, FakeRuntime @@ -440,9 +441,7 @@ def test_pipeline_state_endpoint_can_return_delta_without_snapshot(tmp_path) -> ) with TestClient(app) as client: - response = client.get( - "/iac-code/pipeline/state?contextId=ctx-1&afterSequence=1&includeSnapshot=false" - ) + response = client.get("/iac-code/pipeline/state?contextId=ctx-1&afterSequence=1&includeSnapshot=false") assert response.status_code == 200 data = response.json() @@ -2871,37 +2870,220 @@ async def run_streaming(self, prompt: str): monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", lambda options: runtime) components = create_runtime_components(model="qwen3.6-plus", host="127.0.0.1", port=41242) call_context = ServerCallContext() + stream = None + collector = None - result = await components.handler.on_message_send( - SendMessageRequest( - message=Message( - message_id="msg-1", - role=Role.ROLE_USER, - parts=[Part(text="hello")], - metadata={"iac_code": {"cwd": str(tmp_path)}}, + try: + result = await components.handler.on_message_send( + SendMessageRequest( + message=Message( + message_id="msg-1", + role=Role.ROLE_USER, + parts=[Part(text="hello")], + metadata={"iac_code": {"cwd": str(tmp_path)}}, + ), + configuration=SendMessageConfiguration(accepted_output_modes=["text/plain"], return_immediately=True), ), - configuration=SendMessageConfiguration(accepted_output_modes=["text/plain"], return_immediately=True), - ), - call_context, - ) - assert isinstance(result, Task) - - stream = components.handler.on_subscribe_to_task(SubscribeToTaskRequest(id=result.id), call_context) - first_event = await asyncio.wait_for(anext(stream), timeout=1) - release.set() - remaining_events = [] + call_context, + ) + assert isinstance(result, Task) - async def collect_remaining_events() -> None: + stream = components.handler.on_subscribe_to_task(SubscribeToTaskRequest(id=result.id), call_context) + first_event_ready = asyncio.get_running_loop().create_future() + remaining_events = [] + + async def collect_events() -> None: + try: + async for event in stream: + if not first_event_ready.done(): + first_event_ready.set_result(event) + else: + remaining_events.append(event) + except BaseException as exc: + if not first_event_ready.done(): + first_event_ready.set_exception(exc) + raise + finally: + if not first_event_ready.done(): + first_event_ready.set_exception(AssertionError("subscription ended before the initial task")) + + collector = asyncio.create_task(collect_events(), name="collect-active-task-subscription") + first_event = await asyncio.shield(first_event_ready) + release.set() + await asyncio.shield(collector) + + execution_state = await components.execution_control_service.observe( + context_id=result.context_id, + owner='', + ) + assert isinstance(first_event, Task) + assert first_event.id == result.id + assert "second" in json.dumps([event.__class__.__name__ + str(event) for event in remaining_events]) + assert prompts == ["hello"] + assert execution_state["phase"] == "terminated" + assert execution_state["terminationReason"] == "natural_completion" + assert execution_state["releaseReady"] is True + finally: + release.set() + try: + if collector is not None and not collector.done(): + await asyncio.shield(collector) + finally: + try: + if stream is not None: + await stream.aclose() + finally: + await components.aclose() + + +class NaturalCompletionProjectionGate: + """Hold one newer SDK working projection across the executor cleanup boundary.""" + + def __init__(self, store: A2ATaskStore) -> None: + loop = asyncio.get_running_loop() + self._store = store + self._save = store.save + self._get_task_record = store.get_task_record + self._projection_release = loop.create_future() + self._record_read = loop.create_future() + self._task_id: str | None = None + self._projection_intercepted = False + self._projection_overwrote = False + self.projection_waiting = asyncio.Event() + self.projection_saved = asyncio.Event() + self.final_record_read_waiting = asyncio.Event() + self.agent_release = asyncio.Event() + self.agent_completed = asyncio.Event() + self.remaining_events = [] + + def install(self, monkeypatch) -> None: + monkeypatch.setattr(self._store, "save", self.save) + monkeypatch.setattr(self._store, "get_task_record", self.get_task_record) + + def target(self, task_id: str) -> None: + self._task_id = task_id + + async def run_streaming(self, _prompt: str): + yield TextDeltaEvent(text="first") + await self.agent_release.wait() + yield TextDeltaEvent(text="second") + self.agent_completed.set() + + async def save(self, task, context=None) -> None: + record = self._store._tasks.get(task.id) + should_delay = bool( + not self._projection_intercepted + and task.id == self._task_id + and task.status.state == TaskState.TASK_STATE_WORKING + and record is not None + and record.state == "input-required" + and _task_updated_at_from_sdk_task(task) > record.updated_at + ) + if not should_delay: + await self._save(task, context) + return + + self._projection_intercepted = True + self.projection_waiting.set() + await self._projection_release + await self._save(task, context) + record = self._store._tasks[task.id] + if record.state == "working": + self._projection_overwrote = True + self.projection_saved.set() + if self._projection_overwrote: + await self._record_read + + async def get_task_record(self, task_id: str): + current = self._store._tasks.get(task_id) + if ( + task_id == self._task_id + and self.agent_completed.is_set() + and current is not None + and current.state == "input-required" + and current.active_task is None + and not self.final_record_read_waiting.is_set() + ): + self.final_record_read_waiting.set() + await self.projection_waiting.wait() + await self.projection_saved.wait() + record = await self._get_task_record(task_id) + if self._projection_overwrote and task_id == self._task_id and not self._record_read.done(): + self._record_read.set_result(None) + return record + + async def collect(self, stream) -> None: async for event in stream: - remaining_events.append(event) + self.remaining_events.append(event) - await asyncio.wait_for(collect_remaining_events(), timeout=1) + def release_projection(self) -> None: + if not self._projection_release.done(): + self._projection_release.set_result(None) - assert isinstance(first_event, Task) - assert first_event.id == result.id - assert "second" in json.dumps([event.__class__.__name__ + str(event) for event in remaining_events]) - assert prompts == ["hello"] - await components.aclose() + def release_waiters(self) -> None: + self.agent_release.set() + self.release_projection() + self.projection_waiting.set() + self.projection_saved.set() + if not self._record_read.done(): + self._record_read.set_result(None) + + +@pytest.mark.asyncio +async def test_subscribe_finalizes_when_newer_working_projection_arrives_after_executor_cleanup( + monkeypatch, tmp_path +) -> None: + runtime = FakeRuntime(agent_loop=None, session_id="session-1") + monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", lambda options: runtime) + components = create_runtime_components(model="qwen3.6-plus", host="127.0.0.1", port=41242) + gate = NaturalCompletionProjectionGate(components.task_store) + gate.install(monkeypatch) + runtime.agent_loop = gate + send_context = ServerCallContext() + subscribe_context = ServerCallContext() + collector = None + + try: + result = await components.handler.on_message_send( + SendMessageRequest( + message=Message( + message_id="msg-1", + role=Role.ROLE_USER, + parts=[Part(text="hello")], + metadata={"iac_code": {"cwd": str(tmp_path)}}, + ), + configuration=SendMessageConfiguration(accepted_output_modes=["text/plain"], return_immediately=True), + ), + send_context, + ) + assert isinstance(result, Task) + gate.target(result.id) + + stream = components.handler.on_subscribe_to_task(SubscribeToTaskRequest(id=result.id), subscribe_context) + first_event = await asyncio.wait_for(anext(stream), timeout=1) + collector = asyncio.create_task(gate.collect(stream)) + gate.agent_release.set() + await asyncio.wait_for(gate.projection_waiting.wait(), timeout=1) + await asyncio.wait_for(gate.final_record_read_waiting.wait(), timeout=1) + gate.release_projection() + await asyncio.wait_for(gate.projection_saved.wait(), timeout=1) + await asyncio.wait_for(collector, timeout=1) + + execution_state = await components.execution_control_service.observe( + context_id=result.context_id, + owner="", + ) + assert isinstance(first_event, Task) + assert "second" in json.dumps([event.__class__.__name__ + str(event) for event in gate.remaining_events]) + assert execution_state["phase"] == "terminated" + assert execution_state["terminationReason"] == "natural_completion" + assert execution_state["releaseReady"] is True + finally: + gate.release_waiters() + if collector is not None and not collector.done(): + collector.cancel() + await asyncio.gather(collector, return_exceptions=True) + await components.aclose() @pytest.mark.asyncio @@ -2932,6 +3114,8 @@ async def fail_enqueue(_job) -> None: enqueue_attempted.set() raise OSError("queue unavailable") + stream = None + collector = None try: result = await components.handler.on_message_send( SendMessageRequest( @@ -2957,23 +3141,47 @@ async def fail_enqueue(_job) -> None: components.push_queue.enqueue = fail_enqueue # type: ignore[union-attr, method-assign] stream = components.handler.on_subscribe_to_task(SubscribeToTaskRequest(id=result.id), call_context) - await asyncio.wait_for(anext(stream), timeout=1) + first_event_ready = asyncio.get_running_loop().create_future() + + async def collect_events() -> None: + try: + async for event in stream: + if not first_event_ready.done(): + first_event_ready.set_result(event) + except BaseException as exc: + if not first_event_ready.done(): + first_event_ready.set_exception(exc) + raise + finally: + if not first_event_ready.done(): + first_event_ready.set_exception(AssertionError("subscription ended before the initial task")) + + collector = asyncio.create_task(collect_events(), name="collect-push-failure-subscription") + first_event = await asyncio.shield(first_event_ready) + assert isinstance(first_event, Task) with caplog.at_level("WARNING", logger="iac_code.a2a.push"): release.set() + await asyncio.shield(collector) - async def collect_remaining_events() -> None: - async for _event in stream: - pass - - await asyncio.wait_for(collect_remaining_events(), timeout=1) - await asyncio.wait_for(loop_completed.wait(), timeout=1) - await asyncio.wait_for(enqueue_attempted.wait(), timeout=1) - + assert loop_completed.is_set() + assert enqueue_attempted.is_set() final_task = await components.handler.on_get_task(GetTaskRequest(id=result.id), call_context) assert final_task.status.state != TaskState.TASK_STATE_FAILED assert "Failed to enqueue A2A push notification for task" in caplog.text finally: - await components.aclose() + release.set() + try: + if collector is not None: + if collector.done(): + await asyncio.gather(collector, return_exceptions=True) + else: + await asyncio.shield(collector) + finally: + try: + if stream is not None: + await stream.aclose() + finally: + await components.aclose() def test_create_app_wires_stateful_server_primitives(monkeypatch, tmp_path) -> None: @@ -3322,10 +3530,12 @@ def test_execution_control_endpoints_pause_query_resume_and_recover(tmp_path) -> with TestClient(app) as client: pause = client.post("/iac-code/execution/pause", json=pause_payload) assert pause.status_code == 202 + assert "inputHandoffReady" not in pause.json() pause_id = pause.json()["pauseId"] state = client.get("/iac-code/execution/state?contextId=ctx-1&executionId=exec-1") assert state.status_code == 200 + assert "inputHandoffReady" not in state.json() assert state.json()["phase"] in {"pause_committing", "paused"} assert state.json()["executionStatus"] == "input-required" assert state.json()["streamAvailable"] is False @@ -3334,6 +3544,7 @@ def test_execution_control_endpoints_pause_query_resume_and_recover(tmp_path) -> assert recovery.status_code == 200 assert recovery.json()["outputText"] == ["finished turn"] assert recovery.json()["task"]["id"] == "task-1" + assert "inputHandoffReady" not in recovery.json()["executionControl"] resumed = client.post( "/iac-code/execution/resume", @@ -3346,6 +3557,7 @@ def test_execution_control_endpoints_pause_query_resume_and_recover(tmp_path) -> }, ) assert resumed.status_code == 202 + assert "inputHandoffReady" not in resumed.json() assert resumed.json()["phase"] == "resuming" stale = client.post( @@ -3378,7 +3590,6 @@ def test_execution_control_endpoints_pause_query_resume_and_recover(tmp_path) -> }, ) assert late_timeout.status_code == 409 - terminated = client.post( "/iac-code/execution/terminate", json={ @@ -3390,4 +3601,96 @@ def test_execution_control_endpoints_pause_query_resume_and_recover(tmp_path) -> }, ) assert terminated.status_code == 202 + assert "inputHandoffReady" not in terminated.json() assert terminated.json()["phase"] == "terminating" + + +def test_execution_state_reads_persisted_terminal_snapshot_after_cold_restart(tmp_path) -> None: + persistence_dir = tmp_path / "a2a" + control = ExecutionController( + context_id="ctx-1", + task_id="task-1", + owner="", + cwd=str(tmp_path), + server_instance_id="instance-before-restart", + persistence_path=persistence_dir / "execution-control" / "ctx-1.json", + backup_service=None, + execution_id="exec-1", + ) + control.phase = "terminated" + control.execution_status = "input-required" + control.stream_available = False + control.termination_reason = "natural_completion" + control.backup = {"status": "shared_committed", "generation": 3, "commitId": "commit-3"} + control.release_ready = True + persisted = control.snapshot() + persisted["owner"] = "" + atomic_write_json(persistence_dir / "execution-control" / "ctx-1.json", persisted) + + app = create_app( + host="127.0.0.1", + port=41242, + token=None, + model="qwen3.6-plus", + persistence_dir=persistence_dir, + ) + + with TestClient(app) as client: + state = client.get("/iac-code/execution/state?contextId=ctx-1&executionId=exec-1") + mutation = client.post( + "/iac-code/execution/terminate", + json={ + "contextId": "ctx-1", + "expectedExecutionId": "exec-1", + "requestId": "cold-mutation", + "connectionEpoch": 1, + "reason": "stop_chat", + }, + ) + + assert state.status_code == 200 + assert state.json()["phase"] == "terminated" + assert state.json()["terminationReason"] == "natural_completion" + assert state.json()["releaseReady"] is True + assert "owner" not in state.json() + assert mutation.status_code == 404 + + +def test_execution_state_hides_persisted_terminal_snapshot_from_wrong_owner(tmp_path) -> None: + persistence_dir = tmp_path / "a2a" + control = ExecutionController( + context_id="ctx-1", + task_id="task-1", + owner="bearer", + cwd=str(tmp_path), + server_instance_id="instance-before-restart", + persistence_path=persistence_dir / "execution-control" / "ctx-1.json", + backup_service=None, + execution_id="exec-1", + ) + control.phase = "terminated" + control.execution_status = "completed" + control.stream_available = False + control.termination_reason = "natural_completion" + control.release_ready = True + persisted = control.snapshot() + persisted["owner"] = "bearer" + atomic_write_json(persistence_dir / "execution-control" / "ctx-1.json", persisted) + + app = create_app( + host="127.0.0.1", + port=41242, + basic_username="alice", + basic_password="pass", + token=None, + model="qwen3.6-plus", + persistence_dir=persistence_dir, + ) + + with TestClient(app) as client: + state = client.get( + "/iac-code/execution/state?contextId=ctx-1&executionId=exec-1", + headers={"Authorization": "Basic " + b64encode(b"alice:pass").decode("ascii")}, + ) + + assert state.status_code == 404 diff --git a/tests/a2a/test_execution_control.py b/tests/a2a/test_execution_control.py index 0e99040c..8b66f1c8 100644 --- a/tests/a2a/test_execution_control.py +++ b/tests/a2a/test_execution_control.py @@ -7,6 +7,7 @@ import pytest +from iac_code.a2a import execution_control as execution_control_module from iac_code.a2a.backup import run_sync_fenced, run_sync_fenced_with_cancel_completion from iac_code.a2a.execution_control import ( ExecutionControlConflictError, @@ -102,6 +103,499 @@ async def worker() -> None: await control.close() +@pytest.mark.asyncio +async def test_natural_business_boundary_self_finalizes_without_terminate_request( + tmp_path: Path, +) -> None: + control = _controller(tmp_path) + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + + completion_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + + assert completion_generation is not None + observed = await control.finalize_natural_completion( + task_id="task-1", + completion_generation=completion_generation, + ) + assert observed["phase"] == "terminated" + assert observed["terminationReason"] == "natural_completion" + assert observed["releaseReady"] is True + snapshot = control.snapshot() + assert snapshot["phase"] == "terminated" + assert snapshot["executionStatus"] == "input-required" + assert snapshot["terminationReason"] == "natural_completion" + assert snapshot["backup"] == {"status": "disabled"} + assert snapshot["connectionEpoch"] == -1 + await control.close() + + +@pytest.mark.asyncio +async def test_new_execution_waits_for_natural_finalized_cleanup(tmp_path: Path) -> None: + service = ExecutionControlService(persistence_root=None, backup_service=None) + control = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + completion_generation = await control.detach_task( + current, + execution_status="completed", + natural_completion=True, + ) + assert completion_generation is not None + + cleanup_entered = asyncio.Event() + release_cleanup = asyncio.Event() + + async def finalized_cleanup() -> None: + cleanup_entered.set() + await release_cleanup.wait() + + finalizing = asyncio.create_task( + service.finalize_natural_completion( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + completion_generation=completion_generation, + finalized_cleanup=finalized_cleanup, + ) + ) + await cleanup_entered.wait() + begin_attempted = asyncio.Event() + + async def begin_next_execution() -> ExecutionController: + begin_attempted.set() + return await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + + beginning = asyncio.create_task(begin_next_execution()) + await begin_attempted.wait() + assert not beginning.done() + + release_cleanup.set() + await finalizing + next_control = await beginning + assert next_control is not control + await next_control.detach_task(beginning, execution_status="completed") + await next_control.close() + + +@pytest.mark.asyncio +async def test_old_response_boundary_cannot_finalize_newer_same_task_turn(tmp_path: Path) -> None: + control = _controller(tmp_path) + current = asyncio.current_task() + assert current is not None + + await control.attach_task(current) + first_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + await control.attach_task(current) + second_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + + stale = await control.finalize_natural_completion( + task_id="task-1", + completion_generation=first_generation, + ) + observed = await control.observe_state() + + assert first_generation != second_generation + assert stale["phase"] == "running" + assert observed["phase"] == "running" + + settled = await control.finalize_natural_completion( + task_id="task-1", + completion_generation=second_generation, + ) + assert settled["phase"] == "terminated" + assert settled["terminationReason"] == "natural_completion" + assert settled["releaseReady"] is True + await control.close() + + +@pytest.mark.asyncio +async def test_explicit_termination_preempts_inflight_natural_cleanup(tmp_path: Path) -> None: + cleanup_started = asyncio.Event() + release_cleanup = asyncio.Event() + reasons: list[str] = [] + + async def cleanup(_context_id: str, _task_id: str, reason: str) -> str: + reasons.append(reason) + if reason == "natural_completion": + cleanup_started.set() + await release_cleanup.wait() + return "input-required" if reason == "natural_completion" else "canceled" + + control = ExecutionController( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + server_instance_id="instance-1", + persistence_path=tmp_path / "control.json", + backup_service=None, + termination_cleanup=cleanup, + execution_id="exec-1", + ) + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + completion_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + natural = asyncio.create_task( + control.finalize_natural_completion( + task_id="task-1", + completion_generation=completion_generation, + ) + ) + await asyncio.wait_for(cleanup_started.wait(), timeout=1) + + claimed = await control.terminate( + execution_id="exec-1", + request_id="request-stop-during-natural", + connection_epoch=1, + reason="stop_chat", + ) + release_cleanup.set() + await natural + await _wait_for_condition(lambda: control.release_ready) + retried = await control.terminate( + execution_id="exec-1", + request_id="request-stop-during-natural", + connection_epoch=1, + reason="stop_chat", + ) + + assert claimed["phase"] == "terminating" + assert control.termination_reason == "stop_chat" + assert control.execution_status == "canceled" + assert retried["terminationReason"] == "stop_chat" + assert reasons[0] == "natural_completion" + assert "stop_chat" in reasons + await control.close() + + +@pytest.mark.asyncio +async def test_explicit_termination_claim_wins_before_natural_finalization(tmp_path: Path) -> None: + control = _controller(tmp_path) + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + completion_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + + claimed = await control.terminate( + execution_id="exec-1", + request_id="request-stop", + connection_epoch=1, + reason="stop_chat", + ) + assert completion_generation is not None + observed = await control.finalize_natural_completion( + task_id="task-1", + completion_generation=completion_generation, + ) + + assert claimed["terminationReason"] == "stop_chat" + assert observed["terminationReason"] == "stop_chat" + await _wait_for_condition(lambda: control.release_ready) + assert control.termination_reason == "stop_chat" + await control.close() + + +@pytest.mark.asyncio +async def test_explicit_termination_after_natural_release_publishes_fresh_backup( + tmp_path: Path, +) -> None: + calls: list[tuple[BackupReason, bool]] = [] + + class BackupService: + def backup_session(self, _cwd, _session_id, *, reason, critical) -> BackupResult: + calls.append((reason, critical)) + generation = len(calls) + return BackupResult( + enabled=True, + generation=generation, + commit_id=f"commit-{generation}", + shared_committed=True, + ) + + control = _controller(tmp_path, backup_service=BackupService()) + control.bind_session("session-1") + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + completion_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + assert completion_generation is not None + await control.finalize_natural_completion( + task_id="task-1", + completion_generation=completion_generation, + ) + + claimed = await control.terminate( + execution_id="exec-1", + request_id="request-stop-after-natural", + connection_epoch=1, + reason="stop_chat", + ) + + assert claimed["phase"] == "terminating" + assert claimed["terminationReason"] == "stop_chat" + assert claimed["releaseReady"] is False + await _wait_for_condition(lambda: control.release_ready) + assert calls == [ + (BackupReason.TERMINAL, True), + (BackupReason.TERMINAL, True), + ] + assert control.termination_reason == "stop_chat" + assert control.backup["generation"] == 2 + await control.close() + + +@pytest.mark.asyncio +async def test_explicit_termination_fences_inflight_natural_backup(tmp_path: Path) -> None: + first_backup_started = threading.Event() + release_first_backup = threading.Event() + calls: list[str] = [] + + class BackupService: + def backup_session(self, _cwd, _session_id, *, reason, critical) -> BackupResult: + assert reason is BackupReason.TERMINAL + assert critical is True + calls.append(control.termination_reason or "none") + generation = len(calls) + if generation == 1: + first_backup_started.set() + assert release_first_backup.wait(timeout=2) + return BackupResult( + enabled=True, + generation=generation, + commit_id=f"commit-{generation}", + shared_committed=True, + ) + + control = _controller(tmp_path, backup_service=BackupService()) + control.bind_session("session-1") + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + completion_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + assert completion_generation is not None + natural = asyncio.create_task( + control.finalize_natural_completion( + task_id="task-1", + completion_generation=completion_generation, + ) + ) + assert await asyncio.to_thread(first_backup_started.wait, 1) + + await control.terminate( + execution_id="exec-1", + request_id="request-stop-during-backup", + connection_epoch=1, + reason="stop_chat", + ) + await asyncio.sleep(0.01) + assert calls == ["natural_completion"] + + release_first_backup.set() + await natural + await _wait_for_condition(lambda: control.release_ready) + + assert calls == ["natural_completion", "stop_chat"] + assert control.termination_reason == "stop_chat" + assert control.backup["generation"] == 2 + await control.close() + + +@pytest.mark.asyncio +async def test_initial_input_wait_exposes_only_local_continuation_readiness(tmp_path: Path) -> None: + control = _controller(tmp_path) + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + await control.detach_task(current, execution_status="input-required") + + snapshot = control.snapshot() + + assert snapshot["localInputContinuationReady"] is True + assert snapshot["inputHandoffReady"] is False + assert "localInputContinuationReady" not in control.protocol_snapshot(snapshot) + await control.close() + + +@pytest.mark.asyncio +async def test_natural_business_boundary_requires_critical_shared_backup(tmp_path: Path) -> None: + calls: list[tuple[BackupReason, bool]] = [] + + class BackupService: + def backup_session(self, _cwd, _session_id, *, reason, critical) -> BackupResult: + calls.append((reason, critical)) + return BackupResult( + enabled=True, + generation=7, + commit_id="commit-7", + shared_committed=True, + ) + + control = _controller(tmp_path, backup_service=BackupService()) + control.bind_session("session-1") + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + + completion_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + + assert completion_generation is not None + await control.finalize_natural_completion( + task_id="task-1", + completion_generation=completion_generation, + ) + await _wait_for_condition(lambda: control.release_ready) + assert calls == [(BackupReason.TERMINAL, True)] + assert control.release_ready is True + assert control.backup["status"] == "shared_committed" + await control.close() + + +@pytest.mark.asyncio +async def test_natural_business_boundary_waits_for_disconnect_pause_resume_without_cancel( + tmp_path: Path, +) -> None: + control = _controller(tmp_path) + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + paused = await control.pause( + task_id="task-1", + expected_execution_id="exec-1", + request_id="request-pause", + connection_epoch=1, + reason="client_disconnected", + reconnect_timeout_seconds=10, + ) + assert paused["phase"] == "pausing" + + completion_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + + assert current.cancelled() is False + await _wait_for_condition(lambda: control.phase == "paused") + assert control.termination_reason is None + assert control.pause_id == paused["pauseId"] + assert control.release_ready is False + assert completion_generation is not None + await control.finalize_natural_completion( + task_id="task-1", + completion_generation=completion_generation, + ) + + resumed = await control.resume( + execution_id="exec-1", + pause_id=paused["pauseId"], + request_id="request-resume", + connection_epoch=2, + ) + + assert resumed["phase"] == "resuming" + await _wait_for_condition(lambda: control.phase == "running") + await control.observe_state() + await _wait_for_condition(lambda: control.release_ready) + assert control.phase == "terminated" + assert control.termination_reason == "natural_completion" + assert control.pause_id is None + assert control.release_ready is True + await control.close() + + +@pytest.mark.asyncio +async def test_natural_completion_state_observation_retries_failed_cleanup( + tmp_path: Path, +) -> None: + cleanup_attempts = 0 + + async def cleanup(_context_id: str, _task_id: str, _reason: str) -> str: + nonlocal cleanup_attempts + cleanup_attempts += 1 + if cleanup_attempts == 1: + raise OSError("injected cleanup failure") + return "input-required" + + control = ExecutionController( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + server_instance_id="instance-1", + persistence_path=tmp_path / "control.json", + backup_service=None, + termination_cleanup=cleanup, + execution_id="exec-1", + ) + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + completion_generation = await control.detach_task( + current, + execution_status="input-required", + natural_completion=True, + ) + assert completion_generation is not None + await control.finalize_natural_completion( + task_id="task-1", + completion_generation=completion_generation, + ) + await _wait_for_condition(lambda: control.phase == "terminated" and control.backup["status"] == "blocked") + + observed = await control.observe_state() + + assert observed["phase"] == "terminated" + await _wait_for_condition(lambda: control.release_ready) + assert cleanup_attempts == 2 + assert control.execution_status == "input-required" + await control.close() + + @pytest.mark.asyncio async def test_resume_is_ordered_after_inflight_paused_commit(tmp_path: Path, monkeypatch) -> None: control = _controller(tmp_path) @@ -1069,6 +1563,701 @@ async def test_new_normal_turn_gets_new_execution_identity_but_pipeline_continua assert next_normal_turn.execution_id != first.execution_id assert all(task.done() for task in first._background_tasks) await next_normal_turn.detach_task(current, execution_status="input-required") + + third_normal_turn = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + assert third_normal_turn is not next_normal_turn + assert third_normal_turn.execution_id != next_normal_turn.execution_id + await third_normal_turn.detach_task(current, execution_status="input-required") + await service.close() + + +@pytest.mark.asyncio +async def test_pipeline_continuation_replaces_drained_terminated_control_when_backup_is_blocked( + tmp_path: Path, +) -> None: + service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + first = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await first.detach_task(current, execution_status="input-required") + first.phase = "terminated" + first.execution_status = "canceled" + first.backup = {"status": "blocked", "error": "shared backup unavailable"} + first.release_ready = False + + with pytest.raises(ExecutionControlConflictError): + await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + ) + + admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert admission is not None + continuation = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=admission, + ) + + assert continuation is not first + assert continuation.execution_id != first.execution_id + assert all(task.done() for task in first._background_tasks) + await continuation.detach_task(current, execution_status="input-required") + await service.close() + + +@pytest.mark.asyncio +async def test_recoverable_input_admission_rejects_mismatched_task_and_live_background_work(tmp_path: Path) -> None: + service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + control = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await control.detach_task(current, execution_status="input-required") + control.phase = "terminated" + control.execution_status = "canceled" + control.backup = {"status": "blocked", "error": "shared backup unavailable"} + control.release_ready = False + + assert ( + await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-other", + owner="owner-1", + ) + is None + ) + + background_release = asyncio.Event() + background = control._spawn(background_release.wait(), "test-recovery-admission") + assert ( + await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + is None + ) + background_release.set() + await background + await service.close() + + +@pytest.mark.asyncio +async def test_recoverable_input_wait_wakes_after_last_background_task_finishes(tmp_path: Path) -> None: + service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + control = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await control.detach_task(current, execution_status="input-required") + control.phase = "terminated" + control.execution_status = "canceled" + control.backup = {"status": "blocked", "error": "shared backup unavailable"} + control.release_ready = False + background_release = asyncio.Event() + background = control._spawn(background_release.wait(), "test-recovery-wait") + + waiter = asyncio.create_task( + service.wait_until_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + timeout=1, + ) + ) + await asyncio.sleep(0) + assert not waiter.done() + + background_release.set() + await background + await asyncio.wait_for(waiter, timeout=1) + await service.close() + + +@pytest.mark.asyncio +async def test_recoverable_input_admission_is_single_owner_across_service_instances(tmp_path: Path) -> None: + first_service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + second_service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + + admissions = await asyncio.gather( + first_service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ), + second_service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ), + ) + + assert sum(admission is not None for admission in admissions) == 1 + winner = first_service if admissions[0] is not None else second_service + loser = second_service if winner is first_service else first_service + admission = admissions[0] or admissions[1] + assert admission is not None + + with pytest.raises(ExecutionControlConflictError, match="active in another process"): + await loser.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + ) + + control = await winner.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=admission, + ) + assert ( + await loser.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + is None + ) + with pytest.raises(ExecutionControlConflictError, match="active in another process"): + await loser.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + ) + current = asyncio.current_task() + assert current is not None + await control.detach_task(current, execution_status="input-required") + await first_service.close() + await second_service.close() + + +@pytest.mark.asyncio +async def test_stale_local_terminated_control_cannot_replace_new_shared_running_control(tmp_path: Path) -> None: + stale_service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + recovering_service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + stale_control = await stale_service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await stale_control.detach_task(current, execution_status="input-required") + stale_control.phase = "terminated" + stale_control.execution_status = "canceled" + stale_control.backup = {"status": "blocked", "error": "shared backup unavailable"} + stale_control.release_ready = False + stale_control.revision += 1 + await stale_control._persist_snapshot(stale_control.snapshot()) + + admission = await recovering_service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert admission is not None + recovered = await recovering_service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=admission, + ) + + assert ( + await stale_service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + is None + ) + persisted = json.loads((tmp_path / "execution-control" / "ctx-1.json").read_text(encoding="utf-8")) + assert persisted["phase"] == "running" + assert persisted["executionId"] == recovered.execution_id + assert stale_control.execution_id != recovered.execution_id + + await recovered.detach_task(current, execution_status="input-required") + await stale_service.close() + await recovering_service.close() + + +@pytest.mark.asyncio +async def test_recovered_input_wait_can_handoff_to_another_service_instance(tmp_path: Path) -> None: + source_service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + first_recovery_service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + next_recovery_service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + source = await source_service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await source.detach_task(current, execution_status="input-required") + source.phase = "terminated" + source.execution_status = "canceled" + source.backup = {"status": "blocked", "error": "shared backup unavailable"} + source.release_ready = False + source.revision += 1 + await source._persist_snapshot(source.snapshot()) + + first_admission = await first_recovery_service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert first_admission is not None + first_recovery = await first_recovery_service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=first_admission, + ) + await first_recovery.detach_task(current, execution_status="input-required") + handoff_snapshot = json.loads((tmp_path / "execution-control" / "ctx-1.json").read_text(encoding="utf-8")) + assert handoff_snapshot["inputHandoffReady"] is True + assert handoff_snapshot["executionStatus"] == "input-required" + assert handoff_snapshot["streamAvailable"] is False + + next_admission = await next_recovery_service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert next_admission is not None + next_recovery = await next_recovery_service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=next_admission, + ) + assert next_recovery.execution_id != first_recovery.execution_id + persisted = json.loads((tmp_path / "execution-control" / "ctx-1.json").read_text(encoding="utf-8")) + assert persisted["phase"] == "running" + assert persisted["executionId"] == next_recovery.execution_id + assert persisted["inputHandoffReady"] is False + + await next_recovery.detach_task(current, execution_status="input-required") + await source_service.close() + await first_recovery_service.close() + await next_recovery_service.close() + + +@pytest.mark.asyncio +async def test_failed_input_handoff_commit_requires_admission_and_retries_on_reserve( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + source = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await source.detach_task(current, execution_status="input-required") + source.phase = "terminated" + source.execution_status = "canceled" + source.backup = {"status": "blocked", "error": "shared backup unavailable"} + source.release_ready = False + source.revision += 1 + await source._persist_snapshot(source.snapshot()) + admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert admission is not None + recovered = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=admission, + ) + original_atomic_write_json = execution_control_module.atomic_write_json + + def fail_input_handoff(path: Path, value: dict) -> None: + if value.get("inputHandoffReady") is True: + raise OSError("input handoff write failed") + original_atomic_write_json(path, value) + + monkeypatch.setattr(execution_control_module, "atomic_write_json", fail_input_handoff) + with pytest.raises(OSError, match="input handoff write failed"): + await recovered.detach_task(current, execution_status="input-required") + assert recovered.input_handoff_ready() + persisted_path = tmp_path / "execution-control" / "ctx-1.json" + assert json.loads(persisted_path.read_text(encoding="utf-8"))["inputHandoffReady"] is False + + with pytest.raises(ExecutionControlConflictError, match="active in another process"): + await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + ) + + monkeypatch.setattr(execution_control_module, "atomic_write_json", original_atomic_write_json) + retry_admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert retry_admission is not None + assert json.loads(persisted_path.read_text(encoding="utf-8"))["inputHandoffReady"] is True + retried = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=retry_admission, + ) + assert retried is not recovered + + await retried.detach_task(current, execution_status="input-required") + await service.close() + + +@pytest.mark.asyncio +async def test_in_memory_recovered_input_handoff_still_requires_admission() -> None: + service = ExecutionControlService(persistence_root=None, backup_service=None) + source = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd="/tmp", + ) + current = asyncio.current_task() + assert current is not None + await source.detach_task(current, execution_status="input-required") + source.phase = "terminated" + source.execution_status = "canceled" + source.backup = {"status": "blocked", "error": "shared backup unavailable"} + source.release_ready = False + admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert admission is not None + recovered = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd="/tmp", + continue_input_required=True, + recoverable_input_admission=admission, + ) + await recovered.detach_task(current, execution_status="input-required") + assert recovered.input_handoff_ready() + + with pytest.raises(ExecutionControlConflictError, match="active in another process"): + await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd="/tmp", + continue_input_required=True, + ) + + retry_admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert retry_admission is not None + retried = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd="/tmp", + continue_input_required=True, + recoverable_input_admission=retry_admission, + ) + assert retried is not recovered + + await retried.detach_task(current, execution_status="input-required") + await service.close() + + +@pytest.mark.asyncio +async def test_expired_recovery_admission_does_not_mutate_local_or_shared_control( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + original = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await original.detach_task(current, execution_status="input-required") + original.phase = "terminated" + original.execution_status = "canceled" + original.backup = {"status": "blocked", "error": "shared backup unavailable"} + original.release_ready = False + original.revision += 1 + await original._persist_snapshot(original.snapshot()) + control_path = tmp_path / "execution-control" / "ctx-1.json" + original_document = control_path.read_text(encoding="utf-8") + + admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert admission is not None + admission_path = tmp_path / "execution-control" / ".ctx-1.recoverable-input.json" + expires_at = json.loads(admission_path.read_text(encoding="utf-8"))["expiresAt"] + monkeypatch.setattr("iac_code.a2a.execution_control.time.time", lambda: expires_at + 1) + + with pytest.raises(ExecutionControlConflictError, match="admission is stale"): + await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=admission, + ) + + assert service.get_for_context("ctx-1") is original + assert not original.has_managed_work() + assert control_path.read_text(encoding="utf-8") == original_document + await service.release_recoverable_input_continuation(admission) + await service.close() + + +@pytest.mark.asyncio +async def test_recovery_activation_persistence_failure_keeps_original_control( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + original = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await original.detach_task(current, execution_status="input-required") + original.phase = "terminated" + original.execution_status = "canceled" + original.backup = {"status": "blocked", "error": "shared backup unavailable"} + original.release_ready = False + original.revision += 1 + await original._persist_snapshot(original.snapshot()) + control_path = tmp_path / "execution-control" / "ctx-1.json" + original_document = control_path.read_text(encoding="utf-8") + admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert admission is not None + original_atomic_write_json = execution_control_module.atomic_write_json + + def fail_running_snapshot(path: Path, value: dict) -> None: + if path == control_path and value.get("phase") == "running": + raise OSError("running snapshot write failed") + original_atomic_write_json(path, value) + + monkeypatch.setattr(execution_control_module, "atomic_write_json", fail_running_snapshot) + + with pytest.raises(OSError, match="running snapshot write failed"): + await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=admission, + ) + + assert service.get_for_context("ctx-1") is original + assert not original.has_managed_work() + assert control_path.read_text(encoding="utf-8") == original_document + admission_path = tmp_path / "execution-control" / ".ctx-1.recoverable-input.json" + assert admission_path.exists() + await service.release_recoverable_input_continuation(admission) + await service.close() + + +@pytest.mark.asyncio +async def test_cancelled_recovery_activation_rolls_back_shared_and_local_control( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + original = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await original.detach_task(current, execution_status="input-required") + original.phase = "terminated" + original.execution_status = "canceled" + original.backup = {"status": "blocked", "error": "shared backup unavailable"} + original.release_ready = False + original.revision += 1 + await original._persist_snapshot(original.snapshot()) + control_path = tmp_path / "execution-control" / "ctx-1.json" + original_document = control_path.read_text(encoding="utf-8") + admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert admission is not None + activated = threading.Event() + allow_return = threading.Event() + original_activate = service._recoverable_input_admissions.activate + + def activate_then_block(*args, **kwargs): + result = original_activate(*args, **kwargs) + activated.set() + assert allow_return.wait(timeout=5) + return result + + monkeypatch.setattr(service._recoverable_input_admissions, "activate", activate_then_block) + begin_task = asyncio.create_task( + service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=admission, + ) + ) + assert await asyncio.to_thread(activated.wait, 5) + begin_task.cancel() + allow_return.set() + + with pytest.raises(asyncio.CancelledError): + await begin_task + + assert service.get_for_context("ctx-1") is original + assert control_path.read_text(encoding="utf-8") == original_document + admission_path = tmp_path / "execution-control" / ".ctx-1.recoverable-input.json" + assert admission_path.exists() + await service.release_recoverable_input_continuation(admission) + await service.close() + + +@pytest.mark.asyncio +async def test_admission_unlink_failure_keeps_new_control_consistent_and_release_retries( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + service = ExecutionControlService(persistence_root=tmp_path, backup_service=None) + original = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + ) + current = asyncio.current_task() + assert current is not None + await original.detach_task(current, execution_status="input-required") + original.phase = "terminated" + original.execution_status = "canceled" + original.backup = {"status": "blocked", "error": "shared backup unavailable"} + original.release_ready = False + original.revision += 1 + await original._persist_snapshot(original.snapshot()) + admission = await service.reserve_recoverable_input_continuation( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + ) + assert admission is not None + admission_path = tmp_path / "execution-control" / ".ctx-1.recoverable-input.json" + original_unlink = Path.unlink + failed_once = False + + def fail_admission_unlink_once(path: Path, *args, **kwargs) -> None: + nonlocal failed_once + if path == admission_path and not failed_once: + failed_once = True + raise OSError("admission unlink failed") + original_unlink(path, *args, **kwargs) + + monkeypatch.setattr(Path, "unlink", fail_admission_unlink_once) + recovered = await service.begin_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + cwd=str(tmp_path), + continue_input_required=True, + recoverable_input_admission=admission, + ) + + persisted = json.loads((tmp_path / "execution-control" / "ctx-1.json").read_text(encoding="utf-8")) + assert service.get_for_context("ctx-1") is recovered + assert persisted["phase"] == "running" + assert persisted["executionId"] == recovered.execution_id + assert admission_path.exists() + await service.release_recoverable_input_continuation(admission) + assert not admission_path.exists() + + await recovered.detach_task(current, execution_status="input-required") await service.close() diff --git a/tests/a2a/test_execution_control_regressions.py b/tests/a2a/test_execution_control_regressions.py index eb77aa04..c6f126d6 100644 --- a/tests/a2a/test_execution_control_regressions.py +++ b/tests/a2a/test_execution_control_regressions.py @@ -2,6 +2,7 @@ import asyncio import json +import shutil import threading from concurrent.futures import ThreadPoolExecutor from contextlib import suppress @@ -10,10 +11,12 @@ import pytest from a2a.types import Task, TaskState, TaskStatus +from a2a.utils.errors import InvalidParamsError from starlette.testclient import TestClient from iac_code.a2a.app import create_app from iac_code.a2a.execution_control import ( + NATURAL_COMPLETION_FINALIZED_GENERATION, ExecutionController, ExecutionControlService, bind_execution_control, @@ -21,6 +24,7 @@ reset_execution_control, ) from iac_code.a2a.executor import IacCodeA2AExecutor +from iac_code.a2a.input_required import PermissionInputRegistry from iac_code.a2a.persistence import A2APersistenceStore from iac_code.a2a.pipeline_executor import _cancel_task_safely, _drive_stream_events from iac_code.a2a.task_store import A2ATaskStore @@ -31,9 +35,9 @@ from iac_code.tools.base import ToolContext, ToolRegistry from iac_code.tools.cloud.aliyun.ros_stack import RosStack from iac_code.tools.cloud.aliyun.ros_stack_instances import RosStackInstances -from iac_code.types.stream_events import MessageEndEvent, TextDeltaEvent, Usage +from iac_code.types.stream_events import MessageEndEvent, PermissionRequestEvent, TextDeltaEvent, Usage -from .fakes import FakeAgentLoop, FakeEventQueue, FakeRequestContext, FakeRuntime +from .fakes import FakeAgentLoop, FakeEventQueue, FakeRequestContext, FakeRuntime, pending_future @pytest.fixture(autouse=True) @@ -50,6 +54,306 @@ async def wait(): await asyncio.wait_for(wait(), timeout) +class _ObservedBackupGate: + def __init__(self, backup_session, loop): + self._backup_session = backup_session + self._loop = loop + self.started = asyncio.Event() + self.release = threading.Event() + self.terminal_finished = asyncio.Event() + self.calls = [] + + def backup_session(self, *args, **kwargs): + reason = kwargs.get("reason") + critical = kwargs.get("critical") + self.calls.append((reason, critical)) + if reason == BackupReason.NORMAL_TURN_END: + self._loop.call_soon_threadsafe(self.started.set) + assert self.release.wait(5) + result = self._backup_session(*args, **kwargs) + if reason == BackupReason.TERMINAL: + self._loop.call_soon_threadsafe(self.terminal_finished.set) + return result + + +@pytest.mark.asyncio +async def test_natural_completion_releases_permission_cancel_latch_for_next_turn() -> None: + registry = PermissionInputRegistry() + await registry.cancel_task("task-1", reversible=True) + + class ExecutionControlService: + def set_termination_cleanup(self, _callback) -> None: + return None + + def set_resume_callback(self, _callback) -> None: + return None + + async def finalize_natural_completion(self, **kwargs): + await kwargs["finalized_cleanup"]() + return { + "terminationReason": "natural_completion", + NATURAL_COMPLETION_FINALIZED_GENERATION: 1, + } + + executor = IacCodeA2AExecutor( + task_store=A2ATaskStore(), + model="test", + permission_input_registry=registry, + execution_control_service=ExecutionControlService(), + ) + + await executor.finalize_natural_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + completion_generation=1, + ) + future = pending_future() + pending = await registry.register( + PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "pwd"}, + tool_use_id="tool-1", + response_future=future, + ), + task_id="task-1", + context_id="ctx-1", + resolution_owner=object(), + ) + + assert pending.task_id == "task-1" + assert future.done() is False + + +@pytest.mark.asyncio +async def test_explicit_termination_winning_natural_finalize_keeps_permission_cancel_latch() -> None: + registry = PermissionInputRegistry() + await registry.cancel_task("task-1", reversible=True) + + class ExecutionControlService: + def set_termination_cleanup(self, _callback) -> None: + return None + + def set_resume_callback(self, _callback) -> None: + return None + + async def finalize_natural_completion(self, **_kwargs): + return { + "terminationReason": "explicit_terminate", + NATURAL_COMPLETION_FINALIZED_GENERATION: 1, + } + + executor = IacCodeA2AExecutor( + task_store=A2ATaskStore(), + model="test", + permission_input_registry=registry, + execution_control_service=ExecutionControlService(), + ) + + await executor.finalize_natural_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + completion_generation=1, + ) + future = pending_future() + with pytest.raises(InvalidParamsError, match="cancellation is already in progress"): + await registry.register( + PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "pwd"}, + tool_use_id="tool-1", + response_future=future, + ), + task_id="task-1", + context_id="ctx-1", + resolution_owner=object(), + ) + + +@pytest.mark.asyncio +async def test_concurrent_explicit_cancel_survives_older_natural_finalize() -> None: + registry = PermissionInputRegistry() + await registry.cancel_task("task-1", reversible=True) + finalize_entered = asyncio.Event() + allow_finalize = asyncio.Event() + + class ExecutionControlService: + def set_termination_cleanup(self, _callback) -> None: + return None + + def set_resume_callback(self, _callback) -> None: + return None + + async def finalize_natural_completion(self, **kwargs): + finalize_entered.set() + await allow_finalize.wait() + await kwargs["finalized_cleanup"]() + return { + "terminationReason": "natural_completion", + NATURAL_COMPLETION_FINALIZED_GENERATION: 1, + } + + executor = IacCodeA2AExecutor( + task_store=A2ATaskStore(), + model="test", + permission_input_registry=registry, + execution_control_service=ExecutionControlService(), + ) + finalizing = asyncio.create_task( + executor.finalize_natural_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + completion_generation=1, + ) + ) + await finalize_entered.wait() + await registry.cancel_task("task-1") + allow_finalize.set() + await finalizing + + with pytest.raises(InvalidParamsError, match="cancellation is already in progress"): + await registry.register( + PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "pwd"}, + tool_use_id="tool-1", + response_future=pending_future(), + ), + task_id="task-1", + context_id="ctx-1", + resolution_owner=object(), + ) + + +@pytest.mark.asyncio +async def test_natural_finalize_releases_reversible_close_created_while_control_settles() -> None: + registry = PermissionInputRegistry() + finalize_entered = asyncio.Event() + allow_finalize = asyncio.Event() + + class ExecutionControlService: + def set_termination_cleanup(self, _callback) -> None: + return None + + def set_resume_callback(self, _callback) -> None: + return None + + async def finalize_natural_completion(self, **kwargs): + finalize_entered.set() + await allow_finalize.wait() + await kwargs["finalized_cleanup"]() + return { + "terminationReason": "natural_completion", + NATURAL_COMPLETION_FINALIZED_GENERATION: 1, + } + + executor = IacCodeA2AExecutor( + task_store=A2ATaskStore(), + model="test", + permission_input_registry=registry, + execution_control_service=ExecutionControlService(), + ) + finalizing = asyncio.create_task( + executor.finalize_natural_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + completion_generation=1, + ) + ) + await finalize_entered.wait() + await registry.cancel_task("task-1", reversible=True) + allow_finalize.set() + await finalizing + + pending = await registry.register( + PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "pwd"}, + tool_use_id="tool-1", + response_future=pending_future(), + ), + task_id="task-1", + context_id="ctx-1", + resolution_owner=object(), + ) + assert pending.task_id == "task-1" + + +@pytest.mark.asyncio +async def test_stale_generation_cannot_release_newer_natural_turn_permission_latch(tmp_path) -> None: + control = controller(tmp_path) + current = asyncio.current_task() + assert current is not None + await control.attach_task(current) + first_generation = await control.detach_task( + current, + execution_status="completed", + natural_completion=True, + ) + await control.attach_task(current) + second_generation = await control.detach_task( + current, + execution_status="completed", + natural_completion=True, + ) + assert first_generation is not None + assert second_generation is not None + settled = await control.finalize_natural_completion( + task_id="task-1", + completion_generation=second_generation, + ) + assert settled[NATURAL_COMPLETION_FINALIZED_GENERATION] == second_generation + + registry = PermissionInputRegistry() + await registry.cancel_task("task-1", reversible=True) + + class ExecutionControlService: + def set_termination_cleanup(self, _callback) -> None: + return None + + def set_resume_callback(self, _callback) -> None: + return None + + async def finalize_natural_completion(self, **kwargs): + state = await control.finalize_natural_completion( + task_id=kwargs["task_id"], + completion_generation=kwargs["completion_generation"], + ) + if state.get(NATURAL_COMPLETION_FINALIZED_GENERATION) == kwargs["completion_generation"]: + await kwargs["finalized_cleanup"]() + return state + + executor = IacCodeA2AExecutor( + task_store=A2ATaskStore(), + model="test", + permission_input_registry=registry, + execution_control_service=ExecutionControlService(), + ) + await executor.finalize_natural_execution( + context_id="ctx-1", + task_id="task-1", + owner="owner-1", + completion_generation=first_generation, + ) + + with pytest.raises(InvalidParamsError, match="cancellation is already in progress"): + await registry.register( + PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "pwd"}, + tool_use_id="tool-1", + response_future=pending_future(), + ), + task_id="task-1", + context_id="ctx-1", + resolution_owner=object(), + ) + await control.close() + + def controller(tmp_path, backup=None): return ExecutionController( context_id="ctx-1", @@ -93,6 +397,189 @@ def blocked_write(*args): await store.stop_cleanup_loop() +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("fail_first_task_commit", "runtime_close_mode"), + [(False, "ok"), (True, "ok"), (False, "raises"), (False, "hangs")], +) +async def test_disconnect_timeout_backup_preserves_candidate_selection_for_sandbox_restore( + tmp_path, + monkeypatch, + fail_first_task_commit, + runtime_close_mode, +): + from iac_code.a2a.pipeline_journal import A2APipelineJournal + from iac_code.a2a.pipeline_paths import a2a_pipeline_dir_for_session + from iac_code.a2a.pipeline_snapshot import A2APipelineSnapshotStore, reduce_pipeline_events + + monkeypatch.setenv("IAC_CODE_CONFIG_DIR", str(tmp_path / "config")) + shared = tmp_path / "shared" + monkeypatch.setenv("IAC_CODE_CONFIG_BACKUP_DIR", str(shared)) + backup = SessionBackupService() + service = ExecutionControlService(persistence_root=tmp_path / "a2a", backup_service=backup) + store = A2ATaskStore(persistence=A2APersistenceStore(tmp_path / "a2a"), backup_service=backup) + executor = IacCodeA2AExecutor( + task_store=store, + model="test", + backup_service=backup, + execution_control_service=service, + ) + runtime_close_started = asyncio.Event() + runtime_close_cancelled = asyncio.Event() + lifecycle_order = [] + + async def close_runtime(): + lifecycle_order.append("runtime_close_started") + runtime_close_started.set() + if runtime_close_mode == "raises": + raise RuntimeError("injected runtime close failure") + if runtime_close_mode == "hangs": + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + runtime_close_cancelled.set() + raise + + if runtime_close_mode == "hangs": + monkeypatch.setattr( + "iac_code.a2a.task_store._RUNTIME_CLOSE_TIMEOUT_SECONDS", + 0.01, + raising=False, + ) + + original_persist_snapshots = store._persist_terminated_task_snapshots_strict + + def record_persist_snapshots(*args): + original_persist_snapshots(*args) + lifecycle_order.append("task_committed") + + monkeypatch.setattr(store, "_persist_terminated_task_snapshots_strict", record_persist_snapshots) + original_backup_session = backup.backup_session + + def record_backup_session(*args, **kwargs): + lifecycle_order.append("backup_started") + return original_backup_session(*args, **kwargs) + + monkeypatch.setattr(backup, "backup_session", record_backup_session) + + context = await store.get_or_create_context( + context_id="ctx-1", + cwd=str(tmp_path), + runtime_factory=lambda _session_id: SimpleNamespace(aclose=close_runtime), + ) + context.active_task_id = "task-1" + store.mirror_context(context) + task = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + task.state = "input-required" + store.mirror_task(task) + pipeline_dir = a2a_pipeline_dir_for_session(cwd=str(tmp_path), session_id=context.session_id) + pending_selection = { + "schemaVersion": "1.0", + "extensionUri": "urn:iac-code:a2a:pipeline-events:v1", + "eventId": "evt-selection", + "sequence": 1, + "createdAt": "2026-09-15T10:00:00Z", + "eventType": "input_required", + "scope": "step", + "pipelineRunId": "ctx-1", + "taskId": "task-1", + "contextId": "ctx-1", + "pipelineName": "selling", + "status": "input_required", + "step": {"runId": "step-1", "id": "confirm_and_select", "attempt": 1}, + "input": { + "inputId": "selection-1", + "kind": "candidate_selection", + "prompt": "请选择方案", + "options": [{"name": "方案A", "candidate_index": 0}], + }, + } + journal = A2APipelineJournal(pipeline_dir) + journal.append(pending_selection) + A2APipelineSnapshotStore(pipeline_dir).save(reduce_pipeline_events([pending_selection])) + control = controller(tmp_path, backup) + control.bind_session(context.session_id) + control._termination_cleanup = executor._terminate_detached_execution + service._controls["ctx-1"] = control + store.set_execution_control_provider(service.snapshot_for_context, service.has_active_work) + + try: + await control.pause( + task_id="task-1", + expected_execution_id=control.execution_id, + request_id="pause", + connection_epoch=1, + reason="transport_disconnected", + reconnect_timeout_seconds=60, + ) + await wait_until(lambda: control.phase == "paused") + pause_id = control.pause_id + assert pause_id is not None + original_save_task = store._persistence.save_task + if fail_first_task_commit: + + def fail_task_commit(_snapshot): + raise OSError("injected task commit failure") + + monkeypatch.setattr(store._persistence, "save_task", fail_task_commit) + await control.terminate( + execution_id=control.execution_id, + request_id="disconnect-timeout", + connection_epoch=2, + reason="disconnect_timeout", + pause_id=pause_id, + ) + if fail_first_task_commit: + await wait_until(lambda: control.phase == "terminated" and control.backup["status"] == "blocked") + assert not control.release_ready + assert A2APipelineSnapshotStore(pipeline_dir).load()["status"] == "waiting_input" + monkeypatch.setattr(store._persistence, "save_task", original_save_task) + await control.terminate( + execution_id=control.execution_id, + request_id="disconnect-timeout", + connection_epoch=2, + reason="disconnect_timeout", + pause_id=pause_id, + ) + if runtime_close_mode == "hangs": + # Observe the close timeout itself separately from real snapshot and + # shared-backup I/O; release readiness is not a one-second SLA. + await asyncio.wait_for(runtime_close_started.wait(), 10) + await asyncio.wait_for(runtime_close_cancelled.wait(), 10) + await wait_until(lambda: control.release_ready) + + assert control.phase == "terminated" + assert control.execution_status == "input-required" + assert control.backup["status"] == "shared_committed" + assert runtime_close_started.is_set() + assert lifecycle_order.index("task_committed") < lifecycle_order.index("runtime_close_started") + assert lifecycle_order.index("runtime_close_started") < lifecycle_order.index("backup_started") + assert json.loads(next(shared.rglob("a2a/task.json")).read_text(encoding="utf-8"))["state"] == ( + "input-required" + ) + shared_pipeline_snapshot = json.loads( + next(shared.rglob("a2a/pipeline/a2a-snapshot.json")).read_text(encoding="utf-8") + ) + assert shared_pipeline_snapshot["status"] == "waiting_input" + + session_dir = SessionStorage().session_dir(str(tmp_path), context.session_id) + shutil.rmtree(session_dir) + restored = backup.restore_session(str(tmp_path), context.session_id) + assert restored.restored + restored_pipeline_dir = a2a_pipeline_dir_for_session(cwd=str(tmp_path), session_id=context.session_id) + restored_task = json.loads((session_dir / "a2a" / "task.json").read_text(encoding="utf-8")) + restored_context = json.loads((session_dir / "a2a" / "context.json").read_text(encoding="utf-8")) + assert restored_task["state"] == "input-required" + assert restored_context["active_task_id"] is None + assert A2APipelineSnapshotStore(restored_pipeline_dir).load()["status"] == "waiting_input" + assert all( + event["eventType"] != "pipeline_canceled" for event in A2APipelineJournal(restored_pipeline_dir).read_all() + ) + finally: + await service.close() + await store.stop_cleanup_loop() + + @pytest.mark.asyncio @pytest.mark.parametrize("staged", [False, True]) @pytest.mark.parametrize("mode", ["normal", "pipeline"]) @@ -107,7 +594,8 @@ async def test_terminate_waits_for_bootstrap_and_cleanup_then_backs_up(tmp_path, executor = IacCodeA2AExecutor( task_store=store, model="test", backup_service=backup, execution_control_service=service ) - started, release = threading.Event(), threading.Event() + loop = asyncio.get_running_loop() + started, release = asyncio.Event(), threading.Event() closing, close_release, closed = asyncio.Event(), asyncio.Event(), asyncio.Event() async def close(): @@ -116,8 +604,9 @@ async def close(): closed.set() def factory(options): - started.set() - assert release.wait(5) + loop.call_soon_threadsafe(started.set) + # Safety bound only; the test releases this gate in finally as well. + assert release.wait(20) return FakeRuntime(agent_loop=FakeAgentLoop([]), session_id=options.session_id, aclose=close) monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", factory) @@ -134,7 +623,9 @@ async def publish(): executor.execute(FakeRequestContext(metadata={"iac_code": {"cwd": str(tmp_path)}}), FakeEventQueue()) ) try: - assert await asyncio.to_thread(started.wait, 2) + # Bootstrap performs real persistence before entering the factory. + # Wait for its notification without occupying another executor thread. + await asyncio.wait_for(started.wait(), 10) control = service.get_for_context("ctx-1") assert control.session_id is not None await control.terminate( @@ -145,7 +636,7 @@ async def publish(): assert not control.release_ready assert not execution.done() release.set() - await asyncio.wait_for(closing.wait(), 2) + await asyncio.wait_for(closing.wait(), 10) assert not control.release_ready execution.cancel() # Repeated cancellation must not abandon runtime cleanup. await asyncio.sleep(0.01) @@ -670,16 +1161,8 @@ async def test_finished_normal_turn_keeps_result_during_slow_backup(tmp_path, mo agent_loop=FakeAgentLoop([TextDeltaEvent(text="finished result")]), session_id=options.session_id ), ) - started, release = threading.Event(), threading.Event() - original_backup = backup.backup_session - - def gated_backup(*args, **kwargs): - if kwargs.get("reason") == BackupReason.NORMAL_TURN_END: - started.set() - assert release.wait(5) - return original_backup(*args, **kwargs) - - monkeypatch.setattr(backup, "backup_session", gated_backup) + backup_gate = _ObservedBackupGate(backup.backup_session, asyncio.get_running_loop()) + monkeypatch.setattr(backup, "backup_session", backup_gate.backup_session) publisher = ( asyncio.create_task(publish_staged_backups(SessionBackupStagingWorker(tmp_path / "staging", shared))) if staged @@ -690,7 +1173,7 @@ def gated_backup(*args, **kwargs): executor.execute(FakeRequestContext(metadata={"iac_code": {"cwd": str(tmp_path)}}), queue) ) try: - assert await asyncio.to_thread(started.wait, 2) + await backup_gate.started.wait() record = await store.get_task_record("task-1") assert record.state == "input-required" control = service.get_for_context("ctx-1") @@ -714,8 +1197,8 @@ def gated_backup(*args, **kwargs): ) await wait_until(lambda: control.phase == "terminating") assert not control.release_ready - release.set() - await asyncio.wait_for(execution, 3) + backup_gate.release.set() + await execution expected = "canceled" if termination == "legacy" else "input-required" if termination != "legacy": await wait_until(lambda: control.release_ready) @@ -723,8 +1206,14 @@ def gated_backup(*args, **kwargs): record = await store.get_task_record("task-1") assert record.state == expected assert record.output_text == ["finished result"] - # Legacy cancel only stages its noncritical backup; it has no - # execution-control releaseReady barrier for shared publication. + if termination == "legacy": + assert backup_gate.terminal_finished.is_set() + assert backup_gate.calls == [ + (BackupReason.NORMAL_TURN_END, False), + (BackupReason.TERMINAL, True), + ] + # Legacy cancel stages both backups without an execution-control + # releaseReady barrier for shared publication. await wait_until( lambda: any( json.loads(snapshot.read_text(encoding="utf-8"))["state"] == expected @@ -732,7 +1221,7 @@ def gated_backup(*args, **kwargs): ) ) finally: - release.set() + backup_gate.release.set() execution.cancel() await asyncio.gather(execution, return_exceptions=True) await service.close() diff --git a/tests/a2a/test_executor.py b/tests/a2a/test_executor.py index e4028ce4..92ab09dd 100644 --- a/tests/a2a/test_executor.py +++ b/tests/a2a/test_executor.py @@ -7,11 +7,17 @@ from types import SimpleNamespace import pytest -from a2a.types import Task, TaskState, TaskStatusUpdateEvent +from a2a.server.context import ServerCallContext +from a2a.types import Task, TaskState, TaskStatus, TaskStatusUpdateEvent from a2a.utils.errors import InvalidParamsError from google.protobuf.json_format import MessageToDict from iac_code.a2a.backup import backup_session_async +from iac_code.a2a.execution_control import ( + NaturalCompletionGenerationCarrier, + RecoverableInputAdmissionCarrier, + bind_execution_control, +) from iac_code.a2a.executor import IacCodeA2AExecutor, _normal_handoff_has_backup_ack from iac_code.a2a.exposure import A2AExposureType from iac_code.a2a.input_required import PermissionIdentityValidationError, PermissionResponse @@ -21,6 +27,7 @@ from iac_code.a2a.pipeline_journal import A2APipelineJournal from iac_code.a2a.pipeline_paths import a2a_pipeline_dir_for_session from iac_code.a2a.pipeline_snapshot import A2APipelineSnapshotStore, reduce_pipeline_events +from iac_code.a2a.request_scoped_active_task import PipelineLifecycleEventQueueCarrier from iac_code.a2a.task_store import A2ATaskStore from iac_code.agent.message import ImageBlock, Message, TextBlock from iac_code.commands.registry import CommandRegistry, PromptCommand @@ -49,6 +56,7 @@ from iac_code.services.session_storage import SessionStorage from iac_code.skills.frontmatter import SkillFrontmatter from iac_code.skills.skill_definition import SkillDefinition +from iac_code.tools.cloud.aliyun.ros_client import RosClientFactory from iac_code.types.skill_source import SkillSource from iac_code.types.stream_events import ( MessageEndEvent, @@ -2660,6 +2668,153 @@ async def execute( ) +@pytest.mark.asyncio +async def test_executor_binds_recovered_pipeline_lifecycle_before_delegation( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + monkeypatch.setenv("IAC_CODE_MODE", "pipeline") + observed: dict[str, object] = {} + + class Control: + task_id = "task-1" + + async def checkpoint(self) -> None: + observed["checkpointed"] = True + + async def mark_execution_started(self) -> None: + observed["started"] = True + + async def detach_task( + self, + _task, + *, + execution_status: str, + natural_completion: bool = False, + ) -> None: + observed["detached_as"] = execution_status + observed["natural_completion"] = natural_completion + + class ExecutionControlService: + def set_termination_cleanup(self, _callback) -> None: + return None + + def set_resume_callback(self, _callback) -> None: + return None + + async def begin_execution(self, **kwargs): + observed["begin"] = kwargs + return Control() + + class SpyPipelineExecutor: + def __init__(self, **_kwargs) -> None: + return None + + async def execute(self, *, event_queue, **_kwargs): + observed["events_at_delegation"] = list(event_queue.events) + return True + + monkeypatch.setattr("iac_code.a2a.executor.IacCodeA2APipelineExecutor", SpyPipelineExecutor) + + store = A2ATaskStore(metrics=NoOpA2AMetrics()) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.state = "input-required" + executor = IacCodeA2AExecutor( + task_store=store, + model="qwen3.6-plus", + execution_control_service=ExecutionControlService(), + ) + context = FakeRequestContext( + task_id="task-1", + context_id="ctx-1", + text="确认方案并继续", + metadata={"iac_code": {"cwd": str(tmp_path), "run_mode": "pipeline"}}, + ) + context.current_task = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + RecoverableInputAdmissionCarrier.attach(context, "recovery-1") + PipelineLifecycleEventQueueCarrier.attach(context) + queue = FakeEventQueue() + + await executor.execute(context, queue) + + assert observed["checkpointed"] is True + assert observed["started"] is True + assert observed["begin"]["recoverable_input_admission"] == "recovery-1" + assert observed["natural_completion"] is True + assert PipelineLifecycleEventQueueCarrier.is_bound(context) + events_at_delegation = observed["events_at_delegation"] + assert isinstance(events_at_delegation, list) + assert len(events_at_delegation) == 1 + dumped = dump(events_at_delegation[0]) + assert dumped["taskId"] == "task-1" + assert dumped["contextId"] == "ctx-1" + assert dumped["status"]["state"] == "TASK_STATE_WORKING" + + +@pytest.mark.asyncio +async def test_executor_marks_failed_response_as_natural_completion(monkeypatch, tmp_path: Path) -> None: + observed: dict[str, object] = {} + + class Control: + task_id = "task-1" + + async def detach_task( + self, + _task, + *, + execution_status: str, + natural_completion: bool = False, + ) -> int: + observed["execution_status"] = execution_status + observed["natural_completion"] = natural_completion + return 17 + + class ExecutionControlService: + def set_termination_cleanup(self, _callback) -> None: + return None + + def set_resume_callback(self, _callback) -> None: + return None + + async def finalize_natural_completion(self, **kwargs) -> None: + observed["finalized_generation"] = kwargs["completion_generation"] + + store = A2ATaskStore(metrics=NoOpA2AMetrics()) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + control_service = ExecutionControlService() + executor = IacCodeA2AExecutor( + task_store=store, + model="qwen3.6-plus", + execution_control_service=control_service, + ) + + async def fail_execution(*_args, **_kwargs) -> None: + record.state = "failed" + bind_execution_control(Control()) + + monkeypatch.setattr(executor, "_execute", fail_execution) + context = FakeRequestContext( + task_id="task-1", + context_id="ctx-1", + metadata={"iac_code": {"cwd": str(tmp_path)}}, + ) + context.call_context = ServerCallContext() + NaturalCompletionGenerationCarrier.prepare(context) + assert NaturalCompletionGenerationCarrier.mark_delivered(context.call_context) is None + + await executor.execute(context, FakeEventQueue()) + + assert observed == { + "execution_status": "failed", + "natural_completion": True, + "finalized_generation": 17, + } + assert NaturalCompletionGenerationCarrier.read(context) == 17 + + @pytest.mark.asyncio async def test_executor_hydrates_running_pipeline_task_id_from_sidecar( monkeypatch: pytest.MonkeyPatch, tmp_path: Path @@ -2899,9 +3054,7 @@ async def test_executor_runs_normal_mode_when_iac_code_mode_is_normal( @pytest.mark.asyncio -async def test_normal_mode_ignores_stale_pipeline_name( - monkeypatch: pytest.MonkeyPatch, tmp_path: Path -) -> None: +async def test_normal_mode_ignores_stale_pipeline_name(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: monkeypatch.setenv("IAC_CODE_MODE", "pipeline") loop = FakeAgentLoop([TextDeltaEvent(text="normal")]) runtime = FakeRuntime(agent_loop=loop, session_id="session-1") @@ -4309,9 +4462,7 @@ def test_region_only_metadata_copies_configured_credential_without_mutating_it( ) executor = self._make_executor() - result = executor._resolve_aliyun_credential( - {"iac_code": {"alibaba_cloud_region_id": "cn-beijing"}} - ) + result = executor._resolve_aliyun_credential({"iac_code": {"alibaba_cloud_region_id": "cn-beijing"}}) assert result is not None assert result is not configured @@ -4327,9 +4478,7 @@ def test_region_only_metadata_returns_none_without_configured_credential( monkeypatch.setattr("iac_code.a2a.executor.AliyunCredentials.load", lambda: None) executor = self._make_executor() - result = executor._resolve_aliyun_credential( - {"iac_code": {"alibaba_cloud_region_id": "cn-beijing"}} - ) + result = executor._resolve_aliyun_credential({"iac_code": {"alibaba_cloud_region_id": "cn-beijing"}}) assert result is None @@ -4337,9 +4486,7 @@ def test_region_only_metadata_rejects_invalid_region(self) -> None: executor = self._make_executor() with pytest.raises(InvalidParamsError, match="Unsupported Alibaba Cloud region ID"): - executor._resolve_aliyun_credential( - {"iac_code": {"alibaba_cloud_region_id": "https://example.com"}} - ) + executor._resolve_aliyun_credential({"iac_code": {"alibaba_cloud_region_id": "https://example.com"}}) @pytest.mark.asyncio @@ -5101,7 +5248,9 @@ async def test_suspending_permission_answer_waits_for_owner_then_resumes_once( async def pending_for_response(_response): return pending - async def answer(_response, *, before_delivery=None): + async def answer(_response, *, aliyun_credential=None, before_delivery=None): + assert aliyun_credential.access_key_id == "fresh-sts-id" + assert aliyun_credential.sts_token == "fresh-sts-token" if before_delivery is not None: await before_delivery() return True @@ -5133,7 +5282,18 @@ async def publish(_queue, **kwargs): monkeypatch.setattr(executor, "_publish_status", publish) await executor._execute( - FakeRequestContext(task_id="task-1", context_id="ctx-1"), + FakeRequestContext( + task_id="task-1", + context_id="ctx-1", + metadata={ + "iac_code": { + "alibaba_cloud_access_key_id": "fresh-sts-id", + "alibaba_cloud_access_key_secret": "fresh-sts-secret", + "alibaba_cloud_security_token": "fresh-sts-token", + "alibaba_cloud_region_id": "cn-beijing", + } + }, + ), FakeEventQueue(), context_id="ctx-1", ) @@ -5166,7 +5326,7 @@ async def pending_for_response(_response): calls.append("lookup") return pending - async def answer(_response, *, before_delivery=None): + async def answer(_response, *, aliyun_credential=None, before_delivery=None): del before_delivery calls.append("answer") raise InvalidParamsError("permission boundary has no live owner") @@ -5276,9 +5436,7 @@ async def activate_llm_headers() -> None: ) assert activations == ["activated"] assert restored_storage.exists(cwd, session_id) - assert ( - restored_storage.session_dir(cwd, session_id) / "permission-waits" / f"{boundary_id}.json" - ).is_file() + assert (restored_storage.session_dir(cwd, session_id) / "permission-waits" / f"{boundary_id}.json").is_file() @pytest.mark.asyncio @@ -5445,7 +5603,7 @@ async def test_identity_lookup_failure_keeps_live_permission_pending(monkeypatch async def pending_for_response(_response): return pending - async def answer(_response, *, before_delivery=None): + async def answer(_response, *, aliyun_credential=None, before_delivery=None): del before_delivery raise PermissionIdentityValidationError("InternalError", retryable=True) @@ -5498,7 +5656,7 @@ async def test_rejected_live_permission_does_not_replace_context_llm_headers( async def pending_for_response(_response): return SimpleNamespace() - async def answer(_response, *, before_delivery=None): + async def answer(_response, *, aliyun_credential=None, before_delivery=None): assert before_delivery is not None raise PermissionIdentityValidationError("cloud_execution_identity_changed", retryable=False) @@ -5515,9 +5673,7 @@ async def answer(_response, *, before_delivery=None): FakeEventQueue(), ) - assert await store.resolve_context_llm_headers("ctx-1", None) == { - "Authorization": "Bearer accepted" - } + assert await store.resolve_context_llm_headers("ctx-1", None) == {"Authorization": "Bearer accepted"} @pytest.mark.asyncio @@ -5540,7 +5696,7 @@ async def test_duplicate_live_permission_keeps_first_committed_llm_headers( async def pending_for_response(_response): return pending - async def answer(_response, *, before_delivery=None): + async def answer(_response, *, aliyun_credential=None, before_delivery=None): nonlocal answer_count answer_count += 1 if answer_count == 1: @@ -5574,9 +5730,7 @@ async def claim_continuation(_pending): ) assert answer_count == 2 - assert await store.resolve_context_llm_headers("ctx-1", None) == { - "Authorization": "Bearer first" - } + assert await store.resolve_context_llm_headers("ctx-1", None) == {"Authorization": "Bearer first"} @pytest.mark.asyncio @@ -5598,7 +5752,7 @@ async def test_rejected_sideband_permission_does_not_replace_context_llm_headers async def is_sideband_response(_response): return True - async def answer(_response, *, before_delivery=None): + async def answer(_response, *, aliyun_credential=None, before_delivery=None): assert before_delivery is not None raise InvalidParamsError("permission_resume_invalid: pending permission is not active") @@ -5611,9 +5765,7 @@ async def answer(_response, *, before_delivery=None): metadata={"iac_code": {"llm_headers": {"Authorization": "Bearer rejected"}}}, ) - assert await store.resolve_context_llm_headers("ctx-1", None) == { - "Authorization": "Bearer accepted" - } + assert await store.resolve_context_llm_headers("ctx-1", None) == {"Authorization": "Bearer accepted"} @pytest.mark.asyncio @@ -5686,9 +5838,17 @@ def resolve(self, _boundary_id, **kwargs): ) seen_access_key_ids: list[str | None] = [] + class CapturedRosClient: + def __init__(self, config) -> None: + self.config = config + + monkeypatch.setattr("iac_code.tools.cloud.aliyun.ros_client.RosClient", CapturedRosClient) + def register_cloud_tools(_registry, credentials, _services): credential = credentials.get_provider("aliyun") - seen_access_key_ids.append(credential.access_key_id if credential else None) + client = RosClientFactory.create(credential, region_id="cn-beijing") + seen_access_key_ids.append(client.config.access_key_id) + assert client.config.security_token == "resume-sts-token" monkeypatch.setattr("iac_code.a2a.executor.PermissionWaitCheckpointStore", lambda *_args: CheckpointStore()) monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", lambda _options: runtime) @@ -5738,8 +5898,7 @@ def register_cloud_tools(_registry, credentials, _services): index for index, event in enumerate(queue.events) if isinstance(event, TaskStatusUpdateEvent) - and dump(event).get("metadata", {}).get("iac_code", {}).get("inputReceived", {}).get("decision") - == "allow_once" + and dump(event).get("metadata", {}).get("iac_code", {}).get("inputReceived", {}).get("decision") == "allow_once" ] final_indices = [ index diff --git a/tests/a2a/test_input_required.py b/tests/a2a/test_input_required.py index b7b0454c..3942ab5f 100644 --- a/tests/a2a/test_input_required.py +++ b/tests/a2a/test_input_required.py @@ -3,6 +3,7 @@ import asyncio import json import sys +import threading import pytest from a2a.types import Message, Part, Role, Task, TaskState, TaskStatus @@ -14,6 +15,7 @@ from iac_code.a2a.executor import IacCodeA2AExecutor from iac_code.a2a.input_required import ( PERMISSION_QUERY_PREFIX, + PermissionIdentityValidationError, PermissionInputRegistry, PermissionResponse, parse_permission_response, @@ -35,14 +37,26 @@ PermissionWaitPolicy, build_permission_checkpoint, ) +from iac_code.services.providers.aliyun import AliyunCredential, AliyunCredentials from iac_code.services.session_backup import BackupResult from iac_code.services.session_storage import SessionStorage +from iac_code.tools.cloud.aliyun.ros_client import RosClientFactory from iac_code.types.permissions import PermissionAuditMetadata, PermissionResult from iac_code.types.stream_events import PermissionRequestEvent, SubPipelineStreamEvent from .fakes import FakeEventQueue, pending_future +def _sts_credential(access_key_id: str, token: str) -> AliyunCredential: + return AliyunCredential( + mode="StsToken", + access_key_id=access_key_id, + access_key_secret=f"{access_key_id}-secret", + sts_token=token, + region_id="cn-beijing", + ) + + @pytest.mark.asyncio async def test_cancel_durable_detached_permission_runs_suspend_callback() -> None: registry = PermissionInputRegistry() @@ -83,6 +97,162 @@ async def suspend() -> None: assert await registry.has_pending_task("task-1") is False +@pytest.mark.asyncio +async def test_natural_completion_preserves_normal_input_wait(tmp_path) -> None: + class DisabledBackup: + def initialize_session(self, *_args, **_kwargs): + return None + + def backup_session(self, *_args, **_kwargs): + return BackupResult(enabled=False) + + class Control: + execution_mode = "normal" + + def has_managed_work(self) -> bool: + return False + + class ExecutionControlService: + def set_termination_cleanup(self, _callback) -> None: + return None + + def set_resume_callback(self, _callback) -> None: + return None + + def get_for_context(self, context_id: str): + assert context_id == "ctx-1" + return Control() + + cwd = tmp_path / "workspace" + cwd.mkdir() + store = A2ATaskStore(backup_service=DisabledBackup()) + context = await store.get_or_create_context( + context_id="ctx-1", + cwd=str(cwd), + runtime_factory=lambda _session_id: object(), + ) + context.active_task_id = None + store.mirror_context(context) + task = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + task.state = "input-required" + store.mirror_task(task) + executor = IacCodeA2AExecutor( + task_store=store, + model="fake-model", + backup_service=DisabledBackup(), + execution_control_service=ExecutionControlService(), + ) + + result = await executor._terminate_detached_execution( + "ctx-1", + "task-1", + "natural_completion", + ) + + assert result == "input-required" + assert (await store.get_task_record("task-1")).state == "input-required" + assert (await store.get_context_record("ctx-1")).runtime is None + + +@pytest.mark.asyncio +async def test_natural_completion_closes_completed_runtime_before_release(tmp_path) -> None: + class Runtime: + def __init__(self) -> None: + self.closed = False + + async def aclose(self) -> None: + self.closed = True + + class DisabledBackup: + def initialize_session(self, *_args, **_kwargs): + return None + + def backup_session(self, *_args, **_kwargs): + return BackupResult(enabled=False) + + runtime = Runtime() + cwd = tmp_path / "workspace" + cwd.mkdir() + store = A2ATaskStore(backup_service=DisabledBackup()) + context = await store.get_or_create_context( + context_id="ctx-1", + cwd=str(cwd), + runtime_factory=lambda _session_id: runtime, + ) + context.active_task_id = None + store.mirror_context(context) + task = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + task.state = "completed" + store.mirror_task(task) + executor = IacCodeA2AExecutor( + task_store=store, + model="fake-model", + backup_service=DisabledBackup(), + ) + + result = await executor._terminate_detached_execution( + "ctx-1", + "task-1", + "natural_completion", + ) + + assert result == "completed" + assert runtime.closed is True + assert (await store.get_context_record("ctx-1")).runtime is None + + +@pytest.mark.asyncio +async def test_natural_completion_fails_closed_for_active_permission(tmp_path) -> None: + class DisabledBackup: + def initialize_session(self, *_args, **_kwargs): + return None + + def backup_session(self, *_args, **_kwargs): + return BackupResult(enabled=False) + + cwd = tmp_path / "workspace" + cwd.mkdir() + store = A2ATaskStore(backup_service=DisabledBackup()) + context = await store.get_or_create_context( + context_id="ctx-1", + cwd=str(cwd), + runtime_factory=lambda _session_id: object(), + ) + task = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + task.state = "input-required" + store.mirror_task(task) + registry = PermissionInputRegistry() + pending = await registry.register( + PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "true"}, + tool_use_id="tool-1", + response_future=pending_future(), + ), + task_id="task-1", + context_id="ctx-1", + scope="normal", + ) + pending.continuation = object() + executor = IacCodeA2AExecutor( + task_store=store, + model="fake-model", + permission_input_registry=registry, + backup_service=DisabledBackup(), + ) + + with pytest.raises(RuntimeError, match="active permission wait"): + await executor._terminate_detached_execution( + "ctx-1", + "task-1", + "natural_completion", + ) + + assert await registry.has_pending_task("task-1") is True + assert (await store.get_task_record("task-1")).state == "input-required" + assert context.runtime is not None + + @pytest.mark.asyncio async def test_terminate_detached_pipeline_permission_cancels_sidecar_without_handoff(tmp_path) -> None: from iac_code.a2a.pipeline_paths import a2a_pipeline_dir_for_session @@ -164,6 +334,180 @@ async def suspend() -> None: assert "pipeline_handoff_ready" not in event_types +@pytest.mark.asyncio +@pytest.mark.parametrize( + ( + "reason", + "input_kind", + "input_overrides", + "race_cancel", + "expected_state", + "expected_snapshot_status", + "expected_canceled_events", + ), + [ + ("disconnect_timeout", "candidate_selection", {}, False, "input-required", "waiting_input", 0), + ("client_disconnect_control_timeout", "candidate_selection", {}, False, "input-required", "waiting_input", 0), + ("natural_completion", "candidate_selection", {}, False, "input-required", "waiting_input", 0), + ("disconnect_timeout", "deployment_confirmation", {}, False, "input-required", "waiting_input", 0), + ( + "disconnect_timeout", + "ask_user_question", + {"toolUseId": "ask-1"}, + False, + "input-required", + "waiting_input", + 0, + ), + ( + "disconnect_timeout", + "pipeline_pause_confirmation", + {"paused": True}, + False, + "input-required", + "waiting_input", + 0, + ), + ("disconnect_timeout", "ask_user_question", {}, False, "canceled", "canceled", 1), + ("disconnect_timeout", "pipeline_pause_confirmation", {}, False, "canceled", "canceled", 1), + ( + "disconnect_timeout", + "pipeline_pause_confirmation", + {"paused": False}, + False, + "canceled", + "canceled", + 1, + ), + ("disconnect_timeout", "unknown_input", {}, False, "canceled", "canceled", 1), + ("explicit_terminate", "candidate_selection", {}, False, "canceled", "canceled", 1), + ("disconnect_timeout", "permission", {}, False, "canceled", "canceled", 1), + ("disconnect_timeout", "candidate_selection", {}, True, "canceled", "canceled", 1), + ], +) +async def test_terminate_detached_waiting_input_distinguishes_sandbox_release_from_cancel( + tmp_path, + monkeypatch, + reason, + input_kind, + input_overrides, + race_cancel, + expected_state, + expected_snapshot_status, + expected_canceled_events, +) -> None: + from iac_code.a2a.pipeline_paths import a2a_pipeline_dir_for_session + + class DisabledBackup: + def initialize_session(self, *_args, **_kwargs): + return None + + def backup_session(self, *_args, **_kwargs): + return BackupResult(enabled=False) + + cwd = tmp_path / "workspace" + cwd.mkdir() + store = A2ATaskStore(backup_service=DisabledBackup()) + context = await store.get_or_create_context( + context_id="ctx-1", + cwd=str(cwd), + runtime_factory=lambda _session_id: object(), + ) + context.active_task_id = "task-1" + store.mirror_context(context) + task = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + task.state = "input-required" + store.mirror_task(task) + pipeline_dir = a2a_pipeline_dir_for_session(cwd=str(cwd), session_id=context.session_id) + pending_selection = { + "schemaVersion": "1.0", + "extensionUri": "urn:iac-code:a2a:pipeline-events:v1", + "eventId": "evt-selection", + "sequence": 1, + "createdAt": "2026-09-15T10:00:00Z", + "eventType": "input_required", + "scope": "step", + "pipelineRunId": "ctx-1", + "taskId": "task-1", + "contextId": "ctx-1", + "pipelineName": "selling", + "status": "input_required", + "step": {"runId": "step-1", "id": "confirm_and_select", "attempt": 1}, + "input": { + "inputId": "selection-1", + "kind": input_kind, + "prompt": "请选择方案", + "options": [{"name": "方案A", "candidate_index": 0}], + **input_overrides, + }, + } + journal = A2APipelineJournal(pipeline_dir) + journal.append(pending_selection) + A2APipelineSnapshotStore(pipeline_dir).save(reduce_pipeline_events([pending_selection])) + backup = DisabledBackup() + executor = IacCodeA2AExecutor( + task_store=store, + model="fake-model", + permission_input_registry=PermissionInputRegistry(), + backup_service=backup, + ) + + if race_cancel: + from iac_code.a2a import executor as executor_module + from iac_code.a2a.pipeline_executor import WaitingInputCancelResult, cancel_waiting_input_task_from_sidecar + + entered = threading.Event() + release = threading.Event() + original_check = executor_module.sandbox_release_recoverable_task_id_from_sidecar + + def gated_sidecar_check(**kwargs): + result = original_check(**kwargs) + entered.set() + assert release.wait(3) + return result + + monkeypatch.setattr( + executor_module, + "sandbox_release_recoverable_task_id_from_sidecar", + gated_sidecar_check, + ) + termination = asyncio.create_task(executor._terminate_detached_execution("ctx-1", "task-1", reason)) + try: + assert await asyncio.to_thread(entered.wait, 2) + task_record = await store.get_task_record("task-1") + context_record = await store.get_context_record("ctx-1") + canceled = await asyncio.to_thread( + cancel_waiting_input_task_from_sidecar, + cwd=str(cwd), + session_id=context.session_id, + context_id="ctx-1", + task_id="task-1", + backup_service=backup, + task_store=store, + task_record=task_record, + context_record=context_record, + allow_normal_handoff=False, + ) + assert canceled == WaitingInputCancelResult.CANCELED + assert await store.cancel_inactive_input_required_task(task_id="task-1", context_id="ctx-1") + finally: + release.set() + result = await termination + else: + result = await executor._terminate_detached_execution("ctx-1", "task-1", reason) + + assert result == expected_state + assert (await store.get_task_record("task-1")).state == expected_state + snapshot = A2APipelineSnapshotStore(pipeline_dir).load() + assert snapshot is not None and snapshot["status"] == expected_snapshot_status + committed_canceled_events = [ + event + for event in journal.read_all() + if event["eventType"] == "pipeline_canceled" and event.get("visibility") == "committed" + ] + assert len(committed_canceled_events) == expected_canceled_events + + @pytest.fixture(autouse=True) def _stable_permission_identity(monkeypatch): async def resolve(**_kwargs): @@ -367,28 +711,254 @@ async def test_normal_permission_publishes_input_required_and_resumes_live_futur assert queue.events[-1].status.state == TaskState.TASK_STATE_WORKING +@pytest.mark.asyncio +async def test_permission_answer_refreshes_credential_in_waiting_pipeline_task(monkeypatch) -> None: + registry = PermissionInputRegistry() + old_credential = _sts_credential("old-id", "old-token") + fresh_credential = _sts_credential("fresh-id", "fresh-token") + registered = asyncio.Event() + observed: list[tuple[AliyunCredential | None, object]] = [] + + class CapturedRosClient: + def __init__(self, config) -> None: + self.config = config + + monkeypatch.setattr("iac_code.tools.cloud.aliyun.ros_client.RosClient", CapturedRosClient) + + async def run_pipeline() -> None: + with a2a_request_context(aliyun_credential=old_credential): + future = pending_future() + request = PermissionRequestEvent( + tool_name="ros_stack", + tool_input={"action": "CreateStack"}, + tool_use_id="tool-1", + response_future=future, + ) + pending = await registry.register(request, task_id="task-1", context_id="ctx-1") + registered.pending = pending # type: ignore[attr-defined] + registered.set() + assert await future is True + credential = AliyunCredentials.load() + client = RosClientFactory.create(credential, region_id="cn-beijing") + observed.append((credential, client.config)) + + monkeypatch.setattr("iac_code.a2a.input_required.emit_permission_boundary_audit", lambda *_a, **_k: True) + pipeline_task = asyncio.create_task(run_pipeline()) + await registered.wait() + pending = registered.pending # type: ignore[attr-defined] + await registry.answer( + PermissionResponse( + task_id="task-1", + context_id="ctx-1", + request_task_id="task-1", + input_id=pending.input_id, + tool_use_id="tool-1", + decision="allow_once", + ), + aliyun_credential=fresh_credential, + ) + await pipeline_task + + assert observed[0][0] is old_credential + assert observed[0][1].access_key_id == "fresh-id" + assert observed[0][1].security_token == "fresh-token" + + +@pytest.mark.asyncio +async def test_permission_answer_refreshes_detached_credential_without_cross_session_leak(monkeypatch) -> None: + registry = PermissionInputRegistry() + first_old = _sts_credential("first-old-id", "first-old-token") + second_old = _sts_credential("second-old-id", "second-old-token") + fresh = _sts_credential("first-fresh-id", "first-fresh-token") + + continue_second = asyncio.Event() + second_registered = asyncio.Event() + second_observed: list[object] = [] + + class CapturedRosClient: + def __init__(self, config) -> None: + self.config = config + + monkeypatch.setattr("iac_code.tools.cloud.aliyun.ros_client.RosClient", CapturedRosClient) + + async def register(credential: AliyunCredential, task_id: str, context_id: str, *, wait: bool = False): + with a2a_request_context(aliyun_credential=credential): + pending = await registry.register( + PermissionRequestEvent( + tool_name="ros_stack", + tool_input={"action": "CreateStack"}, + tool_use_id=f"tool-{task_id}", + response_future=pending_future(), + ), + task_id=task_id, + context_id=context_id, + scope="normal", + ) + if wait: + second_registered.set() + await continue_second.wait() + second_observed.append(RosClientFactory.create(AliyunCredentials.load(), "cn-beijing").config) + return pending + + first = await asyncio.create_task(register(first_old, "task-1", "ctx-1")) + second_task = asyncio.create_task(register(second_old, "task-2", "ctx-2", wait=True)) + await second_registered.wait() + observed: list[tuple[AliyunCredential | None, object]] = [] + + async def detached_continuation() -> None: + with a2a_request_context(aliyun_credential=first_old): + credential = AliyunCredentials.load() + observed.append((credential, RosClientFactory.create(credential, "cn-beijing").config)) + + first.continuation = detached_continuation + monkeypatch.setattr("iac_code.a2a.input_required.emit_permission_boundary_audit", lambda *_a, **_k: True) + assert await registry.answer( + PermissionResponse( + task_id="task-1", + context_id="ctx-1", + request_task_id="task-1", + input_id=first.input_id, + tool_use_id="tool-task-1", + decision="allow_once", + ), + aliyun_credential=fresh, + ) + continuation = await registry.claim_continuation(first) + assert continuation is not None + await continuation() + continue_second.set() + await second_task + + assert observed[0][0] is first_old + assert observed[0][1].access_key_id == "first-fresh-id" + assert observed[0][1].security_token == "first-fresh-token" + assert second_observed[0].access_key_id == "second-old-id" + assert second_observed[0].security_token == "second-old-token" + assert second_old.access_key_id == "second-old-id" + assert second_old.sts_token == "second-old-token" + + @pytest.mark.asyncio async def test_permission_mismatch_and_duplicate_reply_fail_closed(monkeypatch) -> None: registry = PermissionInputRegistry() + old_credential = _sts_credential("old-id", "old-token") + fresh_credential = _sts_credential("fresh-id", "fresh-token") + + class CapturedRosClient: + def __init__(self, config) -> None: + self.config = config + + monkeypatch.setattr("iac_code.tools.cloud.aliyun.ros_client.RosClient", CapturedRosClient) request = PermissionRequestEvent( tool_name="bash", tool_input={"cmd": "pwd"}, tool_use_id="tool-1", response_future=pending_future(), ) - pending = await registry.register(request, task_id="task-1", context_id="ctx-1") + with a2a_request_context(aliyun_credential=old_credential): + pending = await registry.register(request, task_id="task-1", context_id="ctx-1") monkeypatch.setattr("iac_code.a2a.input_required.emit_permission_boundary_audit", lambda *_args, **_kwargs: True) wrong = parse_permission_response(_permission_message(input_id=pending.input_id)) assert wrong is not None wrong = type(wrong)(**{**wrong.__dict__, "context_id": "ctx-other"}) with pytest.raises(InvalidParamsError, match="input_response_mismatch"): - await registry.answer(wrong) + await registry.answer(wrong, aliyun_credential=fresh_credential) + assert old_credential.access_key_id == "old-id" parsed = parse_permission_response(_permission_message(decision="deny", input_id=pending.input_id)) assert parsed is not None - assert await registry.answer(parsed) is False + assert await registry.answer(parsed, aliyun_credential=fresh_credential) is False + assert old_credential.access_key_id == "fresh-id" + denied_client = RosClientFactory.create(old_credential, "cn-beijing") + assert denied_client.config.access_key_id == "fresh-id" + assert denied_client.config.security_token == "fresh-token" await registry.complete(pending) with pytest.raises(InvalidParamsError, match="pending permission"): - await registry.answer(parsed) + await registry.answer(parsed, aliyun_credential=_sts_credential("duplicate-id", "duplicate-token")) + assert old_credential.access_key_id == "fresh-id" + + +@pytest.mark.asyncio +async def test_changed_execution_identity_does_not_refresh_waiting_credential(monkeypatch, tmp_path) -> None: + registry = PermissionInputRegistry() + registry.set_permission_wait_coordinator(PermissionWaitCoordinator(PermissionWaitPolicy())) + old_credential = _sts_credential("old-id", "old-token") + request = PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "pwd"}, + tool_use_id="tool-1", + response_future=pending_future(), + ) + with a2a_request_context(aliyun_credential=old_credential): + pending = await registry.register(request, task_id="task-1", context_id="ctx-1", scope="normal") + SessionStorage().ensure_v2_session_dir_for_new_session(str(tmp_path), "session-1") + store = PermissionWaitCheckpointStore(str(tmp_path), "session-1") + record = store.create( + build_permission_checkpoint( + session_id="session-1", + task_id="task-1", + context_id="ctx-1", + input_id=pending.input_id, + tool_use_id="tool-1", + tool_name="bash", + tool_input=request.tool_input, + permission_class="normal", + continuation_frame={ + "assistantMessageRef": "session.jsonl:0", + "assistantMessageDigest": "a" * 64, + "orderedToolUseIds": ["tool-1"], + "currentIndex": 0, + "decisions": [{"toolUseId": "tool-1", "state": "pending", "source": None}], + }, + policy=PermissionWaitPolicy(), + principal_ref="different-principal", + ) + ) + pending.boundary_id = record["boundaryId"] + pending.checkpoint_store = store + registry.activate_durable_boundary(pending, record) + monkeypatch.setattr("iac_code.a2a.input_required.emit_permission_boundary_audit", lambda *_a, **_k: True) + + response = PermissionResponse( + task_id="task-1", + context_id="ctx-1", + request_task_id="task-1", + input_id=pending.input_id, + tool_use_id="tool-1", + decision="allow_once", + ) + with pytest.raises(PermissionIdentityValidationError, match="cloud_execution_identity_changed"): + await registry.answer(response, aliyun_credential=_sts_credential("fresh-id", "fresh-token")) + + assert old_credential.access_key_id == "old-id" + assert request.response_future is not None and not request.response_future.done() + + +@pytest.mark.asyncio +async def test_terminated_permission_does_not_refresh_waiting_credential(monkeypatch) -> None: + registry = PermissionInputRegistry() + old_credential = _sts_credential("old-id", "old-token") + request = PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "pwd"}, + tool_use_id="tool-1", + response_future=pending_future(), + ) + with a2a_request_context(aliyun_credential=old_credential): + pending = await registry.register(request, task_id="task-1", context_id="ctx-1") + monkeypatch.setattr("iac_code.a2a.input_required.emit_permission_boundary_audit", lambda *_a, **_k: True) + await registry.cancel_task("task-1") + response = PermissionResponse( + task_id="task-1", + context_id="ctx-1", + request_task_id="task-1", + input_id=pending.input_id, + tool_use_id="tool-1", + decision="allow_once", + ) + with pytest.raises(InvalidParamsError, match="pending permission"): + await registry.answer(response, aliyun_credential=_sts_credential("fresh-id", "fresh-token")) + + assert old_credential.access_key_id == "old-id" @pytest.mark.asyncio @@ -1329,6 +1899,32 @@ async def test_reversible_terminal_token_cannot_reopen_later_permanent_cancel(mo assert future.result() is False +@pytest.mark.asyncio +async def test_naturally_finished_task_can_register_permission_in_next_turn(monkeypatch) -> None: + registry = PermissionInputRegistry() + await registry.cancel_task("task-1", reversible=True) + closing_tokens = await registry.reversible_closing_tokens("task-1") + assert len(closing_tokens) == 1 + await registry.reopen_task(closing_tokens[0]) + monkeypatch.setattr("iac_code.a2a.input_required.emit_permission_boundary_audit", lambda *_args, **_kwargs: True) + future = pending_future() + + pending = await registry.register( + PermissionRequestEvent( + tool_name="bash", + tool_input={"cmd": "pwd"}, + tool_use_id="tool-1", + response_future=future, + ), + task_id="task-1", + context_id="ctx-1", + resolution_owner=object(), + ) + + assert pending.task_id == "task-1" + assert future.done() is False + + def test_existing_pipeline_inputs_get_unified_projection_without_mutating_legacy_envelope() -> None: envelope = { "eventId": "evt-1", diff --git a/tests/a2a/test_permission_execution_control.py b/tests/a2a/test_permission_execution_control.py index be64dfd4..1363fb38 100644 --- a/tests/a2a/test_permission_execution_control.py +++ b/tests/a2a/test_permission_execution_control.py @@ -226,6 +226,24 @@ def gated_backup(*args, **kwargs): @pytest.mark.asyncio @pytest.mark.parametrize("action", ["resume", "terminate", "approve_resume"]) async def test_normal_permission_resident_timer_keeps_paused_runtime(tmp_path, monkeypatch, action): + timer_started = asyncio.Event() + release_timer = asyncio.Event() + suspend_attempted = asyncio.Event() + + class GatedPermissionWaitCoordinator(PermissionWaitCoordinator): + async def _run_resident_timer(self, owner) -> None: + timer_started.set() + await release_timer.wait() + await super()._run_resident_timer(owner) + + async def suspend_now(self, boundary_id: str) -> bool: + suspend_attempted.set() + return await super().suspend_now(boundary_id) + + monkeypatch.setattr( + "iac_code.services.permission_wait.PermissionWaitCoordinator", + GatedPermissionWaitCoordinator, + ) backup = SessionBackupService() _, store, service = make_executor(tmp_path, backup) executor = IacCodeA2AExecutor( @@ -267,6 +285,7 @@ async def close(): response_task = None try: await executor.execute(FakeRequestContext(metadata={"iac_code": {"cwd": str(tmp_path)}}), FakeEventQueue()) + await timer_started.wait() pending = next(iter(executor._permission_input_registry._pending.values())) control = service.get_for_context("ctx-1") paused = await control.pause( @@ -278,7 +297,16 @@ async def close(): reconnect_timeout_seconds=60, ) await wait_until(lambda: control.phase == "paused") - await asyncio.sleep(1.1) + assert pending.checkpoint_store is not None and pending.boundary_id is not None + pending.checkpoint_store.transaction( + pending.boundary_id, + lambda value: { + **value, + "residentDeadlineAt": format_utc(utc_now() - timedelta(seconds=1)), + }, + ) + release_timer.set() + await suspend_attempted.wait() assert not closed and not future.done() and pending.continuation is not None assert control.phase == "paused" if action == "terminate": @@ -314,6 +342,7 @@ async def close(): await wait_until(lambda: bool(closed)) assert closed == [True] and not ran.is_set() finally: + release_timer.set() if response_task is not None: response_task.cancel() await asyncio.gather(response_task, return_exceptions=True) diff --git a/tests/a2a/test_pipeline_executor.py b/tests/a2a/test_pipeline_executor.py index 037a8604..1b9120a4 100644 --- a/tests/a2a/test_pipeline_executor.py +++ b/tests/a2a/test_pipeline_executor.py @@ -2449,6 +2449,117 @@ async def test_non_retryable_exception_terminal_waits_for_active_interrupt(monke assert await module._register_active_interrupt(runtime) is False +@pytest.mark.asyncio +async def test_direct_route_gate_fence_is_published_only_after_interrupt_registration() -> None: + from iac_code.a2a import pipeline_executor as module + from iac_code.a2a.request_scoped_active_task import DirectPipelineRouteGate, DirectPipelineRouteOutcome + + accepted_runtime = module.A2APipelineRuntime(agent_runtime=_fake_runtime()) + accepted_queue = FakeEventQueue() + accepted_gate = DirectPipelineRouteGate() + + assert await module._register_active_interrupt( + accepted_runtime, + event_queue=accepted_queue, + direct_route_gate=accepted_gate, + ) + assert accepted_gate.outcome is DirectPipelineRouteOutcome.ACTIVE + assert accepted_queue.events == [accepted_gate.marker] + await module._settle_active_interrupt_safely(accepted_runtime) + + terminal_runtime = module.A2APipelineRuntime(agent_runtime=_fake_runtime()) + terminal_runtime.terminal_publication_started = True + terminal_queue = FakeEventQueue() + terminal_gate = DirectPipelineRouteGate() + + assert not await module._register_active_interrupt( + terminal_runtime, + event_queue=terminal_queue, + direct_route_gate=terminal_gate, + ) + assert terminal_gate.outcome is DirectPipelineRouteOutcome.RECOVERY_REQUIRED + assert terminal_queue.events == [] + + +@pytest.mark.asyncio +async def test_direct_route_gate_rebinds_recovered_publisher_to_current_lifecycle_queue() -> None: + from iac_code.a2a import pipeline_executor as module + from iac_code.a2a.request_scoped_active_task import DirectPipelineRouteGate + + stale_queue = FakeEventQueue() + current_queue = FakeEventQueue() + publisher = SimpleNamespace(event_queue=stale_queue) + runtime = module.A2APipelineRuntime( + agent_runtime=_fake_runtime(), + publisher=publisher, + ) + gate = DirectPipelineRouteGate() + + assert await module._register_active_interrupt( + runtime, + event_queue=current_queue, + direct_route_gate=gate, + ) + + assert publisher.event_queue is current_queue + assert current_queue.events == [gate.marker] + assert stale_queue.events == [] + await module._settle_active_interrupt_safely(runtime) + + +@pytest.mark.asyncio +async def test_direct_route_gate_fences_new_pipeline_owner_before_publication( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + from iac_code.a2a.request_scoped_active_task import DirectPipelineRouteGate, DirectPipelineRouteOutcome + + monkeypatch.setenv("IAC_CODE_MODE", "pipeline") + fake_pipeline = FakePipeline( + [ + PipelineEvent( + type=PipelineEventType.PIPELINE_COMPLETED, + step_id=None, + timestamp=1717821601.0, + data={"total_steps": 1}, + ) + ], + session_dir=tmp_path / "sidecar", + ) + monkeypatch.setattr("iac_code.a2a.pipeline_executor.create_pipeline", lambda *args, **kwargs: fake_pipeline) + monkeypatch.setattr("iac_code.a2a.pipeline_executor.create_agent_runtime", lambda options: _fake_runtime()) + + store = A2ATaskStore(metrics=NoOpA2AMetrics()) + executor = IacCodeA2APipelineExecutor( + task_store=store, + model="qwen3.6-plus", + metrics=NoOpA2AMetrics(), + artifact_store=None, + push_notifier=None, + permission_resolver=None, + auto_approve_permissions=False, + thinking_exposure_types=None, + ) + task = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + queue = FakeEventQueue() + gate = DirectPipelineRouteGate() + + await executor.execute( + context=FakeRequestContext(metadata={"iac_code": {"cwd": str(tmp_path)}}), + event_queue=queue, + task=task, + task_id="task-1", + context_id="ctx-1", + cwd=str(tmp_path), + prompt='{"selected_candidate_index": 0}', + direct_route_gate=gate, + ) + + assert gate.outcome is DirectPipelineRouteOutcome.ACTIVE + assert queue.events[0] is gate.marker + assert any(isinstance(event, TaskStatusUpdateEvent) for event in queue.events[1:]) + + @pytest.mark.asyncio async def test_cancel_before_outbound_registration_aborts_and_joins_worker( monkeypatch: pytest.MonkeyPatch, @@ -2763,9 +2874,10 @@ async def no_terminal_handoff(*args, **kwargs): ) ) await asyncio.wait_for(terminal_decision_started.wait(), timeout=1) + retry_queue = FakeEventQueue() interrupt = asyncio.create_task( executor._route_active_pipeline_interrupt( - FakeEventQueue(), + retry_queue, task=task, ctx=ctx, task_id="task-1", @@ -2773,6 +2885,7 @@ async def no_terminal_handoff(*args, **kwargs): cwd=str(tmp_path), pipeline_input=normalize_pipeline_user_input("change course"), preserve_task_record=True, + bind_publisher_event_queue=True, ) ) await asyncio.sleep(0) @@ -2781,6 +2894,9 @@ async def no_terminal_handoff(*args, **kwargs): await consumer assert pipeline.handler_calls == 0 + retry_status = _status_events(retry_queue)[-1]["status"] + assert retry_status["state"] == "TASK_STATE_INPUT_REQUIRED" + assert retry_status["message"]["parts"] == [{"text": RETRY_TEXT}] assert publisher.calls == [ ("single", PipelineEventType.PIPELINE_COMPLETED), ] @@ -4272,11 +4388,13 @@ async def test_terminal_backup_blocked_reopens_permission_registration_after_pen registry = executor._permission_input_registry class ResumedPermissionOwner: - async def resolve_permission(self, pending, response) -> bool: + async def resolve_permission(self, pending, response, *, before_delivery=None) -> bool: await registry.claim(pending, response) approved = response.decision == "allow_once" future = pending.request.response_future assert future is not None + if before_delivery is not None: + await before_delivery() future.set_result(approved) await registry.complete(pending) return approved @@ -8642,6 +8760,48 @@ async def handle_user_interrupt(self, message: str) -> SimpleNamespace: assert task.state == "working" +@pytest.mark.asyncio +async def test_base_lifecycle_active_interrupt_publishes_binding_frame_before_routing( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + from iac_code.a2a import pipeline_executor as module + + stale_queue = FakeEventQueue() + current_queue = FakeEventQueue() + publisher = SimpleNamespace(event_queue=stale_queue) + runtime = module.A2APipelineRuntime( + agent_runtime=_fake_runtime(), + pipeline=object(), + publisher=publisher, + ) + ctx = SimpleNamespace(runtime=runtime) + task = SimpleNamespace(task_id="task-1", context_id="ctx-1", state="input-required") + executor = _pipeline_executor() + route_registered = AsyncMock(return_value=True) + monkeypatch.setattr(executor, "_route_registered_active_pipeline_interrupt", route_registered) + + routed = await executor._route_active_pipeline_interrupt( + current_queue, + task=task, + ctx=ctx, + task_id="task-1", + context_id="ctx-1", + cwd=str(tmp_path), + pipeline_input='{"selected_candidate_index": 0}', + preserve_task_record=True, + bind_publisher_event_queue=True, + ) + + assert routed is True + assert publisher.event_queue is current_queue + assert _status_events(current_queue)[0]["status"]["state"] == "TASK_STATE_WORKING" + assert _status_events(current_queue)[0]["taskId"] == "task-1" + assert _status_events(current_queue)[0]["contextId"] == "ctx-1" + assert stale_queue.events == [] + route_registered.assert_awaited_once() + + @pytest.mark.asyncio async def test_executor_routes_running_sidecar_pending_ask_to_ask_resume( monkeypatch: pytest.MonkeyPatch, @@ -9205,6 +9365,9 @@ async def test_pipeline_executor_resumes_waiting_input_after_stale_active_task_r ) -> None: from iac_code.a2a.pipeline_executor import IacCodeA2APipelineExecutor from iac_code.a2a.pipeline_paths import a2a_pipeline_dir_for_session + from iac_code.a2a.request_scoped_active_task import ( + PipelineLifecycleEventQueueCarrier, + ) cwd = tmp_path / "workspace" cwd.mkdir() @@ -9265,13 +9428,16 @@ async def test_pipeline_executor_resumes_waiting_input_after_stale_active_task_r monkeypatch.setattr(executor, "_create_pipeline", lambda **_kwargs: fake_pipeline) queue = FakeEventQueue() + request_context = FakeRequestContext( + task_id=task_id, + context_id=context_id, + text="0", + metadata={"iac_code": {"cwd": str(cwd)}}, + ) + PipelineLifecycleEventQueueCarrier.attach(request_context) + await executor.execute( - context=FakeRequestContext( - task_id=task_id, - context_id=context_id, - text="0", - metadata={"iac_code": {"cwd": str(cwd)}}, - ), + context=request_context, event_queue=queue, task=task, task_id=task_id, @@ -9281,6 +9447,7 @@ async def test_pipeline_executor_resumes_waiting_input_after_stale_active_task_r ) assert fake_pipeline.resume_prompts == ["0"] + assert _status_events(queue)[0]["status"]["state"] == "TASK_STATE_WORKING" assert task.state == "input-required" assert ctx.active_task_id is None assert all( diff --git a/tests/a2a/test_resource_selector.py b/tests/a2a/test_resource_selector.py index 0647e49e..4b0588d4 100644 --- a/tests/a2a/test_resource_selector.py +++ b/tests/a2a/test_resource_selector.py @@ -25,6 +25,7 @@ from iac_code.agent.message import Message as AgentMessage from iac_code.resource_selector.profiles import PROFILE_HASH, get_profile from iac_code.resource_selector.tools import SelectCloudResourceTool +from iac_code.services.session_backup import BackupReason from iac_code.services.session_storage import SessionStorage from iac_code.tools.base import ToolContext from iac_code.types.stream_events import ( @@ -272,6 +273,17 @@ async def test_live_resource_selection_can_be_answered_while_input_required_is_p event = selection_event(future=asyncio.get_running_loop().create_future()) resumed = asyncio.Event() + class RecordingBackup: + def __init__(self): + self.calls = [] + + async def backup(self, _service, _cwd, _session_id, *, reason, critical, **_kwargs): + self.calls.append((reason, critical)) + return None + + backup = RecordingBackup() + monkeypatch.setattr("iac_code.a2a.executor.backup_session_async", backup.backup) + class ImmediateAnswerLoop: async def run_streaming(self, _prompt): yield event @@ -301,16 +313,30 @@ async def publish_status(target_queue, **kwargs): response=response(), ) ) - await asyncio.sleep(0) + # Observe answer delivery without waiting for the continuation, + # which needs execute() to release the context lock first. + await asyncio.shield(event.response_future) monkeypatch.setattr(executor, "_publish_status", publish_status) - await executor.execute(FakeRequestContext(metadata={"iac_code": {"cwd": str(tmp_path)}}), queue) - assert answer_task is not None - await asyncio.wait_for(answer_task, timeout=1) - - assert resumed.is_set() - assert not await executor._resource_selection_registry.has_pending_task("task-1") + try: + await executor.execute(FakeRequestContext(metadata={"iac_code": {"cwd": str(tmp_path)}}), queue) + assert answer_task is not None + await asyncio.shield(answer_task) + + assert event.response_future.done() + assert resumed.is_set() + assert not await executor._resource_selection_registry.has_pending_task("task-1") + record = await executor._task_store.get_task_record("task-1") + assert "selection accepted" in record.output_text + assert (BackupReason.NORMAL_TURN_END, False) in backup.calls + finally: + if answer_task is not None: + if not answer_task.done(): + answer_task.cancel() + await asyncio.gather(answer_task, return_exceptions=True) + await executor._resource_selection_registry.cancel_task("task-1") + await executor._task_store.stop_cleanup_loop() @pytest.mark.asyncio diff --git a/tests/a2a/test_task_store.py b/tests/a2a/test_task_store.py index e9200d6a..733287b1 100644 --- a/tests/a2a/test_task_store.py +++ b/tests/a2a/test_task_store.py @@ -986,17 +986,75 @@ async def test_late_sdk_working_event_does_not_overwrite_active_executor_finaliz assert task.status.state == TaskState.TASK_STATE_WORKING +@pytest.mark.asyncio +async def test_late_sdk_working_event_does_not_overwrite_running_execution_terminal_boundary(tmp_path) -> None: + persistence = A2APersistenceStore(tmp_path) + store = A2ATaskStore(metrics=NoOpA2AMetrics(), persistence=persistence) + store.set_execution_control_provider( + lambda _context_id: { + "phase": "running", + "taskId": "task-1", + "executionStatus": "input-required", + }, + None, + ) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.state = "input-required" + record.updated_at = 20 + record.active_task = None + store._mirror_task(record) + + await store.save(sdk_task("task-1", state=TaskState.TASK_STATE_WORKING, updated_at=30)) + + assert record.state == "input-required" + assert persistence.load_task("task-1").state == "input-required" + + +@pytest.mark.asyncio +async def test_sdk_working_event_starts_new_execution_after_terminal_boundary(tmp_path) -> None: + persistence = A2APersistenceStore(tmp_path) + store = A2ATaskStore(metrics=NoOpA2AMetrics(), persistence=persistence) + store.set_execution_control_provider( + lambda _context_id: { + "phase": "running", + "taskId": "task-1", + "executionStatus": "working", + }, + None, + ) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.state = "input-required" + record.updated_at = 20 + record.active_task = None + store._mirror_task(record) + + await store.save(sdk_task("task-1", state=TaskState.TASK_STATE_WORKING, updated_at=30)) + + assert record.state == "working" + assert persistence.load_task("task-1").state == "working" + + @pytest.mark.asyncio async def test_stale_sdk_state_does_not_replace_newer_sdk_visible_state() -> None: store = A2ATaskStore(metrics=NoOpA2AMetrics()) await store.save(sdk_task("task-1", state=TaskState.TASK_STATE_WORKING, updated_at=10)) await store.save(sdk_task("task-1", state=TaskState.TASK_STATE_INPUT_REQUIRED, updated_at=20)) - await store.save(sdk_task("task-1", state=TaskState.TASK_STATE_WORKING, updated_at=15)) + delayed = sdk_task("task-1", state=TaskState.TASK_STATE_WORKING, updated_at=15) + delayed.status.message.CopyFrom( + Message( + message_id="delayed-output", + role=Role.ROLE_AGENT, + parts=[Part(text="preserve me for history")], + ) + ) + await store.save(delayed) task = await store.get("task-1") assert task is not None assert task.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + assert delayed.status.state == TaskState.TASK_STATE_WORKING + assert delayed.status.message.parts[0].text == "preserve me for history" @pytest.mark.asyncio @@ -1448,6 +1506,7 @@ async def test_save_attaches_current_execution_control_metadata() -> None: "executionId": "exec-1", "phase": "pausing", "revision": 2, + "inputHandoffReady": True, }, lambda: True, ) @@ -1603,6 +1662,33 @@ async def test_cancel_inactive_input_required_task_updates_internal_and_sdk_stat assert await store.cancel_inactive_input_required_task(task_id="task-1", context_id="ctx-1") is True +@pytest.mark.asyncio +async def test_commit_inactive_execution_task_expected_state_mismatch_preserves_snapshots(tmp_path) -> None: + persistence = A2APersistenceStore(tmp_path / "a2a") + store = A2ATaskStore(metrics=NoOpA2AMetrics(), persistence=persistence) + context = await store.get_or_create_context( + context_id="ctx-1", + cwd=str(tmp_path), + runtime_factory=lambda _session_id: object(), + ) + context.active_task_id = "task-1" + store.mirror_context(context) + await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + await store.save(sdk_task("task-1", state=TaskState.TASK_STATE_INPUT_REQUIRED)) + + committed = await store.commit_inactive_execution_task( + task_id="task-1", + context_id="ctx-1", + expected_state="canceled", + ) + + assert committed is False + assert (await store.get_task_record("task-1")).state == "input-required" + assert (await store.get_context_record("ctx-1")).active_task_id == "task-1" + assert persistence.load_task("task-1").state == "input-required" + assert persistence.load_context("ctx-1").active_task_id == "task-1" + + @pytest.mark.asyncio async def test_cancel_inactive_input_required_task_retries_strict_session_snapshot(tmp_path, monkeypatch) -> None: from iac_code.a2a import task_store as task_store_module diff --git a/tests/a2a/test_transport_dispatcher.py b/tests/a2a/test_transport_dispatcher.py index ea063d18..b618b43b 100644 --- a/tests/a2a/test_transport_dispatcher.py +++ b/tests/a2a/test_transport_dispatcher.py @@ -4,16 +4,40 @@ import json import shutil import threading +import uuid from types import SimpleNamespace import httpx import pytest +from a2a.server.agent_execution import RequestContext +from a2a.server.agent_execution.active_task import _RequestCompleted, _RequestStarted +from a2a.server.agent_execution.active_task_registry import ActiveTaskRegistry from a2a.server.context import ServerCallContext +from a2a.server.events.event_queue_v2 import EventQueueSource, QueueShutDown from a2a.server.request_handlers import DefaultRequestHandler -from a2a.types import Message, Part, Role, SubscribeToTaskRequest, Task, TaskState, TaskStatus, TaskStatusUpdateEvent +from a2a.types import ( + Message, + Part, + Role, + SendMessageRequest, + SubscribeToTaskRequest, + Task, + TaskState, + TaskStatus, + TaskStatusUpdateEvent, +) +from a2a.utils.errors import InvalidParamsError +from google.protobuf.json_format import ParseDict from google.protobuf.struct_pb2 import Value +from iac_code.a2a.execution_control import ( + ExecutionControlService, + NaturalCompletionGenerationCarrier, + RecoverableInputAdmissionCarrier, + RecoverableInputAdmissionLease, +) from iac_code.a2a.input_required import PERMISSION_QUERY_PREFIX +from iac_code.a2a.persistence import A2APersistenceStore from iac_code.a2a.pipeline_journal import A2APipelineJournal from iac_code.a2a.pipeline_paths import a2a_pipeline_dir_for_session from iac_code.a2a.pipeline_snapshot import A2APipelineSnapshotStore, reduce_pipeline_events @@ -25,6 +49,12 @@ pipeline_transport_delivery_tracking_enabled, register_pipeline_transport_delivery, ) +from iac_code.a2a.request_scoped_active_task import ( + DirectPipelineRouteGateCarrier, + PipelineLifecycleEventQueueCarrier, + RequestScopedActiveTask, + RequestScopedActiveTaskRegistry, +) from iac_code.a2a.task_store import A2ATaskStore from iac_code.a2a.transports.dispatcher import ( A2AJsonRpcDispatcher, @@ -38,344 +68,2191 @@ from iac_code.services.session_storage import SessionStorage from iac_code.types.stream_events import PermissionRequestEvent, TextDeltaEvent -from .fakes import FakeAgentLoop, FakeRuntime, pending_future +from .fakes import FakeAgentLoop, FakeEventQueue, FakeRuntime, pending_future _STREAM_TEST_TIMEOUT = 5 @pytest.mark.asyncio -async def test_closed_transport_tracker_reports_closed_stage_on_registration() -> None: - tracker = create_pipeline_transport_delivery_tracker() - close_pipeline_transport_delivery_tracker(tracker) - stages: list[str] = [] +async def test_setup_active_task_attaches_admission_to_the_queued_request_context(monkeypatch) -> None: + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = A2ATaskStore() + call_context = ServerCallContext() + request_context = SimpleNamespace() + handler._stage_recoverable_input_admission(call_context, "recovery-1") - with bind_pipeline_transport_delivery_tracker(tracker): - completion = register_pipeline_transport_delivery( - object(), - stage_observer=lambda stage, _at_ns: stages.append(stage), - ) + async def sdk_setup(_handler, _params, observed_call_context): + assert observed_call_context is call_context + return object(), request_context - with pytest.raises(PipelineTransportDeliveryClosedError): - await completion - assert stages == ["registered", "closed"] + monkeypatch.setattr(DefaultRequestHandler, "_setup_active_task", sdk_setup) + _, result_context = await handler._setup_active_task(object(), call_context) + + assert result_context is request_context + assert RecoverableInputAdmissionCarrier.read(request_context) == "recovery-1" + assert PipelineLifecycleEventQueueCarrier.read(request_context) is True + assert call_context.state == {"iac_code.recoverable_input_admission": "recovery-1"} @pytest.mark.asyncio -async def test_dispatcher_handles_unary_v03_message(monkeypatch, tmp_path) -> None: - loop = FakeAgentLoop([TextDeltaEvent(text="hello from dispatcher")]) +async def test_request_scoped_active_task_ignores_old_terminal_before_its_request_start(monkeypatch) -> None: + request_id = uuid.uuid4() + request_enqueued = asyncio.Event() + call_context = ServerCallContext() + request_context = RequestContext(call_context=call_context, task_id="task-1", context_id="ctx-1") + active_task = RequestScopedActiveTask( + agent_executor=SimpleNamespace(), + task_id="task-1", + task_manager=SimpleNamespace(), + ) - def factory(options): - return FakeRuntime(agent_loop=loop, session_id=options.session_id) + async def enqueue_request(_request_context) -> uuid.UUID: + request_enqueued.set() + return request_id - monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", factory) - components = create_runtime_components(model="qwen3.6-plus", host="127.0.0.1", port=41242) - dispatcher = A2AJsonRpcDispatcher(components) + monkeypatch.setattr(active_task, "enqueue_request", enqueue_request) + stream = active_task.subscribe(request=request_context) + first_event = asyncio.create_task(anext(stream)) + await asyncio.wait_for(request_enqueued.wait(), timeout=_STREAM_TEST_TIMEOUT) - response = await dispatcher.dispatch( - { - "jsonrpc": "2.0", - "id": "1", - "method": "message/send", - "params": { - "message": { - "messageId": "msg-1", - "role": "user", - "parts": [{"kind": "text", "text": "hello"}], - "metadata": {"iac_code": {"cwd": str(tmp_path)}}, - }, - "configuration": {"acceptedOutputModes": ["text/plain"]}, - }, - } + old_terminal = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + current_update = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + canonical_working = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + queued_events = ( + (old_terminal, None), + (_RequestStarted(request_id, request_context), None), + (old_terminal, canonical_working), + (current_update, canonical_working), ) + for queued_event in queued_events: + await active_task._event_queue_subscribers.enqueue_event(queued_event) - assert response["id"] == "1" - assert response["result"]["status"]["state"] == "input-required" - session_id = components.task_store._contexts[response["result"]["contextId"]].session_id - assert response["result"]["metadata"]["iac_code"]["iacCodeSessionId"] == session_id - assert loop.prompts == ["hello"] - await components.aclose() + assert await asyncio.wait_for(first_event, timeout=_STREAM_TEST_TIMEOUT) is current_update + next_event = asyncio.create_task(anext(stream)) + await active_task._event_queue_subscribers.enqueue_event((_RequestCompleted(request_id), None)) + with pytest.raises(StopAsyncIteration): + await asyncio.wait_for(next_event, timeout=_STREAM_TEST_TIMEOUT) + + await active_task._event_queue_agent.close(immediate=True) + await active_task._event_queue_subscribers.close(immediate=True) @pytest.mark.asyncio -async def test_dispatcher_rejects_explicit_invalid_run_mode(tmp_path) -> None: - components = create_runtime_components(model="qwen3.6-plus", host="127.0.0.1", port=41242) - dispatcher = A2AJsonRpcDispatcher(components) +async def test_request_scoped_active_task_isolates_two_concurrent_subscribers(monkeypatch) -> None: + request_ids = [uuid.uuid4(), uuid.uuid4()] + contexts = [ + RequestContext(call_context=ServerCallContext(), task_id="task-1", context_id="ctx-1") for _ in request_ids + ] + both_enqueued = asyncio.Event() + active_task = RequestScopedActiveTask( + agent_executor=SimpleNamespace(), + task_id="task-1", + task_manager=SimpleNamespace(), + ) + enqueued = 0 + + async def enqueue_request(request_context) -> uuid.UUID: + nonlocal enqueued + index = contexts.index(request_context) + enqueued += 1 + if enqueued == 2: + both_enqueued.set() + return request_ids[index] + + monkeypatch.setattr(active_task, "enqueue_request", enqueue_request) + streams = [active_task.subscribe(request=context) for context in contexts] + first_events = [asyncio.create_task(anext(stream)) for stream in streams] + await asyncio.wait_for(both_enqueued.wait(), timeout=_STREAM_TEST_TIMEOUT) + + old_terminal = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + updates = [ + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=state), + ) + for state in (TaskState.TASK_STATE_WORKING, TaskState.TASK_STATE_INPUT_REQUIRED) + ] + events = ( + old_terminal, + _RequestStarted(request_ids[0], contexts[0]), + updates[0], + _RequestCompleted(request_ids[0]), + _RequestStarted(request_ids[1], contexts[1]), + updates[1], + _RequestCompleted(request_ids[1]), + ) + for event in events: + await active_task._event_queue_subscribers.enqueue_event((event, None)) - response = await dispatcher.dispatch( - { - "jsonrpc": "2.0", - "id": "invalid-run-mode", - "method": "message/send", - "params": { - "message": { - "messageId": "msg-invalid-run-mode", - "role": "user", - "parts": [{"kind": "text", "text": "hello"}], - "metadata": {"iac_code": {"cwd": str(tmp_path), "run_mode": "pipline"}}, - }, - "configuration": {"acceptedOutputModes": ["text/plain"]}, - }, - } + assert await asyncio.wait_for(first_events[0], timeout=_STREAM_TEST_TIMEOUT) is updates[0] + assert await asyncio.wait_for(first_events[1], timeout=_STREAM_TEST_TIMEOUT) is updates[1] + for stream in streams: + with pytest.raises(StopAsyncIteration): + await asyncio.wait_for(anext(stream), timeout=_STREAM_TEST_TIMEOUT) + + await active_task._event_queue_agent.close(immediate=True) + await active_task._event_queue_subscribers.close(immediate=True) + + +@pytest.mark.asyncio +async def test_request_enqueue_failure_keeps_admission_owned_by_transport() -> None: + released: list[str] = [] + call_context = ServerCallContext() + call_context.state["iac_code.recoverable_input_admission"] = "recovery-1" + request_context = RequestContext(call_context=call_context, task_id="task-1", context_id="ctx-1") + + async def release(admission: str) -> None: + released.append(admission) + + lease = RecoverableInputAdmissionLease( + "recovery-1", + acknowledge_enqueue=lambda token: call_context.state.pop("iac_code.recoverable_input_admission", None) == token, + release=release, + ) + RecoverableInputAdmissionCarrier.attach(request_context, lease) + active_task = RequestScopedActiveTask( + agent_executor=SimpleNamespace(), + task_id="task-1", + task_manager=SimpleNamespace(), ) + active_task._request_queue.shutdown(immediate=True) - assert response["id"] == "invalid-run-mode" - assert response["error"]["code"] == -32602 - assert response["error"]["message"] == "Unsupported run mode." - await components.aclose() + stream = active_task.subscribe(request=request_context) + with pytest.raises(QueueShutDown): + await anext(stream) + + assert call_context.state == {"iac_code.recoverable_input_admission": "recovery-1"} + assert released == [] + assert active_task._reference_count == 0 + assert active_task._event_queue_subscribers._sinks == set() + await active_task._event_queue_agent.close(immediate=True) + await active_task._event_queue_subscribers.close(immediate=True) @pytest.mark.asyncio -async def test_dispatcher_stream_yields_events(monkeypatch, tmp_path) -> None: - loop = FakeAgentLoop([TextDeltaEvent(text="streamed")]) - runtime = FakeRuntime(agent_loop=loop, session_id="session-1") - monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", lambda options: runtime) - components = create_runtime_components(model="qwen3.6-plus", host="127.0.0.1", port=41242) - dispatcher = A2AJsonRpcDispatcher(components) +async def test_request_scoped_active_task_surfaces_producer_error_before_request_start(monkeypatch) -> None: + request_id = uuid.uuid4() + request_enqueued = asyncio.Event() + request_context = RequestContext( + call_context=ServerCallContext(), + task_id="task-1", + context_id="ctx-1", + ) + active_task = RequestScopedActiveTask( + agent_executor=SimpleNamespace(), + task_id="task-1", + task_manager=SimpleNamespace(), + ) - events = [ - event - async for event in dispatcher.dispatch_stream( - { - "jsonrpc": "2.0", - "id": "2", - "method": "message/stream", - "params": { - "message": { - "messageId": "msg-2", - "role": "user", - "parts": [{"kind": "text", "text": "hello"}], - "metadata": {"iac_code": {"cwd": str(tmp_path)}}, - }, - "configuration": {"acceptedOutputModes": ["text/plain"]}, - }, - } - ) - ] + async def enqueue_request(_request_context) -> uuid.UUID: + request_enqueued.set() + return request_id - assert any(event["result"]["status"]["state"] == "working" for event in events) - assert events[-1]["result"]["status"]["state"] == "input-required" - await components.aclose() + monkeypatch.setattr(active_task, "enqueue_request", enqueue_request) + stream = active_task.subscribe(request=request_context) + result = asyncio.create_task(anext(stream)) + await asyncio.wait_for(request_enqueued.wait(), timeout=_STREAM_TEST_TIMEOUT) + await active_task._event_queue_subscribers.enqueue_event((RuntimeError("producer failed"), None)) + await active_task._event_queue_subscribers.test_only_join_incoming_queue() + await active_task._event_queue_subscribers.close(immediate=True) + + with pytest.raises(RuntimeError, match="producer failed"): + await asyncio.wait_for(result, timeout=_STREAM_TEST_TIMEOUT) + + await active_task._event_queue_agent.close(immediate=True) @pytest.mark.asyncio -async def test_handler_reconciles_terminal_task_when_pipeline_sidecar_is_waiting_input( - monkeypatch: pytest.MonkeyPatch, - tmp_path, -) -> None: - monkeypatch.setenv("IAC_CODE_MODE", "pipeline") - cwd = tmp_path / "workspace" - cwd.mkdir() - context_id = "ctx-1" - task_id = "task-1" +async def test_request_scoped_registry_releases_admission_after_executor_failure() -> None: + released: list[str] = [] + admission_released = asyncio.Event() + executed = asyncio.Event() + + class FailingExecutor: + async def execute(self, _request_context, _event_queue) -> None: + executed.set() + raise RuntimeError("executor failed") + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + async def release(admission: str) -> None: + released.append(admission) + admission_released.set() + + call_context = ServerCallContext() + call_context.state["iac_code.recoverable_input_admission"] = "recovery-1" + request_context = RequestContext(call_context=call_context, task_id="task-1", context_id="ctx-1") + lease = RecoverableInputAdmissionLease( + "recovery-1", + acknowledge_enqueue=lambda token: call_context.state.pop("iac_code.recoverable_input_admission", None) == token, + release=release, + ) + RecoverableInputAdmissionCarrier.attach(request_context, lease) + registry = RequestScopedActiveTaskRegistry(agent_executor=FailingExecutor(), task_store=A2ATaskStore()) + active_task = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + + stream = active_task.subscribe(request=request_context) + with pytest.raises(RuntimeError, match="executor failed"): + await asyncio.wait_for(anext(stream), timeout=_STREAM_TEST_TIMEOUT) + await asyncio.wait_for(executed.wait(), timeout=_STREAM_TEST_TIMEOUT) + await asyncio.wait_for(admission_released.wait(), timeout=_STREAM_TEST_TIMEOUT) + + assert call_context.state == {} + assert released == ["recovery-1"] + + +@pytest.mark.asyncio +async def test_request_scoped_registry_replaces_finished_task_after_durable_reopen() -> None: + class IdleExecutor: + async def execute(self, _request_context, _event_queue) -> None: + return None + + async def cancel(self, _request_context, _event_queue) -> None: + return None + call_context = ServerCallContext() store = A2ATaskStore() - ctx = await store.get_or_create_context( - context_id=context_id, - cwd=str(cwd), - runtime_factory=lambda session_id: SimpleNamespace(session_id=session_id), + reopened = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), ) - ctx.active_task_id = task_id - store.mirror_context(ctx) + await store.save(reopened, call_context) + registry = RequestScopedActiveTaskRegistry(agent_executor=IdleExecutor(), task_store=store) + finished = RequestScopedActiveTask( + agent_executor=IdleExecutor(), + task_id="task-1", + task_manager=SimpleNamespace(), + ) + finished._is_finished.set() + registry._active_tasks["task-1"] = finished + + replacement = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + + assert replacement is not finished + assert await registry.get("task-1") is replacement + + replacement._producer_task.cancel() + replacement._consumer_task.cancel() + await asyncio.gather(replacement._producer_task, replacement._consumer_task, return_exceptions=True) + await replacement._event_queue_agent.close(immediate=True) + await replacement._event_queue_subscribers.close(immediate=True) + + +@pytest.mark.asyncio +async def test_request_scoped_registry_retires_unfinished_sdk_lifecycle_for_durable_recovery() -> None: + class IdleExecutor: + async def execute(self, _request_context, _event_queue) -> None: + return None + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + call_context = ServerCallContext() + store = A2ATaskStore() await store.save( Task( - id=task_id, - context_id=context_id, - status=TaskStatus(state=TaskState.TASK_STATE_FAILED), + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), ), call_context, ) + registry = RequestScopedActiveTaskRegistry(agent_executor=IdleExecutor(), task_store=store) + stale = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + assert stale._is_finished.is_set() is False + assert stale._producer_task is not None and not stale._producer_task.done() + assert stale._consumer_task is not None and not stale._consumer_task.done() + stale._task_manager._current_task = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_CANCELED), + ) + canonical = await store.get("task-1", call_context) + assert canonical is not None + assert canonical.status.state == TaskState.TASK_STATE_INPUT_REQUIRED - pending_input = { - "inputId": "input-confirm_and_select-1", - "kind": "candidate_selection", - "prompt": "请选择方案", - "options": [{"name": "方案A", "candidate_index": 0}], - } - pending_event = { - "schemaVersion": "1.0", - "extensionUri": "urn:iac-code:a2a:pipeline-events:v1", - "eventId": "evt-selection", - "sequence": 1, - "createdAt": "2026-06-08T10:00:00Z", - "eventType": "input_required", - "scope": "step", - "pipelineRunId": context_id, - "taskId": task_id, - "contextId": context_id, - "pipelineName": "selling", - "status": "input_required", - "step": {"runId": "step-confirm_and_select-1", "id": "confirm_and_select", "attempt": 1}, - "input": pending_input, - "data": pending_input, - } - pipeline_dir = a2a_pipeline_dir_for_session(cwd=str(cwd), session_id=ctx.session_id) - A2APipelineJournal(pipeline_dir).append(pending_event) - A2APipelineSnapshotStore(pipeline_dir).save(reduce_pipeline_events([pending_event])) - observed: dict[str, int] = {} + await registry.retire_for_recovery("task-1") - async def sdk_send(_handler, _params, sdk_context): - task = await store.get(task_id, sdk_context) - assert task is not None - observed["state"] = task.status.state - return task + assert await registry.get("task-1") is None + assert stale._is_finished.is_set() is True + assert stale._producer_task.done() + assert stale._consumer_task.done() + + replacement = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + assert replacement is not stale + recovered = await replacement.get_task() + assert recovered.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + + await registry.retire_for_recovery("task-1") + + +@pytest.mark.asyncio +async def test_recovery_replacement_rejects_a_contender_before_the_admitted_request() -> None: + observed: list[str | None] = [] + + class RecordingExecutor: + async def execute(self, request_context, event_queue) -> None: + observed.append(RecoverableInputAdmissionCarrier.read(request_context)) + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + call_context = ServerCallContext() + store = A2ATaskStore() + await store.save( + Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + registry = RequestScopedActiveTaskRegistry(agent_executor=RecordingExecutor(), task_store=store) + assert ( + await registry.reconcile_and_replace_for_recovery( + "task-1", + call_context=call_context, + context_id="ctx-1", + acquire_admission=lambda: asyncio.sleep(0, result="recovery-1"), + ) + == "recovery-1" + ) + replacement = await registry.get("task-1") + assert replacement is not None + contender = RequestContext(call_context=ServerCallContext(), task_id="task-1", context_id="ctx-1") + with pytest.raises(InvalidParamsError, match="recovery continuation is reserved"): + await anext(replacement.subscribe(request=contender)) + + admitted = RequestContext(call_context=call_context, task_id="task-1", context_id="ctx-1") + RecoverableInputAdmissionCarrier.attach(admitted, "recovery-1") + events = [event async for event in replacement.subscribe(request=admitted)] + + assert [event.status.state for event in events] == [TaskState.TASK_STATE_WORKING] + assert observed == ["recovery-1"] + assert replacement._recovery_admission is None + await registry.retire_for_recovery("task-1") + + +@pytest.mark.asyncio +async def test_recovery_claim_blocks_an_old_lifecycle_contender_before_admission() -> None: + class IdleExecutor: + async def execute(self, _request_context, _event_queue) -> None: + return None + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + call_context = ServerCallContext() + store = A2ATaskStore() + await store.save( + Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + registry = RequestScopedActiveTaskRegistry(agent_executor=IdleExecutor(), task_store=store) + old = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + contender_context = RequestContext( + call_context=ServerCallContext(), + task_id="task-1", + context_id="ctx-1", + ) + + await old._lock.acquire() + contender = asyncio.create_task(old.enqueue_request(contender_context)) + await asyncio.sleep(0) + replacement_task = asyncio.create_task( + registry.reconcile_and_replace_for_recovery( + "task-1", + call_context=call_context, + context_id="ctx-1", + acquire_admission=lambda: asyncio.sleep(0, result="recovery-1"), + ) + ) + await asyncio.sleep(0) + old._lock.release() + + with pytest.raises(InvalidParamsError, match="recovery replacement is pending"): + await contender + assert await replacement_task == "recovery-1" + replacement = await registry.get("task-1") + assert replacement is not None and replacement is not old + await registry.cancel_recovery_reservation("task-1", "recovery-1") + + +@pytest.mark.asyncio +async def test_recovery_replacement_waits_for_an_already_enqueued_old_lifecycle_request() -> None: + executed = asyncio.Event() + working_save_started = asyncio.Event() + release_working_save = asyncio.Event() + + class GatedTaskStore(A2ATaskStore): + async def save(self, task, context=None) -> None: + if task.status.state == TaskState.TASK_STATE_WORKING: + working_save_started.set() + await release_working_save.wait() + await super().save(task, context) + + class RecordingExecutor: + async def execute(self, _request_context, event_queue) -> None: + executed.set() + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + call_context = ServerCallContext() + store = GatedTaskStore() + await store.save( + Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + registry = RequestScopedActiveTaskRegistry(agent_executor=RecordingExecutor(), task_store=store) + old = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + contender_context = RequestContext( + call_context=ServerCallContext(), + task_id="task-1", + context_id="ctx-1", + ) + + await old._request_lock.acquire() + await old.enqueue_request(contender_context) + while old._request_queue.qsize() > 0: + await asyncio.sleep(0) + + async def acquire_admission() -> str | None: + task = await store.get("task-1", call_context) + assert task is not None + return None if task.status.state == TaskState.TASK_STATE_WORKING else "recovery-1" + + released: list[str] = [] + + replacement_task = asyncio.create_task( + registry.reconcile_and_replace_for_recovery( + "task-1", + call_context=call_context, + context_id="ctx-1", + acquire_admission=acquire_admission, + release_admission=lambda token: asyncio.sleep(0, result=released.append(token)), + ) + ) + await asyncio.sleep(0) + old._request_lock.release() + + try: + await asyncio.wait_for(executed.wait(), timeout=_STREAM_TEST_TIMEOUT) + await asyncio.wait_for(working_save_started.wait(), timeout=_STREAM_TEST_TIMEOUT) + await asyncio.sleep(0) + assert replacement_task.done() is False + release_working_save.set() + assert await asyncio.wait_for(replacement_task, timeout=_STREAM_TEST_TIMEOUT) is None + assert await registry.get("task-1") is old + assert old._is_finished.is_set() is False + assert released == [] + finally: + release_working_save.set() + if not replacement_task.done(): + replacement_task.cancel() + await asyncio.gather(replacement_task, return_exceptions=True) + await registry.retire_for_recovery("task-1") + + +@pytest.mark.asyncio +async def test_recovery_drain_fails_closed_when_the_old_consumer_exits_before_projection() -> None: + save_started = asyncio.Event() + release_save = asyncio.Event() + + class FailingTaskStore(A2ATaskStore): + async def save(self, task, context=None) -> None: + if task.status.state == TaskState.TASK_STATE_WORKING: + save_started.set() + await release_save.wait() + raise RuntimeError("projection failed") + await super().save(task, context) + + class WorkingExecutor: + async def execute(self, _request_context, event_queue) -> None: + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + call_context = ServerCallContext() + store = FailingTaskStore() + await store.save( + Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + registry = RequestScopedActiveTaskRegistry(agent_executor=WorkingExecutor(), task_store=store) + old = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + await old.enqueue_request(RequestContext(call_context=ServerCallContext(), task_id="task-1", context_id="ctx-1")) + await asyncio.wait_for(save_started.wait(), timeout=_STREAM_TEST_TIMEOUT) + while old.has_unfinished_requests(): + await asyncio.sleep(0) + assert old._request_lock.locked() + + replacement_task = asyncio.create_task( + registry.reconcile_and_replace_for_recovery( + "task-1", + call_context=call_context, + context_id="ctx-1", + acquire_admission=lambda: asyncio.sleep(0, result="recovery-1"), + ) + ) + release_save.set() + + try: + with pytest.raises(InvalidParamsError, match="ended before the accepted request settled"): + await asyncio.wait_for(replacement_task, timeout=_STREAM_TEST_TIMEOUT) + assert old._consumer_task is not None and old._consumer_task.done() + finally: + release_save.set() + if not replacement_task.done(): + replacement_task.cancel() + await asyncio.gather(replacement_task, return_exceptions=True) + await registry.retire_for_recovery("task-1") + + +@pytest.mark.asyncio +async def test_recovery_drain_for_one_task_does_not_block_registry_operations_for_another() -> None: + task_a_started = asyncio.Event() + release_task_a = asyncio.Event() + + class BlockingExecutor: + async def execute(self, request_context, event_queue) -> None: + if request_context.task_id != "task-a": + return + task_a_started.set() + await release_task_a.wait() + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-a", + context_id="ctx-a", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + ) + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + call_context = ServerCallContext() + store = A2ATaskStore() + for task_id, context_id in (("task-a", "ctx-a"), ("task-b", "ctx-b")): + await store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + registry = RequestScopedActiveTaskRegistry(agent_executor=BlockingExecutor(), task_store=store) + old_a = await registry.get_or_create( + "task-a", + call_context=call_context, + context_id="ctx-a", + create_task_if_missing=True, + ) + await old_a.enqueue_request(RequestContext(call_context=ServerCallContext(), task_id="task-a", context_id="ctx-a")) + await asyncio.wait_for(task_a_started.wait(), timeout=_STREAM_TEST_TIMEOUT) + + async def acquire_a() -> str | None: + task = await store.get("task-a", call_context) + assert task is not None + return None if task.status.state == TaskState.TASK_STATE_WORKING else "recovery-1" + + recovery_a = asyncio.create_task( + registry.reconcile_and_replace_for_recovery( + "task-a", + call_context=call_context, + context_id="ctx-a", + acquire_admission=acquire_a, + ) + ) + await asyncio.sleep(0) + + try: + task_b = await asyncio.wait_for( + registry.get_or_create( + "task-b", + call_context=call_context, + context_id="ctx-b", + create_task_if_missing=True, + ), + timeout=0.5, + ) + assert task_b.task_id == "task-b" + finally: + release_task_a.set() + await asyncio.wait_for(recovery_a, timeout=_STREAM_TEST_TIMEOUT) + await registry.retire_for_recovery("task-a") + await registry.retire_for_recovery("task-b") + + +@pytest.mark.asyncio +async def test_cancelled_recovery_replacement_finishes_retiring_old_lifecycle() -> None: + class IdleExecutor: + async def execute(self, _request_context, _event_queue) -> None: + return None + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + call_context = ServerCallContext() + store = A2ATaskStore() + await store.save( + Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + registry = RequestScopedActiveTaskRegistry(agent_executor=IdleExecutor(), task_store=store) + old = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + close_started = asyncio.Event() + release_close = asyncio.Event() + close_finished = asyncio.Event() + + async def gated_close(*, immediate: bool) -> None: + assert immediate is True + close_started.set() + await release_close.wait() + close_finished.set() + + old._event_queue_agent.close = gated_close + released: list[str] = [] + replacement = asyncio.create_task( + registry.reconcile_and_replace_for_recovery( + "task-1", + call_context=call_context, + context_id="ctx-1", + acquire_admission=lambda: asyncio.sleep(0, result="recovery-1"), + release_admission=lambda token: asyncio.sleep(0, result=released.append(token)), + ) + ) + await close_started.wait() + + replacement.cancel() + await asyncio.sleep(0) + assert replacement.done() is False + + release_close.set() + with pytest.raises(asyncio.CancelledError): + await replacement + assert close_finished.is_set() + assert old._producer_task.done() + assert old._consumer_task.done() + assert await registry.get("task-1") is None + assert released == ["recovery-1"] + + +@pytest.mark.asyncio +async def test_stale_sdk_status_event_restores_task_manager_to_canonical_projection() -> None: + class IdleExecutor: + async def execute(self, _request_context, _event_queue) -> None: + return None + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + def status(state: TaskState.ValueType, seconds: int) -> TaskStatus: + value = TaskStatus(state=state) + value.timestamp.seconds = seconds + return value + + call_context = ServerCallContext() + store = A2ATaskStore() + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_INPUT_REQUIRED, 100)), + call_context, + ) + registry = RequestScopedActiveTaskRegistry(agent_executor=IdleExecutor(), task_store=store) + active_task = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.active_task = asyncio.current_task() + record.updated_at = 200 + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_WORKING, 250)), + call_context, + ) + + await active_task._event_queue_agent.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=status(TaskState.TASK_STATE_CANCELED, 150), + ) + ) + await active_task._event_queue_agent.test_only_join_incoming_queue() + canonical = await store.get("task-1", call_context) + + assert canonical is not None + assert canonical.status.state == TaskState.TASK_STATE_WORKING + assert active_task._task_manager._current_task.status.state == TaskState.TASK_STATE_WORKING + assert active_task._is_finished.is_set() is False + + await active_task._event_queue_agent.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=status(TaskState.TASK_STATE_CANCELED, 300), + ) + ) + await active_task._event_queue_agent.test_only_join_incoming_queue() + canonical = await store.get("task-1", call_context) + assert canonical is not None + assert canonical.status.state == TaskState.TASK_STATE_CANCELED + assert active_task._task_manager._current_task.status.state == TaskState.TASK_STATE_CANCELED + assert active_task._is_finished.is_set() is True + await registry.retire_for_recovery("task-1") + + +@pytest.mark.asyncio +async def test_stale_sdk_status_rebuilds_projection_when_owner_cache_is_older() -> None: + def status(state: TaskState.ValueType, seconds: int) -> TaskStatus: + value = TaskStatus(state=state) + value.timestamp.seconds = seconds + return value + + call_context = ServerCallContext() + store = A2ATaskStore() + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_CANCELED, 100)), + call_context, + ) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.state = "working" + record.updated_at = 200 + incoming = Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_FAILED, 150)) + + await store.save(incoming, call_context) + visible = await store.get("task-1", call_context) + + assert incoming.status.state == TaskState.TASK_STATE_WORKING + assert incoming.status.timestamp.seconds == 200 + assert visible is not None + assert visible.status.state == TaskState.TASK_STATE_WORKING + assert visible.status.timestamp.seconds == 200 + + +@pytest.mark.asyncio +async def test_request_scoped_registry_old_cleanup_cannot_remove_replacement() -> None: + registry = RequestScopedActiveTaskRegistry(agent_executor=SimpleNamespace(), task_store=A2ATaskStore()) + finished = RequestScopedActiveTask( + agent_executor=SimpleNamespace(), + task_id="task-1", + task_manager=SimpleNamespace(), + ) + replacement = RequestScopedActiveTask( + agent_executor=SimpleNamespace(), + task_id="task-1", + task_manager=SimpleNamespace(), + ) + registry._active_tasks["task-1"] = finished + + await registry._lock.acquire() + try: + registry._on_active_task_cleanup(finished) + registry._active_tasks["task-1"] = replacement + finally: + registry._lock.release() + await asyncio.gather(*tuple(registry._cleanup_tasks)) + + assert await registry.get("task-1") is replacement + + await finished._event_queue_agent.close(immediate=True) + await finished._event_queue_subscribers.close(immediate=True) + await replacement._event_queue_agent.close(immediate=True) + await replacement._event_queue_subscribers.close(immediate=True) + + +@pytest.mark.asyncio +async def test_reused_sdk_producer_reads_admission_from_each_request_context() -> None: + observed: list[str | None] = [] + + class RecordingExecutor: + async def execute(self, request_context, event_queue) -> None: + observed.append(RecoverableInputAdmissionCarrier.read(request_context)) + if request_context.current_task is None: + await event_queue.enqueue_event( + Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + ) + else: + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + ) + + async def cancel(self, _request_context, event_queue) -> None: + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_CANCELED), + ) + ) + + call_context = ServerCallContext() + store = A2ATaskStore() + registry = ActiveTaskRegistry(agent_executor=RecordingExecutor(), task_store=store) + active_task = await registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + + for admission in ("recovery-first", "recovery-second"): + request_context = RequestContext( + call_context=call_context, + task_id="task-1", + context_id="ctx-1", + ) + RecoverableInputAdmissionCarrier.attach(request_context, admission) + events = [event async for event in active_task.subscribe(request=request_context)] + assert any( + isinstance(event, (Task, TaskStatusUpdateEvent)) + and event.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + for event in events + ) + + assert observed == ["recovery-first", "recovery-second"] + await active_task.cancel(call_context) + + +@pytest.mark.asyncio +async def test_closed_transport_tracker_reports_closed_stage_on_registration() -> None: + tracker = create_pipeline_transport_delivery_tracker() + close_pipeline_transport_delivery_tracker(tracker) + stages: list[str] = [] + + with bind_pipeline_transport_delivery_tracker(tracker): + completion = register_pipeline_transport_delivery( + object(), + stage_observer=lambda stage, _at_ns: stages.append(stage), + ) + + with pytest.raises(PipelineTransportDeliveryClosedError): + await completion + assert stages == ["registered", "closed"] + + +@pytest.mark.asyncio +async def test_dispatcher_handles_unary_v03_message(monkeypatch, tmp_path) -> None: + loop = FakeAgentLoop([TextDeltaEvent(text="hello from dispatcher")]) + + def factory(options): + return FakeRuntime(agent_loop=loop, session_id=options.session_id) + + monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", factory) + components = create_runtime_components(model="qwen3.6-plus", host="127.0.0.1", port=41242) + dispatcher = A2AJsonRpcDispatcher(components) + + response = await dispatcher.dispatch( + { + "jsonrpc": "2.0", + "id": "1", + "method": "message/send", + "params": { + "message": { + "messageId": "msg-1", + "role": "user", + "parts": [{"kind": "text", "text": "hello"}], + "metadata": {"iac_code": {"cwd": str(tmp_path)}}, + }, + "configuration": {"acceptedOutputModes": ["text/plain"]}, + }, + } + ) + + assert response["id"] == "1" + assert response["result"]["status"]["state"] == "input-required" + session_id = components.task_store._contexts[response["result"]["contextId"]].session_id + assert response["result"]["metadata"]["iac_code"]["iacCodeSessionId"] == session_id + assert loop.prompts == ["hello"] + await components.aclose() + + +@pytest.mark.asyncio +async def test_dispatcher_rejects_explicit_invalid_run_mode(tmp_path) -> None: + components = create_runtime_components(model="qwen3.6-plus", host="127.0.0.1", port=41242) + dispatcher = A2AJsonRpcDispatcher(components) + + response = await dispatcher.dispatch( + { + "jsonrpc": "2.0", + "id": "invalid-run-mode", + "method": "message/send", + "params": { + "message": { + "messageId": "msg-invalid-run-mode", + "role": "user", + "parts": [{"kind": "text", "text": "hello"}], + "metadata": {"iac_code": {"cwd": str(tmp_path), "run_mode": "pipline"}}, + }, + "configuration": {"acceptedOutputModes": ["text/plain"]}, + }, + } + ) + + assert response["id"] == "invalid-run-mode" + assert response["error"]["code"] == -32602 + assert response["error"]["message"] == "Unsupported run mode." + await components.aclose() + + +@pytest.mark.asyncio +async def test_dispatcher_stream_yields_events(monkeypatch, tmp_path) -> None: + loop = FakeAgentLoop([TextDeltaEvent(text="streamed")]) + runtime = FakeRuntime(agent_loop=loop, session_id="session-1") + monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", lambda options: runtime) + components = create_runtime_components(model="qwen3.6-plus", host="127.0.0.1", port=41242) + dispatcher = A2AJsonRpcDispatcher(components) + + events = [ + event + async for event in dispatcher.dispatch_stream( + { + "jsonrpc": "2.0", + "id": "2", + "method": "message/stream", + "params": { + "message": { + "messageId": "msg-2", + "role": "user", + "parts": [{"kind": "text", "text": "hello"}], + "metadata": {"iac_code": {"cwd": str(tmp_path)}}, + }, + "configuration": {"acceptedOutputModes": ["text/plain"]}, + }, + } + ) + ] + + assert any(event["result"]["status"]["state"] == "working" for event in events) + assert events[-1]["result"]["status"]["state"] == "input-required" + await components.aclose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ( + "phase", + "release_ready", + "control_task_id", + "admission_allowed", + "sidecar_task_id", + "expected_task_state", + "expected_persisted_state", + "expected_active_task_id", + ), + [ + pytest.param( + "terminated", + True, + "task-1", + True, + "task-1", + TaskState.TASK_STATE_INPUT_REQUIRED, + "input-required", + None, + id="release-ready", + ), + pytest.param( + "terminated", + False, + "task-1", + True, + "task-1", + TaskState.TASK_STATE_INPUT_REQUIRED, + "input-required", + None, + id="terminated-backup-blocked", + ), + pytest.param( + "terminating", + False, + "task-1", + False, + "task-1", + TaskState.TASK_STATE_FAILED, + "failed", + "task-1", + id="termination-in-flight", + ), + pytest.param( + "terminated", + False, + "task-other", + False, + "task-1", + TaskState.TASK_STATE_FAILED, + "failed", + "task-1", + id="blocked-control-task-mismatch", + ), + pytest.param( + "terminated", + False, + "task-1", + False, + "task-1", + TaskState.TASK_STATE_FAILED, + "failed", + "task-1", + id="blocked-control-live-background-work", + ), + pytest.param( + "terminated", + False, + "task-1", + True, + "task-other", + TaskState.TASK_STATE_FAILED, + "failed", + "task-1", + id="sidecar-task-mismatch", + ), + ], +) +async def test_handler_reconciles_terminal_task_when_pipeline_sidecar_is_waiting_input( + monkeypatch: pytest.MonkeyPatch, + tmp_path, + phase: str, + release_ready: bool, + control_task_id: str, + admission_allowed: bool, + sidecar_task_id: str, + expected_task_state: int, + expected_persisted_state: str, + expected_active_task_id: str | None, +) -> None: + monkeypatch.setenv("IAC_CODE_MODE", "pipeline") + cwd = tmp_path / "workspace" + cwd.mkdir() + context_id = "ctx-1" + task_id = "task-1" + call_context = ServerCallContext() + persistence = A2APersistenceStore(tmp_path / "a2a") + store = A2ATaskStore(persistence=persistence) + ctx = await store.get_or_create_context( + context_id=context_id, + cwd=str(cwd), + runtime_factory=lambda session_id: SimpleNamespace(session_id=session_id), + ) + ctx.active_task_id = task_id + store.mirror_context(ctx) + await store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_FAILED), + ), + call_context, + ) + issued_admissions: list[str] = [] + released_admissions: list[str] = [] + + async def reserve_admission(*, context_id: str, task_id: str, owner: str) -> str | None: + assert context_id == "ctx-1" + assert task_id == "task-1" + if not admission_allowed: + return None + admission = f"recovery-{owner}-{task_id}" + issued_admissions.append(admission) + return admission + + async def release_admission(admission: str) -> None: + released_admissions.append(admission) + + store.set_execution_control_provider( + lambda _context_id: { + "taskId": control_task_id, + "phase": phase, + "releaseReady": release_ready, + "backup": {"status": "shared_committed" if release_ready else "blocked"}, + }, + lambda: phase == "terminating" or not release_ready, + reserve_admission, + release_admission, + ) + + pending_input = { + "inputId": "input-confirm_and_select-1", + "kind": "candidate_selection", + "prompt": "请选择方案", + "options": [{"name": "方案A", "candidate_index": 0}], + } + pending_event = { + "schemaVersion": "1.0", + "extensionUri": "urn:iac-code:a2a:pipeline-events:v1", + "eventId": "evt-selection", + "sequence": 1, + "createdAt": "2026-06-08T10:00:00Z", + "eventType": "input_required", + "scope": "step", + "pipelineRunId": context_id, + "taskId": sidecar_task_id, + "contextId": context_id, + "pipelineName": "selling", + "status": "input_required", + "step": {"runId": "step-confirm_and_select-1", "id": "confirm_and_select", "attempt": 1}, + "input": pending_input, + "data": pending_input, + } + pipeline_dir = a2a_pipeline_dir_for_session(cwd=str(cwd), session_id=ctx.session_id) + A2APipelineJournal(pipeline_dir).append(pending_event) + A2APipelineSnapshotStore(pipeline_dir).save(reduce_pipeline_events([pending_event])) + observed: dict[str, int] = {} + + async def sdk_send(_handler, _params, sdk_context): + task = await store.get(task_id, sdk_context) + assert task is not None + observed["state"] = task.status.state + return task + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send", sdk_send) + retired_sdk_tasks: list[str] = [] + + class RecoveryAwareRegistry: + async def reconcile_and_replace_for_recovery( + self, recovered_task_id: str, *, acquire_admission, **_kwargs + ) -> str | None: + admission = await acquire_admission() + if admission is not None: + retired_sdk_tasks.append(recovered_task_id) + return admission + + async def cancel_recovery_reservation(self, _task_id: str, _admission: str) -> None: + return None + + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler._active_task_registry = RecoveryAwareRegistry() + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + params = SimpleNamespace(message=SimpleNamespace(task_id=task_id, context_id=context_id)) + + if expected_task_state == TaskState.TASK_STATE_INPUT_REQUIRED: + original_save_task = persistence.save_task + + def fail_recovery_task_save(snapshot): + if snapshot.state == "input-required": + raise OSError("temporary recovery write failure") + original_save_task(snapshot) + + monkeypatch.setattr(persistence, "save_task", fail_recovery_task_save) + with pytest.raises(OSError, match="temporary recovery write failure"): + await handler.on_message_send(params, call_context) + assert observed == {} + failed_recovery_task = await store.get(task_id, call_context) + assert failed_recovery_task is not None + assert failed_recovery_task.status.state == TaskState.TASK_STATE_FAILED + assert persistence.load_task(task_id).state == "failed" + assert persistence.load_context(context_id).active_task_id == task_id + assert (await store.get_context_record(context_id)).active_task_id == task_id + monkeypatch.setattr(persistence, "save_task", original_save_task) + + if not admission_allowed and sidecar_task_id == task_id: + with pytest.raises(InvalidParamsError, match="already being recovered"): + await handler.on_message_send(params, call_context) + assert observed == {} + assert issued_admissions == [] + assert released_admissions == [] + assert retired_sdk_tasks == [] + assert persistence.load_task(task_id).state == expected_persisted_state + assert persistence.load_context(context_id).active_task_id == expected_active_task_id + return + + result = await handler.on_message_send(params, call_context) + + assert isinstance(result, Task) + assert observed["state"] == expected_task_state + assert result.status.state == expected_task_state + persisted_task = persistence.load_task(task_id) + assert persisted_task is not None + assert persisted_task.state == expected_persisted_state + persisted_context = persistence.load_context(context_id) + assert persisted_context is not None + assert persisted_context.active_task_id == expected_active_task_id + assert (await store.get_context_record(context_id)).active_task_id == expected_active_task_id + session_dir = SessionStorage().session_dir(str(cwd), ctx.session_id) + context_snapshot = json.loads((session_dir / "a2a" / "context.json").read_text(encoding="utf-8")) + assert context_snapshot["active_task_id"] == expected_active_task_id + if expected_task_state == TaskState.TASK_STATE_INPUT_REQUIRED: + assert len(issued_admissions) == 2 + assert released_admissions == issued_admissions + else: + assert issued_admissions == [] + assert released_admissions == [] + assert retired_sdk_tasks == ([task_id] if expected_task_state == TaskState.TASK_STATE_INPUT_REQUIRED else []) + + if expected_task_state == TaskState.TASK_STATE_INPUT_REQUIRED: + for seconds, late_state in enumerate( + ( + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_COMPLETED, + ), + start=1, + ): + late_terminal = Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=late_state), + ) + late_terminal.status.timestamp.FromSeconds(seconds) + await store.save(late_terminal, call_context) + + visible_task = await store.get(task_id, call_context) + assert visible_task is not None + assert visible_task.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + persisted_task = persistence.load_task(task_id) + assert persisted_task is not None + assert persisted_task.state == "input-required" + persisted_context = persistence.load_context(context_id) + assert persisted_context is not None + assert persisted_context.active_task_id is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "late_state", + [ + TaskState.TASK_STATE_FAILED, + TaskState.TASK_STATE_CANCELED, + TaskState.TASK_STATE_COMPLETED, + ], +) +async def test_recovered_running_execution_rejects_older_terminal_projection(tmp_path, late_state: int) -> None: + cwd = tmp_path / "workspace" + cwd.mkdir() + context_id = "ctx-1" + task_id = "task-1" + call_context = ServerCallContext() + persistence = A2APersistenceStore(tmp_path / "a2a") + store = A2ATaskStore(persistence=persistence) + service = ExecutionControlService(persistence_root=tmp_path / "a2a", backup_service=None) + store.set_execution_control_provider( + service.snapshot_for_context, + service.has_active_work, + service.reserve_recoverable_input_continuation, + service.release_recoverable_input_continuation, + ) + context_record = await store.get_or_create_context( + context_id=context_id, + cwd=str(cwd), + runtime_factory=lambda session_id: SimpleNamespace(session_id=session_id), + ) + context_record.active_task_id = task_id + store.mirror_context(context_record) + failed = Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_FAILED), + ) + failed.status.timestamp.FromSeconds(100) + await store.save(failed, call_context) + owner = store.owner_for_context(call_context) + + original = await service.begin_execution( + context_id=context_id, + task_id=task_id, + owner=owner, + cwd=str(cwd), + ) + current = asyncio.current_task() + assert current is not None + await original.detach_task(current, execution_status="input-required") + original.phase = "terminated" + original.execution_status = "canceled" + original.backup = {"status": "blocked", "error": "shared backup unavailable"} + original.release_ready = False + original.revision += 1 + await original._persist_snapshot(original.snapshot()) + + recovered_task = Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + recovered_task.status.timestamp.FromSeconds(200) + admission = await store.reconcile_recoverable_input_required_task( + recovered_task, + context_record, + call_context, + ) + assert admission is not None + recovered = await service.begin_execution( + context_id=context_id, + task_id=task_id, + owner=owner, + cwd=str(cwd), + continue_input_required=True, + recoverable_input_admission=admission, + ) + assert recovered.phase == "running" + + late_terminal = Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=late_state), + ) + late_terminal.status.timestamp.FromSeconds(150) + await store.save(late_terminal, call_context) + + visible = await store.get(task_id, call_context) + assert visible is not None + assert visible.status.state == TaskState.TASK_STATE_INPUT_REQUIRED + assert persistence.load_task(task_id).state == "input-required" + assert persistence.load_context(context_id).active_task_id is None + + await recovered.detach_task(current, execution_status="input-required") + await service.close() + + +@pytest.mark.asyncio +async def test_handler_does_not_reconcile_a_running_sidecar_owner( + monkeypatch: pytest.MonkeyPatch, + tmp_path, +) -> None: + monkeypatch.setenv("IAC_CODE_MODE", "pipeline") + cwd = tmp_path / "workspace" + cwd.mkdir() + context_id = "ctx-1" + task_id = "task-1" + call_context = ServerCallContext() + persistence = A2APersistenceStore(tmp_path / "a2a") + store = A2ATaskStore(persistence=persistence) + context_record = await store.get_or_create_context( + context_id=context_id, + cwd=str(cwd), + runtime_factory=lambda session_id: SimpleNamespace(session_id=session_id), + ) + context_record.active_task_id = task_id + store.mirror_context(context_record) + await store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_FAILED), + ), + call_context, + ) + admission_calls: list[str] = [] + + async def reserve_admission(**_kwargs) -> str: + admission_calls.append("called") + return "unexpected" + + store.set_execution_control_provider(None, None, reserve_admission, None) + include_running_values: list[bool] = [] + + def running_sidecar_only(*, include_running: bool, **_kwargs) -> str | None: + include_running_values.append(include_running) + return task_id if include_running else None + + monkeypatch.setattr( + "iac_code.a2a.transports.dispatcher.recoverable_task_id_from_sidecar", + running_sidecar_only, + ) + + async def sdk_send(_handler, _params, sdk_context): + return await store.get(task_id, sdk_context) + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send", sdk_send) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + params = SimpleNamespace(message=SimpleNamespace(task_id=task_id, context_id=context_id)) + + result = await handler.on_message_send(params, call_context) + + assert isinstance(result, Task) + assert result.status.state == TaskState.TASK_STATE_FAILED + assert include_running_values == [False] + assert admission_calls == [] + assert persistence.load_task(task_id).state == "failed" + assert persistence.load_context(context_id).active_task_id == task_id + + +@pytest.mark.asyncio +async def test_handler_rejects_recovery_owned_by_another_process_before_sdk( + monkeypatch: pytest.MonkeyPatch, + tmp_path, +) -> None: + monkeypatch.setenv("IAC_CODE_MODE", "pipeline") + cwd = tmp_path / "workspace" + cwd.mkdir() + persistence_root = tmp_path / "a2a" + persistence = A2APersistenceStore(persistence_root) + context_id = "ctx-1" + task_id = "task-1" + call_context = ServerCallContext() + first_service = ExecutionControlService(persistence_root=persistence_root, backup_service=None) + second_service = ExecutionControlService(persistence_root=persistence_root, backup_service=None) + first_store = A2ATaskStore(persistence=persistence) + first_store.set_execution_control_provider( + first_service.snapshot_for_context, + first_service.has_active_work, + first_service.reserve_recoverable_input_continuation, + first_service.release_recoverable_input_continuation, + ) + context_record = await first_store.get_or_create_context( + context_id=context_id, + cwd=str(cwd), + runtime_factory=lambda session_id: SimpleNamespace(session_id=session_id), + ) + context_record.active_task_id = None + first_store.mirror_context(context_record) + await first_store.save( + Task( + id=task_id, + context_id=context_id, + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + admission = await first_service.reserve_recoverable_input_continuation( + context_id=context_id, + task_id=task_id, + owner=first_store.owner_for_context(call_context), + ) + assert admission is not None + + second_store = A2ATaskStore(persistence=A2APersistenceStore(persistence_root)) + second_store.set_execution_control_provider( + second_service.snapshot_for_context, + second_service.has_active_work, + second_service.reserve_recoverable_input_continuation, + second_service.release_recoverable_input_continuation, + ) + monkeypatch.setattr( + "iac_code.a2a.transports.dispatcher.recoverable_task_id_from_sidecar", + lambda **_kwargs: task_id, + ) + sdk_called = False + + async def sdk_send(_handler, _params, _context): + nonlocal sdk_called + sdk_called = True + return None + + async def hydrate(_params) -> None: + return None + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send", sdk_send) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = second_store + handler._active_task_registry = None + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + params = SimpleNamespace(message=SimpleNamespace(task_id=task_id, context_id=context_id)) + + try: + with pytest.raises(InvalidParamsError, match="already being recovered"): + await handler.on_message_send(params, call_context) + assert sdk_called is False + assert persistence.load_task(task_id).state == "input-required" + assert persistence.load_context(context_id).active_task_id is None + recovered = await first_service.begin_execution( + context_id=context_id, + task_id=task_id, + owner=first_store.owner_for_context(call_context), + cwd=str(cwd), + continue_input_required=True, + recoverable_input_admission=admission, + ) + assert recovered.phase == "running" + current = asyncio.current_task() + assert current is not None + await recovered.detach_task(current, execution_status="input-required") + finally: + await first_service.release_recoverable_input_continuation(admission) + await first_service.close() + await second_service.close() + + +@pytest.mark.asyncio +async def test_handler_restores_backup_before_hydrating_omitted_pipeline_task_id( + monkeypatch: pytest.MonkeyPatch, + tmp_path, +) -> None: + monkeypatch.setenv("IAC_CODE_MODE", "pipeline") + monkeypatch.setenv("IAC_CODE_CONFIG_DIR", str(tmp_path / "config")) + monkeypatch.setenv("IAC_CODE_CONFIG_BACKUP_DIR", str(tmp_path / "backup")) + cwd = tmp_path / "workspace" + cwd.mkdir() + context_id = "ctx-restore" + task_id = "task-restore" + store = A2ATaskStore() + ctx = await store.get_or_create_context( + context_id=context_id, + cwd=str(cwd), + runtime_factory=lambda session_id: SimpleNamespace(session_id=session_id), + ) + storage = SessionStorage() + storage.save(str(cwd), ctx.session_id, []) + pending_input = { + "inputId": "input-confirm_and_select-1", + "kind": "candidate_selection", + "prompt": "请选择方案", + "options": [{"name": "方案A", "candidate_index": 0}], + } + pending_event = { + "schemaVersion": "1.0", + "extensionUri": "urn:iac-code:a2a:pipeline-events:v1", + "eventId": "evt-selection", + "sequence": 1, + "createdAt": "2026-06-08T10:00:00Z", + "eventType": "input_required", + "scope": "step", + "pipelineRunId": context_id, + "taskId": task_id, + "contextId": context_id, + "pipelineName": "selling", + "status": "input_required", + "step": {"runId": "step-confirm_and_select-1", "id": "confirm_and_select", "attempt": 1}, + "input": pending_input, + "data": pending_input, + } + pipeline_dir = a2a_pipeline_dir_for_session(cwd=str(cwd), session_id=ctx.session_id) + A2APipelineJournal(pipeline_dir).append(pending_event) + A2APipelineSnapshotStore(pipeline_dir).save(reduce_pipeline_events([pending_event])) + backup_service = SessionBackupService(storage, retry_delays=()) + backup_service.initialize_session(str(cwd), ctx.session_id) + backup_service.backup_session(str(cwd), ctx.session_id, reason=BackupReason.INPUT_REQUIRED, critical=True) + primary_session_dir = storage.session_dir(str(cwd), ctx.session_id) + shutil.rmtree(primary_session_dir) + + class FakeExecutor: + async def _reconcile_session_before_route(self, *, context_id: str, cwd: str): + assert context_id == "ctx-restore" + return await asyncio.to_thread(backup_service.reconcile_session, cwd, ctx.session_id) + + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler.agent_executor = FakeExecutor() + params = SimpleNamespace(message=SimpleNamespace(task_id=None, context_id=context_id)) + + await handler._hydrate_recoverable_pipeline_task_id(params) + + assert params.message.task_id == task_id + assert storage.session_dir(str(cwd), ctx.session_id).is_dir() + + +@pytest.mark.asyncio +async def test_dispatcher_stream_backpressures_asgi_until_consumer_resumes() -> None: + first_chunk_consumed = asyncio.Event() + + async def app(_scope, _receive, send) -> None: + await send( + { + "type": "http.response.start", + "status": 200, + "headers": [(b"content-type", b"text/event-stream")], + } + ) + await send( + { + "type": "http.response.body", + "body": b'data: {"id":"1","result":{"index":1}}\n\n', + "more_body": True, + } + ) + first_chunk_consumed.set() + await send( + { + "type": "http.response.body", + "body": b'data: {"id":"1","result":{"index":2}}\n\n', + "more_body": False, + } + ) + + dispatcher = A2AJsonRpcDispatcher(SimpleNamespace(app=app)) + stream = dispatcher.dispatch_stream({"jsonrpc": "2.0", "id": "1"}) + + first = await anext(stream) + assert first["result"]["index"] == 1 + assert first_chunk_consumed.is_set() is False + + second = await anext(stream) + assert first_chunk_consumed.is_set() is True + assert second["result"]["index"] == 2 + with pytest.raises(StopAsyncIteration): + await anext(stream) + + await dispatcher.aclose() + + +@pytest.mark.asyncio +async def test_streaming_asgi_transport_cancels_app_before_response_start() -> None: + app_started = asyncio.Event() + app_cancelled = asyncio.Event() + + async def app(_scope, _receive, _send) -> None: + app_started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + app_cancelled.set() + raise + + transport = _StreamingASGITransport(app) + async with httpx.AsyncClient(transport=transport, base_url="http://transport.local") as client: + request_task = asyncio.create_task(client.get("/")) + await app_started.wait() + + request_task.cancel() + with pytest.raises(asyncio.CancelledError): + await request_task + + assert app_cancelled.is_set() + assert not any(task.get_name() == "a2a-streaming-asgi-dispatch" and not task.done() for task in asyncio.all_tasks()) + + +@pytest.mark.asyncio +async def test_message_stream_acknowledges_transport_delivery_only_when_resumed(monkeypatch) -> None: + observed: dict[str, asyncio.Future[None]] = {} + stages: list[str] = [] + update = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + + async def sdk_stream(_handler, _params, _context): + observed["completion"] = register_pipeline_transport_delivery( + update, + stage_observer=lambda stage, _at_ns: stages.append(stage), + ) + yield update + + async def hydrate(_params) -> None: + return None + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send_stream", sdk_stream) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = object() + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + params = SimpleNamespace(message=SimpleNamespace(task_id=None)) + + stream = handler.on_message_send_stream(params, SimpleNamespace()) + assert await anext(stream) is update + assert observed["completion"].done() is False + assert stages == ["registered", "dequeued"] + + with pytest.raises(StopAsyncIteration): + await anext(stream) + + assert observed["completion"].done() is True + assert stages == ["registered", "dequeued", "acknowledged"] + assert pipeline_transport_delivery_tracking_enabled() is False + + +@pytest.mark.asyncio +async def test_message_stream_finalizes_natural_execution_only_after_stream_exhaustion(monkeypatch) -> None: + call_context = ServerCallContext() + store = A2ATaskStore() + update = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + order: list[str] = [] + + async def sdk_stream(_handler, _params, _context): + order.append("yield") + yield update + NaturalCompletionGenerationCarrier.attach(SimpleNamespace(call_context=_context), 7) + order.append("exhausted") + + class Executor: + async def finalize_natural_execution(self, **kwargs) -> None: + assert kwargs["completion_generation"] == 7 + order.append("finalized") + + async def hydrate(_params) -> None: + return None + + async def reconcile(_params, _context) -> None: + return None + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send_stream", sdk_stream) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler.agent_executor = Executor() + handler._active_task_registry = None + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + + stream = handler.on_message_send_stream(params, call_context) + assert await anext(stream) is update + assert order == ["yield"] + with pytest.raises(StopAsyncIteration): + await anext(stream) + + assert order == ["yield", "exhausted", "finalized"] + + +@pytest.mark.asyncio +async def test_message_stream_disconnect_does_not_naturally_finalize_execution(monkeypatch) -> None: + call_context = ServerCallContext() + store = A2ATaskStore() + first = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + finalized = False + + async def sdk_stream(_handler, _params, _context): + yield first + NaturalCompletionGenerationCarrier.attach(SimpleNamespace(call_context=_context), 7) + yield TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + + class Executor: + async def finalize_natural_execution(self, **_kwargs) -> None: + nonlocal finalized + finalized = True + + async def hydrate(_params) -> None: + return None + + async def reconcile(_params, _context) -> None: + return None + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send_stream", sdk_stream) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler.agent_executor = Executor() + handler._active_task_registry = None + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + + stream = handler.on_message_send_stream(params, call_context) + assert await anext(stream) is first + await stream.aclose() + + assert finalized is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize("state", [TaskState.TASK_STATE_INPUT_REQUIRED, TaskState.TASK_STATE_COMPLETED]) +async def test_message_stream_boundary_aclose_finalizes_after_executor_detaches(monkeypatch, state) -> None: + call_context = ServerCallContext() + store = A2ATaskStore() + boundary = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=state), + ) + detach_entered = asyncio.Event() + allow_detach = asyncio.Event() + finalized: list[int] = [] + + class Executor: + async def finalize_natural_execution(self, **kwargs) -> None: + finalized.append(kwargs["completion_generation"]) + + executor = Executor() + + async def sdk_stream(_handler, _params, _context): + try: + yield boundary + finally: + detach_entered.set() + await allow_detach.wait() + NaturalCompletionGenerationCarrier.attach(SimpleNamespace(call_context=_context), 7) + + async def hydrate(_params) -> None: + return None + + async def reconcile(_params, _context) -> None: + return None + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send_stream", sdk_stream) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler.agent_executor = executor + handler._active_task_registry = None + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + + stream = handler.on_message_send_stream(params, call_context) + assert await anext(stream) is boundary + closing = asyncio.create_task(stream.aclose()) + await detach_entered.wait() + assert finalized == [] + allow_detach.set() + await closing + + assert finalized == [7] + + +@pytest.mark.asyncio +async def test_message_stream_boundary_aclose_finalizes_generation_attached_before_delivery(monkeypatch) -> None: + call_context = ServerCallContext() + boundary = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + finalized: list[int] = [] + + async def sdk_stream(_handler, _params, _context): + NaturalCompletionGenerationCarrier.attach(SimpleNamespace(call_context=_context), 7) + yield boundary + + class Executor: + async def finalize_natural_execution(self, **kwargs) -> None: + finalized.append(kwargs["completion_generation"]) + + async def hydrate(_params) -> None: + return None + + async def reconcile(_params, _context) -> None: + return None + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send_stream", sdk_stream) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = A2ATaskStore() + handler.agent_executor = Executor() + handler._active_task_registry = None + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + + stream = handler.on_message_send_stream(params, call_context) + assert await anext(stream) is boundary + await stream.aclose() + + assert finalized == [7] + + +@pytest.mark.asyncio +async def test_message_stream_pending_permission_boundary_aclose_without_generation_does_not_finalize( + monkeypatch, +) -> None: + call_context = ServerCallContext() + boundary = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ) + finalized = False + + async def sdk_stream(_handler, _params, _context): + yield boundary + + class Executor: + async def finalize_natural_execution(self, **_kwargs) -> None: + nonlocal finalized + finalized = True + + async def hydrate(_params) -> None: + return None + + async def reconcile(_params, _context) -> None: + return None + + monkeypatch.setattr(DefaultRequestHandler, "on_message_send_stream", sdk_stream) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = A2ATaskStore() + handler.agent_executor = Executor() + handler._active_task_registry = None + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + + stream = handler.on_message_send_stream(params, call_context) + assert await anext(stream) is boundary + await stream.aclose() + + assert finalized is False + + +@pytest.mark.asyncio +async def test_detached_permission_response_finalizes_after_terminal_producer_drains() -> None: + call_context = ServerCallContext() + request_context = SimpleNamespace(call_context=call_context) + NaturalCompletionGenerationCarrier.prepare(request_context) + producer_can_finish = asyncio.Event() + finalized = asyncio.Event() + + class RequestContextBuilder: + async def build(self, **_kwargs): + return request_context + + class Executor: + async def execute(self, _request_context, event_queue) -> None: + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + ) + await producer_can_finish.wait() + NaturalCompletionGenerationCarrier.attach(request_context, 7) + + async def finalize_natural_execution(self, **kwargs) -> None: + assert kwargs["completion_generation"] == 7 + finalized.set() - monkeypatch.setattr(DefaultRequestHandler, "on_message_send", sdk_send) handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) - handler.task_store = store - handler._validate_extensions = lambda _context: None - handler._validate_pipeline_message_request = lambda _params: None - params = SimpleNamespace(message=SimpleNamespace(task_id=task_id, context_id=context_id)) + handler.task_store = A2ATaskStore() + handler.agent_executor = Executor() + handler._request_context_builder = RequestContextBuilder() + handler._detached_message_producers = set() + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + task = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) - result = await handler.on_message_send(params, call_context) + stream = handler._on_inactive_permission_send_stream(params, call_context, task=task) + event = await anext(stream) + assert event.status.state == TaskState.TASK_STATE_COMPLETED + await stream.aclose() + producer_can_finish.set() - assert isinstance(result, Task) - assert observed["state"] == TaskState.TASK_STATE_INPUT_REQUIRED - assert result.status.state == TaskState.TASK_STATE_INPUT_REQUIRED - session_dir = SessionStorage().session_dir(str(cwd), ctx.session_id) - context_snapshot = json.loads((session_dir / "a2a" / "context.json").read_text(encoding="utf-8")) - assert context_snapshot["active_task_id"] is None + await asyncio.wait_for(finalized.wait(), timeout=1) + await asyncio.gather(*handler._detached_message_producers) @pytest.mark.asyncio -async def test_handler_restores_backup_before_hydrating_omitted_pipeline_task_id( - monkeypatch: pytest.MonkeyPatch, - tmp_path, -) -> None: - monkeypatch.setenv("IAC_CODE_MODE", "pipeline") - monkeypatch.setenv("IAC_CODE_CONFIG_DIR", str(tmp_path / "config")) - monkeypatch.setenv("IAC_CODE_CONFIG_BACKUP_DIR", str(tmp_path / "backup")) - cwd = tmp_path / "workspace" - cwd.mkdir() - context_id = "ctx-restore" - task_id = "task-restore" - store = A2ATaskStore() - ctx = await store.get_or_create_context( - context_id=context_id, - cwd=str(cwd), - runtime_factory=lambda session_id: SimpleNamespace(session_id=session_id), - ) - storage = SessionStorage() - storage.save(str(cwd), ctx.session_id, []) - pending_input = { - "inputId": "input-confirm_and_select-1", - "kind": "candidate_selection", - "prompt": "请选择方案", - "options": [{"name": "方案A", "candidate_index": 0}], - } - pending_event = { - "schemaVersion": "1.0", - "extensionUri": "urn:iac-code:a2a:pipeline-events:v1", - "eventId": "evt-selection", - "sequence": 1, - "createdAt": "2026-06-08T10:00:00Z", - "eventType": "input_required", - "scope": "step", - "pipelineRunId": context_id, - "taskId": task_id, - "contextId": context_id, - "pipelineName": "selling", - "status": "input_required", - "step": {"runId": "step-confirm_and_select-1", "id": "confirm_and_select", "attempt": 1}, - "input": pending_input, - "data": pending_input, - } - pipeline_dir = a2a_pipeline_dir_for_session(cwd=str(cwd), session_id=ctx.session_id) - A2APipelineJournal(pipeline_dir).append(pending_event) - A2APipelineSnapshotStore(pipeline_dir).save(reduce_pipeline_events([pending_event])) - backup_service = SessionBackupService(storage, retry_delays=()) - backup_service.initialize_session(str(cwd), ctx.session_id) - backup_service.backup_session(str(cwd), ctx.session_id, reason=BackupReason.INPUT_REQUIRED, critical=True) - primary_session_dir = storage.session_dir(str(cwd), ctx.session_id) - shutil.rmtree(primary_session_dir) +async def test_detached_permission_disconnect_without_terminal_generation_does_not_finalize() -> None: + call_context = ServerCallContext() + request_context = SimpleNamespace(call_context=call_context) + NaturalCompletionGenerationCarrier.prepare(request_context) + finalized = False - class FakeExecutor: - async def _reconcile_session_before_route(self, *, context_id: str, cwd: str): - assert context_id == "ctx-restore" - return await asyncio.to_thread(backup_service.reconcile_session, cwd, ctx.session_id) + class Executor: + async def finalize_natural_execution(self, **_kwargs) -> None: + nonlocal finalized + finalized = True - handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) - handler.task_store = store - handler.agent_executor = FakeExecutor() - params = SimpleNamespace(message=SimpleNamespace(task_id=None, context_id=context_id)) + async def completed_producer() -> None: + return None - await handler._hydrate_recoverable_pipeline_task_id(params) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = A2ATaskStore() + handler.agent_executor = Executor() + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + completed = object() + queue: asyncio.Queue[object] = asyncio.Queue() + await queue.put(completed) + + await handler._drain_inactive_permission_response( + queue, + asyncio.create_task(completed_producer()), + completed, + params=params, + context=call_context, + ) - assert params.message.task_id == task_id - assert storage.session_dir(str(cwd), ctx.session_id).is_dir() + assert finalized is False @pytest.mark.asyncio -async def test_dispatcher_stream_backpressures_asgi_until_consumer_resumes() -> None: - first_chunk_consumed = asyncio.Event() - - async def app(_scope, _receive, send) -> None: - await send( - { - "type": "http.response.start", - "status": 200, - "headers": [(b"content-type", b"text/event-stream")], - } - ) - await send( - { - "type": "http.response.body", - "body": b'data: {"id":"1","result":{"index":1}}\n\n', - "more_body": True, - } - ) - first_chunk_consumed.set() - await send( - { - "type": "http.response.body", - "body": b'data: {"id":"1","result":{"index":2}}\n\n', - "more_body": False, - } - ) +async def test_detached_permission_failed_producer_does_not_finalize() -> None: + finalized = False - dispatcher = A2AJsonRpcDispatcher(SimpleNamespace(app=app)) - stream = dispatcher.dispatch_stream({"jsonrpc": "2.0", "id": "1"}) + class Executor: + async def finalize_natural_execution(self, **_kwargs) -> None: + nonlocal finalized + finalized = True - first = await anext(stream) - assert first["result"]["index"] == 1 - assert first_chunk_consumed.is_set() is False + async def failed_producer() -> None: + raise RuntimeError("continuation failed") - second = await anext(stream) - assert first_chunk_consumed.is_set() is True - assert second["result"]["index"] == 2 - with pytest.raises(StopAsyncIteration): - await anext(stream) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = A2ATaskStore() + handler.agent_executor = Executor() + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + completed = object() + queue: asyncio.Queue[object] = asyncio.Queue() + await queue.put(completed) + + await handler._drain_inactive_permission_response( + queue, + asyncio.create_task(failed_producer()), + completed, + params=params, + context=ServerCallContext(), + ) - await dispatcher.aclose() + assert finalized is False @pytest.mark.asyncio -async def test_streaming_asgi_transport_cancels_app_before_response_start() -> None: - app_started = asyncio.Event() - app_cancelled = asyncio.Event() +async def test_cancelled_detached_permission_drain_cancels_producer_without_finalizing() -> None: + finalized = False + producer_cancelled = asyncio.Event() - async def app(_scope, _receive, _send) -> None: - app_started.set() + class Executor: + async def finalize_natural_execution(self, **_kwargs) -> None: + nonlocal finalized + finalized = True + + async def blocked_producer() -> None: try: await asyncio.Event().wait() except asyncio.CancelledError: - app_cancelled.set() + producer_cancelled.set() raise - transport = _StreamingASGITransport(app) - async with httpx.AsyncClient(transport=transport, base_url="http://transport.local") as client: - request_task = asyncio.create_task(client.get("/")) - await app_started.wait() - - request_task.cancel() - with pytest.raises(asyncio.CancelledError): - await request_task + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = A2ATaskStore() + handler.agent_executor = Executor() + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) + completed = object() + queue: asyncio.Queue[object] = asyncio.Queue() + producer = asyncio.create_task(blocked_producer()) + drain = asyncio.create_task( + handler._drain_inactive_permission_response( + queue, + producer, + completed, + params=params, + context=ServerCallContext(), + ) + ) + await asyncio.sleep(0) + drain.cancel() - assert app_cancelled.is_set() - assert not any(task.get_name() == "a2a-streaming-asgi-dispatch" and not task.done() for task in asyncio.all_tasks()) + with pytest.raises(asyncio.CancelledError): + await drain + assert producer_cancelled.is_set() + assert finalized is False @pytest.mark.asyncio -async def test_message_stream_acknowledges_transport_delivery_only_when_resumed(monkeypatch) -> None: +async def test_message_stream_does_not_acknowledge_transport_delivery_when_closed(monkeypatch) -> None: observed: dict[str, asyncio.Future[None]] = {} stages: list[str] = [] update = TaskStatusUpdateEvent( @@ -404,21 +2281,17 @@ async def hydrate(_params) -> None: stream = handler.on_message_send_stream(params, object()) assert await anext(stream) is update - assert observed["completion"].done() is False - assert stages == ["registered", "dequeued"] - - with pytest.raises(StopAsyncIteration): - await anext(stream) + await stream.aclose() - assert observed["completion"].done() is True - assert stages == ["registered", "dequeued", "acknowledged"] + assert isinstance(observed["completion"].exception(), PipelineTransportDeliveryClosedError) + assert stages == ["registered", "dequeued", "closed"] assert pipeline_transport_delivery_tracking_enabled() is False @pytest.mark.asyncio -async def test_message_stream_does_not_acknowledge_transport_delivery_when_closed(monkeypatch) -> None: - observed: dict[str, asyncio.Future[None]] = {} - stages: list[str] = [] +async def test_pipeline_message_stream_does_not_bind_subscriber_delivery_to_producer(monkeypatch) -> None: + monkeypatch.setenv("IAC_CODE_MODE", "pipeline") + tracking_states = [] update = TaskStatusUpdateEvent( task_id="task-1", context_id="ctx-1", @@ -426,10 +2299,7 @@ async def test_message_stream_does_not_acknowledge_transport_delivery_when_close ) async def sdk_stream(_handler, _params, _context): - observed["completion"] = register_pipeline_transport_delivery( - update, - stage_observer=lambda stage, _at_ns: stages.append(stage), - ) + tracking_states.append(pipeline_transport_delivery_tracking_enabled()) yield update async def hydrate(_params) -> None: @@ -447,42 +2317,65 @@ async def hydrate(_params) -> None: assert await anext(stream) is update await stream.aclose() - assert isinstance(observed["completion"].exception(), PipelineTransportDeliveryClosedError) - assert stages == ["registered", "dequeued", "closed"] + assert tracking_states == [False] assert pipeline_transport_delivery_tracking_enabled() is False @pytest.mark.asyncio -async def test_pipeline_message_stream_does_not_bind_subscriber_delivery_to_producer(monkeypatch) -> None: - monkeypatch.setenv("IAC_CODE_MODE", "pipeline") - tracking_states = [] +async def test_message_stream_retires_stale_sdk_lifecycle_after_durable_recovery(monkeypatch) -> None: + call_context = ServerCallContext() + store = A2ATaskStore() + await store.save( + Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) update = TaskStatusUpdateEvent( task_id="task-1", context_id="ctx-1", status=TaskStatus(state=TaskState.TASK_STATE_WORKING), ) + retired: list[str] = [] async def sdk_stream(_handler, _params, _context): - tracking_states.append(pipeline_transport_delivery_tracking_enabled()) yield update async def hydrate(_params) -> None: return None + async def reconcile(_params, _context) -> str: + return "recovery-1" + + class RecoveryAwareRegistry: + async def reconcile_and_replace_for_recovery(self, task_id: str, *, acquire_admission, **_kwargs) -> str | None: + admission = await acquire_admission() + if admission is not None: + retired.append(task_id) + return admission + + async def cancel_recovery_reservation(self, _task_id: str, _admission: str) -> None: + return None + + async def get(self, _task_id: str): + return None + monkeypatch.setattr(DefaultRequestHandler, "on_message_send_stream", sdk_stream) handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) - handler.task_store = object() + handler.task_store = store + handler._active_task_registry = RecoveryAwareRegistry() handler._validate_extensions = lambda _context: None handler._validate_pipeline_message_request = lambda _params: None handler._hydrate_recoverable_pipeline_task_id = hydrate - params = SimpleNamespace(message=SimpleNamespace(task_id=None)) + handler._reconcile_recoverable_pipeline_task = reconcile + params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) - stream = handler.on_message_send_stream(params, object()) - assert await anext(stream) is update - await stream.aclose() + events = await _collect_async(handler.on_message_send_stream(params, call_context)) - assert tracking_states == [False] - assert pipeline_transport_delivery_tracking_enabled() is False + assert events == [update] + assert retired == ["task-1"] @pytest.mark.asyncio @@ -536,10 +2429,102 @@ async def fail_active_stream(*_args, **_kwargs): handler._on_active_message_send_stream = fail_active_stream params = SimpleNamespace(message=SimpleNamespace(task_id="task-1", context_id="ctx-1")) - events = await _collect_async(handler.on_message_send_stream(params, call_context)) + events = await _collect_async(handler.on_message_send_stream(params, call_context)) + + assert events == [update] + assert sdk_stream_called is True + + +@pytest.mark.asyncio +async def test_input_required_base_stream_rebinds_publisher_after_sdk_lifecycle_finished() -> None: + from iac_code.a2a import pipeline_executor as pipeline_executor_module + + call_context = ServerCallContext() + store = A2ATaskStore() + await store.save( + Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_INPUT_REQUIRED), + ), + call_context, + ) + owner_release = asyncio.Event() + domain_owner = asyncio.create_task(owner_release.wait()) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.active_task = domain_owner + stale_queue = FakeEventQueue() + runtime = pipeline_executor_module.A2APipelineRuntime( + agent_runtime=SimpleNamespace(), + publisher=SimpleNamespace(event_queue=stale_queue), + ) + update = TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + observed: dict[str, object] = {} + + class Executor: + async def execute(self, request_context, event_queue) -> None: + gate = DirectPipelineRouteGateCarrier.read(request_context) + observed["gate"] = gate + observed["registered"] = await pipeline_executor_module._register_active_interrupt( + runtime, + event_queue=event_queue, + direct_route_gate=gate, + bind_publisher_event_queue=PipelineLifecycleEventQueueCarrier.read(request_context), + ) + try: + await runtime.publisher.event_queue.enqueue_event(update) + finally: + await pipeline_executor_module._settle_active_interrupt_safely(runtime) + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + handler = IacCodeRequestHandler( + agent_executor=Executor(), + task_store=store, + agent_card=SimpleNamespace(capabilities=SimpleNamespace(streaming=True, extensions=[])), + ) + + async def hydrate(_params) -> None: + return None + + async def reconcile(_params, _context) -> None: + return None + + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + assert await handler._active_task_registry.get("task-1") is None + message = Message( + message_id="message-1", + task_id="task-1", + context_id="ctx-1", + role=Role.ROLE_USER, + parts=[Part(text='{"selected_candidate_index": 0}')], + ) + ParseDict({"iac_code": {"run_mode": "pipeline"}}, message.metadata) + + try: + events = await asyncio.wait_for( + _collect_async( + handler.on_message_send_stream( + SendMessageRequest(message=message), + call_context, + ) + ), + timeout=_STREAM_TEST_TIMEOUT, + ) + finally: + owner_release.set() + await domain_owner + await handler._active_task_registry.retire_for_recovery("task-1") + assert observed == {"gate": None, "registered": True} assert events == [update] - assert sdk_stream_called is True + assert stale_queue.events == [] @pytest.mark.asyncio @@ -614,6 +2599,638 @@ async def fail_sdk_stream(*_args, **_kwargs): assert message.task_id == "task-1" +@pytest.mark.asyncio +async def test_active_message_route_ignores_old_terminal_events_around_its_request_boundary() -> None: + call_context = ServerCallContext() + store = A2ATaskStore() + task = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + await store.save(task, call_context) + record = await store.get_or_create_task(task_id=task.id, context_id=task.context_id) + record.active_task = asyncio.current_task() + + old_terminal = Task( + id=task.id, + context_id=task.context_id, + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + current_update = TaskStatusUpdateEvent( + task_id=task.id, + context_id=task.context_id, + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + subscribers = EventQueueSource(create_default_sink=False) + + class ForwardingAgentQueue: + def __init__(self) -> None: + self.boundary_enqueued = False + + async def enqueue_event(self, event) -> None: + if not self.boundary_enqueued: + self.boundary_enqueued = True + await subscribers.enqueue_event((old_terminal, task)) + await subscribers.enqueue_event((event, None)) + await subscribers.enqueue_event((old_terminal, task)) + return + await subscribers.enqueue_event((event, task)) + + async def test_only_join_incoming_queue(self) -> None: + await subscribers.test_only_join_incoming_queue() + + class ActiveTask: + def __init__(self) -> None: + self.task_id = task.id + self.direct_message_lock = asyncio.Lock() + self._lock = asyncio.Lock() + self._is_finished = asyncio.Event() + self._reference_count = 0 + self._event_queue_agent = ForwardingAgentQueue() + self._event_queue_subscribers = subscribers + + async def _maybe_cleanup(self) -> None: + return None + + active_task = ActiveTask() + + class ActiveTaskRegistry: + async def get(self, _task_id): + return active_task + + class RequestContextBuilder: + async def build(self, **_kwargs): + return SimpleNamespace() + + class AgentExecutor: + async def execute(self, _request_context, event_queue) -> None: + await event_queue.enqueue_event(current_update) + + async def hydrate(_params) -> None: + return None + + async def reconcile(_params, _context) -> None: + return None + + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler.agent_executor = AgentExecutor() + handler._request_context_builder = RequestContextBuilder() + handler._active_task_registry = ActiveTaskRegistry() + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + message = Message( + message_id="message-1", + task_id=task.id, + context_id=task.context_id, + role=Role.ROLE_USER, + parts=[Part(text="continue")], + ) + params = SimpleNamespace(message=message, configuration=None) + + try: + events = await asyncio.wait_for( + _collect_async(handler.on_message_send_stream(params, call_context)), + timeout=_STREAM_TEST_TIMEOUT, + ) + finally: + await subscribers.close(immediate=True) + + assert events == [current_update] + assert active_task._reference_count == 0 + assert not active_task.direct_message_lock.locked() + + +@pytest.mark.asyncio +async def test_active_pipeline_reentry_delivers_events_after_sdk_lifecycle_replacement() -> None: + from iac_code.a2a import pipeline_executor as pipeline_executor_module + + call_context = ServerCallContext() + store = A2ATaskStore() + task = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + await store.save(task, call_context) + record = await store.get_or_create_task(task_id=task.id, context_id=task.context_id) + record.active_task = asyncio.current_task() + update = TaskStatusUpdateEvent( + task_id=task.id, + context_id=task.context_id, + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + subscribers = EventQueueSource(create_default_sink=False) + stale_queue = FakeEventQueue() + + class ForwardingAgentQueue: + async def enqueue_event(self, event) -> None: + await subscribers.enqueue_event((event, task)) + + async def test_only_join_incoming_queue(self) -> None: + await subscribers.test_only_join_incoming_queue() + + class ActiveTask: + def __init__(self) -> None: + self.task_id = task.id + self.direct_message_lock = asyncio.Lock() + self._lock = asyncio.Lock() + self._is_finished = asyncio.Event() + self._reference_count = 0 + self._event_queue_agent = ForwardingAgentQueue() + self._event_queue_subscribers = subscribers + + async def _maybe_cleanup(self) -> None: + return None + + active_task = ActiveTask() + + class ActiveTaskRegistry: + async def get(self, _task_id): + return active_task + + class RequestContextBuilder: + async def build(self, **_kwargs): + return SimpleNamespace() + + runtime = pipeline_executor_module.A2APipelineRuntime( + agent_runtime=SimpleNamespace(), + publisher=SimpleNamespace(event_queue=stale_queue), + ) + + class AgentExecutor: + async def execute(self, request_context, event_queue) -> None: + gate = DirectPipelineRouteGateCarrier.read(request_context) + assert gate is not None + assert await pipeline_executor_module._register_active_interrupt( + runtime, + event_queue=event_queue, + direct_route_gate=gate, + ) + try: + await runtime.publisher.event_queue.enqueue_event(update) + finally: + await pipeline_executor_module._settle_active_interrupt_safely(runtime) + + async def hydrate(_params) -> None: + return None + + async def reconcile(_params, _context) -> None: + return None + + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler.agent_executor = AgentExecutor() + handler._request_context_builder = RequestContextBuilder() + handler._active_task_registry = ActiveTaskRegistry() + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + message = Message( + message_id="message-1", + task_id=task.id, + context_id=task.context_id, + role=Role.ROLE_USER, + parts=[Part(text='{"selected_candidate_index": 0}')], + ) + ParseDict({"iac_code": {"run_mode": "pipeline"}}, message.metadata) + + try: + events = await asyncio.wait_for( + _collect_async( + handler.on_message_send_stream( + SimpleNamespace(message=message, configuration=None), + call_context, + ) + ), + timeout=_STREAM_TEST_TIMEOUT, + ) + finally: + await subscribers.close(immediate=True) + + assert events == [update] + assert stale_queue.events == [] + assert active_task._reference_count == 0 + assert not active_task.direct_message_lock.locked() + + +@pytest.mark.asyncio +async def test_terminal_winning_direct_route_recovers_same_request_without_old_terminal() -> None: + def status(state: TaskState.ValueType, seconds: int) -> TaskStatus: + value = TaskStatus(state=state) + value.timestamp.seconds = seconds + return value + + class RoutingExecutor: + def __init__(self, owner_release: asyncio.Event) -> None: + self.calls = 0 + self.waited = False + self._owner_release = owner_release + + async def execute(self, request_context, event_queue) -> None: + self.calls += 1 + gate = DirectPipelineRouteGateCarrier.read(request_context) + if gate is not None: + gate.require_recovery() + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=status(TaskState.TASK_STATE_CANCELED, 150), + ) + ) + return + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=status(TaskState.TASK_STATE_WORKING, 300), + ) + ) + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + async def wait_until_recoverable_pipeline_input(self, *, context_id: str, task_id: str) -> None: + assert (context_id, task_id) == ("ctx-1", "task-1") + self.waited = True + self._owner_release.set() + + call_context = ServerCallContext() + store = A2ATaskStore() + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_INPUT_REQUIRED, 200)), + call_context, + ) + owner_release = asyncio.Event() + owner = asyncio.create_task(owner_release.wait()) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.active_task = owner + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_WORKING, 250)), + call_context, + ) + executor = RoutingExecutor(owner_release) + handler = IacCodeRequestHandler( + agent_executor=executor, + task_store=store, + agent_card=SimpleNamespace(capabilities=SimpleNamespace(streaming=True)), + ) + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + + async def hydrate(_params) -> None: + return None + + reconcile_calls = 0 + + async def reconcile(_params, _context) -> str | None: + nonlocal reconcile_calls + reconcile_calls += 1 + if reconcile_calls == 1: + return None + owner_release.set() + await owner + record.active_task = None + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_INPUT_REQUIRED, 260)), + call_context, + ) + return "recovery-1" + + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + old_lifecycle = await handler._active_task_registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + message = Message( + message_id="message-1", + task_id="task-1", + context_id="ctx-1", + role=Role.ROLE_USER, + parts=[Part(text="select candidate 0")], + ) + ParseDict({"iac_code": {"run_mode": "pipeline"}}, message.metadata) + + try: + events = await asyncio.wait_for( + _collect_async( + handler.on_message_send_stream( + SendMessageRequest(message=message), + call_context, + ) + ), + timeout=_STREAM_TEST_TIMEOUT, + ) + replacement = await handler._active_task_registry.get("task-1") + finally: + owner_release.set() + await asyncio.gather(owner, return_exceptions=True) + await handler._active_task_registry.retire_for_recovery("task-1") + + assert [event.status.state for event in events] == [TaskState.TASK_STATE_WORKING] + assert executor.calls == 2 + assert executor.waited is True + assert replacement is not None and replacement is not old_lifecycle + assert old_lifecycle._is_finished.is_set() + + +@pytest.mark.asyncio +async def test_cancelled_direct_stream_finishes_recovery_for_the_same_pipeline_request() -> None: + def status(state: TaskState.ValueType, seconds: int) -> TaskStatus: + value = TaskStatus(state=state) + value.timestamp.seconds = seconds + return value + + direct_started = asyncio.Event() + release_direct = asyncio.Event() + recovered = asyncio.Event() + owner_release = asyncio.Event() + outer_cleanup_started = asyncio.Event() + outer_cleanup_finished = asyncio.Event() + detached_admission_staged = asyncio.Event() + admission_released = asyncio.Event() + + class RoutingExecutor: + async def execute(self, request_context, event_queue) -> None: + gate = DirectPipelineRouteGateCarrier.read(request_context) + if gate is not None: + direct_started.set() + await release_direct.wait() + gate.require_recovery() + return + recovered.set() + await event_queue.enqueue_event( + TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=status(TaskState.TASK_STATE_WORKING, 300), + ) + ) + + async def cancel(self, _request_context, _event_queue) -> None: + return None + + async def wait_until_recoverable_pipeline_input(self, *, context_id: str, task_id: str) -> None: + assert (context_id, task_id) == ("ctx-1", "task-1") + owner_release.set() + + call_context = ServerCallContext( + state={"ordinary": {"value": 1}}, + tenant="tenant-1", + requested_extensions={"urn:test:extension"}, + ) + store = A2ATaskStore() + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_INPUT_REQUIRED, 200)), + call_context, + ) + owner = asyncio.create_task(owner_release.wait()) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.active_task = owner + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_WORKING, 250)), + call_context, + ) + handler = IacCodeRequestHandler( + agent_executor=RoutingExecutor(), + task_store=store, + agent_card=SimpleNamespace(capabilities=SimpleNamespace(streaming=True)), + ) + handler._validate_extensions = lambda _context: None + handler._validate_pipeline_message_request = lambda _params: None + + async def hydrate(_params) -> None: + return None + + reconcile_calls = 0 + + async def reconcile(_params, _context) -> str | None: + nonlocal reconcile_calls + reconcile_calls += 1 + if reconcile_calls == 1: + return None + owner_release.set() + await owner + record.active_task = None + await store.save( + Task(id="task-1", context_id="ctx-1", status=status(TaskState.TASK_STATE_INPUT_REQUIRED, 260)), + call_context, + ) + return "recovery-1" + + handler._hydrate_recoverable_pipeline_task_id = hydrate + handler._reconcile_recoverable_pipeline_task = reconcile + release_contexts: list[object] = [] + release_recovery = handler._release_untransferred_recovery + staged_contexts: list[ServerCallContext] = [] + stage_recovery = handler._stage_recoverable_input_admission + acknowledged: list[tuple[object, str, bool]] = [] + acknowledge_enqueue = handler._acknowledge_recoverable_input_enqueue + setup_active_task = handler._setup_active_task + released_admissions: list[str] = [] + + def track_stage(stage_context, admission) -> None: + stage_recovery(stage_context, admission) + if admission == "recovery-1": + staged_contexts.append(stage_context) + detached_admission_staged.set() + + def track_acknowledge(ack_context, admission) -> bool: + result = acknowledge_enqueue(ack_context, admission) + acknowledged.append((ack_context, admission, result)) + return result + + async def synchronize_setup(setup_params, setup_context): + await outer_cleanup_finished.wait() + return await setup_active_task(setup_params, setup_context) + + async def track_admission_release(admission: str | None) -> None: + if admission is not None: + released_admissions.append(admission) + admission_released.set() + + async def track_release(release_params, release_context) -> None: + release_contexts.append(release_context) + if release_context is call_context: + outer_cleanup_started.set() + await detached_admission_staged.wait() + await release_recovery(release_params, release_context) + outer_cleanup_finished.set() + return + await release_recovery(release_params, release_context) + + handler._stage_recoverable_input_admission = track_stage + handler._acknowledge_recoverable_input_enqueue = track_acknowledge + handler._setup_active_task = synchronize_setup + handler._release_untransferred_recovery = track_release + store.release_recoverable_input_admission = track_admission_release + old_lifecycle = await handler._active_task_registry.get_or_create( + "task-1", + call_context=call_context, + context_id="ctx-1", + create_task_if_missing=True, + ) + message = Message( + message_id="message-1", + task_id="task-1", + context_id="ctx-1", + role=Role.ROLE_USER, + parts=[Part(text="select candidate 0")], + ) + ParseDict({"iac_code": {"run_mode": "pipeline"}}, message.metadata) + stream_task = asyncio.create_task( + _collect_async(handler.on_message_send_stream(SendMessageRequest(message=message), call_context)) + ) + + try: + await asyncio.wait_for(direct_started.wait(), timeout=_STREAM_TEST_TIMEOUT) + stream_task.cancel() + await asyncio.wait_for(outer_cleanup_started.wait(), timeout=_STREAM_TEST_TIMEOUT) + release_direct.set() + await asyncio.wait_for(detached_admission_staged.wait(), timeout=_STREAM_TEST_TIMEOUT) + await asyncio.wait_for(outer_cleanup_finished.wait(), timeout=_STREAM_TEST_TIMEOUT) + with pytest.raises(asyncio.CancelledError): + await stream_task + await asyncio.wait_for(recovered.wait(), timeout=_STREAM_TEST_TIMEOUT) + replacement = await handler._active_task_registry.get("task-1") + assert replacement is not None and replacement is not old_lifecycle + finally: + release_direct.set() + owner_release.set() + await asyncio.gather(owner, return_exceptions=True) + detached = tuple(handler._detached_message_producers) + if detached: + await asyncio.wait_for(asyncio.gather(*detached), timeout=_STREAM_TEST_TIMEOUT) + await handler._active_task_registry.retire_for_recovery("task-1") + await asyncio.wait_for(admission_released.wait(), timeout=_STREAM_TEST_TIMEOUT) + + assert call_context in release_contexts + assert len(staged_contexts) == 1 + detached_context = staged_contexts[0] + assert detached_context is not call_context + assert detached_context.state is not call_context.state + assert detached_context.state["ordinary"] == {"value": 1} + assert detached_context.user is call_context.user + assert detached_context.tenant == "tenant-1" + assert detached_context.requested_extensions == {"urn:test:extension"} + assert detached_context.requested_extensions is not call_context.requested_extensions + assert acknowledged == [(detached_context, "recovery-1", True)] + assert released_admissions == ["recovery-1"] + + +@pytest.mark.asyncio +async def test_active_message_stream_serializes_direct_requests() -> None: + task = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + updates = [ + TaskStatusUpdateEvent( + task_id=task.id, + context_id=task.context_id, + status=TaskStatus(state=state), + ) + for state in (TaskState.TASK_STATE_WORKING, TaskState.TASK_STATE_INPUT_REQUIRED) + ] + started = [asyncio.Event(), asyncio.Event()] + releases = [asyncio.Event(), asyncio.Event()] + subscribers = EventQueueSource(create_default_sink=False) + + class ForwardingAgentQueue: + async def enqueue_event(self, event) -> None: + await subscribers.enqueue_event((event, task)) + + async def test_only_join_incoming_queue(self) -> None: + await subscribers.test_only_join_incoming_queue() + + class ActiveTask: + def __init__(self) -> None: + self.task_id = task.id + self.direct_message_lock = asyncio.Lock() + self._lock = asyncio.Lock() + self._is_finished = asyncio.Event() + self._reference_count = 0 + self._event_queue_agent = ForwardingAgentQueue() + self._event_queue_subscribers = subscribers + + async def _maybe_cleanup(self) -> None: + return None + + class RequestContextBuilder: + async def build(self, *, params, **_kwargs): + return SimpleNamespace(index=int(params.message.message_id[-1])) + + class AgentExecutor: + def __init__(self) -> None: + self.running = 0 + self.max_running = 0 + + async def execute(self, request_context, event_queue) -> None: + index = request_context.index + self.running += 1 + self.max_running = max(self.max_running, self.running) + started[index].set() + try: + await releases[index].wait() + await event_queue.enqueue_event(updates[index]) + finally: + self.running -= 1 + + executor = AgentExecutor() + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.agent_executor = executor + handler._request_context_builder = RequestContextBuilder() + active_task = ActiveTask() + params = [ + SimpleNamespace( + message=SimpleNamespace(message_id=f"message-{index}", context_id=task.context_id), + configuration=None, + ) + for index in range(2) + ] + first_consumer = asyncio.create_task( + _collect_async(handler._on_active_message_send_stream(params[0], object(), task=task, active_task=active_task)) + ) + consumers = [first_consumer] + + try: + await asyncio.wait_for(started[0].wait(), timeout=_STREAM_TEST_TIMEOUT) + second_consumer = asyncio.create_task( + _collect_async( + handler._on_active_message_send_stream(params[1], object(), task=task, active_task=active_task) + ) + ) + consumers.append(second_consumer) + await asyncio.sleep(0) + assert not started[1].is_set() + releases[0].set() + await asyncio.wait_for(started[1].wait(), timeout=_STREAM_TEST_TIMEOUT) + releases[1].set() + events = await asyncio.wait_for(asyncio.gather(*consumers), timeout=_STREAM_TEST_TIMEOUT) + finally: + for release in releases: + release.set() + for consumer in consumers: + if not consumer.done(): + consumer.cancel() + await asyncio.gather(*consumers, return_exceptions=True) + await subscribers.close(immediate=True) + + assert events == [[updates[0]], [updates[1]]] + assert executor.max_running == 1 + assert active_task._reference_count == 0 + assert not active_task.direct_message_lock.locked() + + @pytest.mark.asyncio async def test_text_gateway_sideband_permission_response_hydrates_task_and_returns_short_ack(monkeypatch) -> None: call_context = ServerCallContext() @@ -1232,7 +3849,7 @@ async def hanging_sdk_subscription(self, params, context): events = await asyncio.wait_for( _collect_async(handler.on_subscribe_to_task(SubscribeToTaskRequest(id="task-1"), call_context)), - timeout=0.5, + timeout=_STREAM_TEST_TIMEOUT, ) assert [event.status.state for event in events] == [ @@ -1488,12 +4105,17 @@ async def tap(self) -> FakeTappedQueue: class FakeActiveTask: def __init__(self) -> None: self.task_id = "task-1" + self.direct_message_lock = asyncio.Lock() self._lock = asyncio.Lock() self._is_finished = asyncio.Event() self._reference_count = 0 - self._event_queue_agent = SimpleNamespace() + self._event_queue_agent = SimpleNamespace(enqueue_event=self._enqueue_event) self._event_queue_subscribers = FakeSubscribers(FakeTappedQueue()) + @staticmethod + async def _enqueue_event(_event) -> None: + return None + async def _maybe_cleanup(self) -> None: return None @@ -1523,6 +4145,7 @@ async def consume() -> None: await asyncio.wait_for(producer_cancelled.wait(), timeout=_STREAM_TEST_TIMEOUT) await asyncio.gather(*cleanup_tasks, return_exceptions=True) assert active_task._reference_count == 0 + assert not active_task.direct_message_lock.locked() @pytest.mark.asyncio diff --git a/tests/a2a_e2e/test_execution_control_scenarios.py b/tests/a2a_e2e/test_execution_control_scenarios.py index 1c7378b1..5f2aadfa 100644 --- a/tests/a2a_e2e/test_execution_control_scenarios.py +++ b/tests/a2a_e2e/test_execution_control_scenarios.py @@ -9,6 +9,16 @@ import pytest +# A scenario includes process startup, control round trips (including real +# disconnect-deadline checks), and durable completion. These phases share the +# overall watchdog; keep the per-step scenario wait budget at 20 seconds. +# In particular, Windows pipeline setup can consume most of the old 25s budget +# before the resumed tool is released. Leave time for real backup/stream drain. +WAIT_TIMEOUT = 20 +OVERALL_TIMEOUT = 3 * WAIT_TIMEOUT +PROCESS_TIMEOUT = OVERALL_TIMEOUT + 10 # Allow the runner to clean up and write audits. +TEST_TIMEOUT = PROCESS_TIMEOUT + 10 + CASES = ( ("warm-resume-pausing", "normal"), ("warm-resume-pausing", "pipeline"), @@ -33,7 +43,7 @@ ) @pytest.mark.integration -@pytest.mark.timeout(40) +@pytest.mark.timeout(TEST_TIMEOUT) @pytest.mark.parametrize(("scenario", "mode"), CASES, ids=["{}-{}".format(*case) for case in CASES]) def test_execution_control_scenario(tmp_path: Path, scenario: str, mode: str) -> None: repo_root = Path(__file__).resolve().parents[2] @@ -50,14 +60,14 @@ def test_execution_control_scenario(tmp_path: Path, scenario: str, mode: str) -> "--mode", mode, "--timeout", - "20", + str(WAIT_TIMEOUT), "--overall-timeout", - "25", + str(OVERALL_TIMEOUT), ], cwd=repo_root, text=True, capture_output=True, - timeout=35, + timeout=PROCESS_TIMEOUT, check=False, ) summary_path = run_dir / "summary.json" @@ -70,7 +80,7 @@ def test_execution_control_scenario(tmp_path: Path, scenario: str, mode: str) -> (run_dir / "runner-error.txt").read_text(encoding="utf-8") if (run_dir / "runner-error.txt").exists() else "", ) assert summary is not None and summary["status"] == "passed" - assert 0 < summary["elapsedSeconds"] < 35 + assert 0 < summary["elapsedSeconds"] < PROCESS_TIMEOUT assert summary["server"]["returnCode"] is not None assert summary["server"]["forcedKill"] is False for artifact in ( diff --git a/tests/mcp/test_client.py b/tests/mcp/test_client.py index 55968d70..5d86eec2 100644 --- a/tests/mcp/test_client.py +++ b/tests/mcp/test_client.py @@ -532,9 +532,11 @@ async def test_remote_headers_helper_cleans_up_process_on_cancellation(tmp_path: import time pid_file = sys.argv[1] -with open(pid_file, "w", encoding="utf-8") as handle: +temporary_pid_file = f"{pid_file}.tmp" +with open(temporary_pid_file, "w", encoding="utf-8") as handle: handle.write(str(os.getpid())) handle.flush() +os.replace(temporary_pid_file, pid_file) time.sleep(10) """, ) diff --git a/tests/skill_bridge/test_alicloud_ros_agent_bridge.py b/tests/skill_bridge/test_alicloud_ros_agent_bridge.py index be3773e4..48ba83f6 100644 --- a/tests/skill_bridge/test_alicloud_ros_agent_bridge.py +++ b/tests/skill_bridge/test_alicloud_ros_agent_bridge.py @@ -2584,13 +2584,26 @@ def test_manager_idle_countdown_starts_after_sse_worker_exits(monkeypatch, tmp_p monkeypatch.setenv(bridge.STATE_DIR_ENV, str(tmp_path / "state")) workspace = tmp_path / "workspace" workspace.mkdir() - # Invoke the current interpreter as the fake CLI and let it execute the - # positional ``ros`` script from the worker cwd. This avoids depending on - # Windows batch-file launch behavior in a manager lifecycle test. - fake_cli = Path(sys.executable) + # This stdlib-only process test needs Popen.pid to identify the interpreter + # itself. Windows venv python.exe is a redirector with a different PID from + # the worker it launches; use the base interpreter for the whole process tree. + python_executable = getattr(sys, "_base_executable", None) or sys.executable + monkeypatch.setattr(bridge, "sys", SimpleNamespace(executable=python_executable)) + monkeypatch.delenv("__PYVENV_LAUNCHER__", raising=False) + fake_cli = Path(python_executable) + worker_pid_path = workspace / "worker-pid" + release_worker = workspace / "release-worker" (workspace / "ros").write_text( - "import json, time\n" - + "time.sleep(0.6)\n" + "import json, os, time\n" + + "from pathlib import Path\n" + + "pid_path = Path({!r})\n".format(str(worker_pid_path)) + + "pid_path.with_suffix('.tmp').write_text(str(os.getppid()), encoding='utf-8')\n" + + "pid_path.with_suffix('.tmp').replace(pid_path)\n" + + "release = Path({!r})\n".format(str(release_worker)) + + "deadline = time.monotonic() + 25\n" + + "while not release.exists():\n" + + " assert time.monotonic() < deadline, 'test did not release worker'\n" + + " time.sleep(0.05)\n" + "event = {'result': {'statusUpdate': {'taskId': 'task-1', 'contextId': 'session-1', " + "'status': {'state': 'TASK_STATE_INPUT_REQUIRED', 'message': {'role': 'ROLE_AGENT', " + "'parts': [{'text': 'done'}]}}, 'metadata': {'iac_code': {'assistantFinal': " @@ -2601,36 +2614,64 @@ def test_manager_idle_countdown_starts_after_sse_worker_exits(monkeypatch, tmp_p # Leave enough scheduling headroom for a loaded Windows xdist runner; this # test is about when the idle countdown starts, not sub-second timing. - manager = bridge.ensure_manager(5.0) - started = bridge._manager_request( - manager, - "/start", - { - "workspace": str(workspace), - "prompt": "explain VPC", - "mode": "normal", - "transport": "aliyun_cli", - "endpoint": "ros.aliyuncs.com", - "regionId": "cn-hangzhou", - "aliyunPath": str(fake_cli), - }, - ) - time.sleep(0.35) - assert bridge._pid_alive(manager["pid"]) - - _root, job_path, _spool = bridge._job_paths(started["jobId"]) - deadline = time.monotonic() + 3 - while time.monotonic() < deadline: - if not isinstance(bridge._load_state_json(job_path).get("workerPid"), int): - break - time.sleep(0.02) - assert bridge._load_state_json(job_path)["state"] == "turn-completed" - assert bridge._pid_alive(manager["pid"]) - time.sleep(0.08) - assert bridge._pid_alive(manager["pid"]) - - _wait_for_pid_exit(manager["pid"], timeout=8.0) - assert not bridge._pid_alive(manager["pid"]) + idle_seconds = 5.0 + manager = bridge.ensure_manager(idle_seconds) + try: + started = bridge._manager_request( + manager, + "/start", + { + "workspace": str(workspace), + "prompt": "explain VPC", + "mode": "normal", + "transport": "aliyun_cli", + "endpoint": "ros.aliyuncs.com", + "regionId": "cn-hangzhou", + "aliyunPath": str(fake_cli), + }, + ) + # /start returns only after workerPid has been committed. A fixed fake-CLI + # sleep can finish before that commit on Windows and leave a stale PID. + # Hold the registered worker beyond the idle threshold to test the actual + # contract: active work prevents manager idle shutdown. + root, job_path, _spool = bridge._job_paths(started["jobId"]) + with bridge.StateLock(root / ".job.lock"): + assert bridge._load_state_json(job_path)["workerPid"] == started["workerPid"] + deadline = time.monotonic() + 10 + while not worker_pid_path.exists(): + assert time.monotonic() < deadline, "Fake CLI did not start" + time.sleep(0.05) + assert int(worker_pid_path.read_text(encoding="utf-8")) == started["workerPid"] + time.sleep(idle_seconds + 0.5) + assert bridge._pid_alive(started["workerPid"]) + assert bridge._pid_alive(manager["pid"]) + release_worker.touch() + + deadline = time.monotonic() + 10 + while True: + # Use the worker's cross-process lock: an unlocked read can race + # atomic replacement of job.json and fail with PermissionError on Windows. + with bridge.StateLock(root / ".job.lock"): + job = bridge._load_state_json(job_path) + if not isinstance(job.get("workerPid"), int): + break + assert time.monotonic() < deadline, { + "state": job.get("state"), + "workerPid": job.get("workerPid"), + "workerStartedAt": job.get("workerStartedAt"), + "workerExitedAt": job.get("workerExitedAt"), + } + time.sleep(0.02) + assert job["state"] == "turn-completed" + assert bridge._pid_alive(manager["pid"]) + time.sleep(0.08) + assert bridge._pid_alive(manager["pid"]) + + _wait_for_pid_exit(manager["pid"], timeout=8.0) + assert not bridge._pid_alive(manager["pid"]) + finally: + release_worker.touch() + _wait_for_pid_exit(manager["pid"], timeout=8.0) def test_manager_failed_start_removes_record_and_terminates_spawn(monkeypatch, tmp_path: Path) -> None: diff --git a/tests/web/test_dynamic_mcp_commands.py b/tests/web/test_dynamic_mcp_commands.py index e8e35d0e..f38d170f 100644 --- a/tests/web/test_dynamic_mcp_commands.py +++ b/tests/web/test_dynamic_mcp_commands.py @@ -5,6 +5,8 @@ import httpx import pytest +from iac_code.web.session_manager import WebTurnAdmissionLock + class _PromptProvider: async def get_prompt(self, args: str, _context) -> str: @@ -50,6 +52,38 @@ async def aclose(self) -> None: self.closed = True +class _TrackedDynamicRuntime(_DynamicRuntime): + def __init__(self) -> None: + super().__init__() + self.closed_event = asyncio.Event() + + async def aclose(self) -> None: + await super().aclose() + self.closed_event.set() + + +class _ObservedTurnAdmissionLock(WebTurnAdmissionLock): + def __init__(self) -> None: + super().__init__() + self.first_waiter = asyncio.Event() + self.second_waiter = asyncio.Event() + self._pending_waiters = 0 + + async def acquire(self): + contended = self.locked() + if contended: + self._pending_waiters += 1 + if self._pending_waiters == 1: + self.first_waiter.set() + elif self._pending_waiters == 2: + self.second_waiter.set() + try: + return await super().acquire() + finally: + if contended: + self._pending_waiters -= 1 + + @pytest.mark.asyncio async def test_suggestions_use_dynamic_mcp_registry_and_close_ephemeral_runtimes(tmp_path, monkeypatch) -> None: from iac_code.web.app import create_app @@ -432,57 +466,68 @@ async def test_interrupting_dynamic_command_during_owner_handoff_cleans_reservat from iac_code.web.app import create_app from iac_code.web.session_manager import WebSessionManager - runtime_started = threading.Event() + loop = asyncio.get_running_loop() + runtime_started = asyncio.Event() release_runtime = threading.Event() + runtime = _TrackedDynamicRuntime() def create_runtime(_options): - runtime_started.set() + loop.call_soon_threadsafe(runtime_started.set) release_runtime.wait() - return _DynamicRuntime() + return runtime monkeypatch.setattr("iac_code.web.runtime.create_agent_runtime", create_runtime) manager = WebSessionManager(projects_dir=tmp_path / "projects", cwd=tmp_path) session = manager.create_session(session_id="dynamic-command-handoff-cancel") + admission_lock = _ObservedTurnAdmissionLock() + session.turn_admission_lock = admission_lock app = create_app(session_manager=manager) + command_task = None + interrupt_task = None - async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://testserver") as client: - command_task = asyncio.create_task( - client.post( - f"/api/sessions/{session.session_id}/commands", - json={"command": "/mcp__remote__review details"}, + try: + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://testserver") as client: + command_task = asyncio.create_task( + client.post( + f"/api/sessions/{session.session_id}/commands", + json={"command": "/mcp__remote__review details"}, + ) ) - ) - runtime_did_start = await asyncio.wait_for(asyncio.to_thread(runtime_started.wait, 1), timeout=2) - assert runtime_did_start is True + await runtime_started.wait() - await session.turn_admission_lock.acquire() - interrupt_task = asyncio.create_task( - client.post( - f"/api/sessions/{session.session_id}/interrupt", - json={"message": ""}, + await admission_lock.acquire() + interrupt_task = asyncio.create_task( + client.post( + f"/api/sessions/{session.session_id}/interrupt", + json={"message": ""}, + ) ) - ) - for _ in range(100): - if len(session.turn_admission_lock._waiters or ()) >= 1: - break - await asyncio.sleep(0.01) - assert len(session.turn_admission_lock._waiters or ()) == 1 - release_runtime.set() - for _ in range(100): - if len(session.turn_admission_lock._waiters or ()) >= 2: - break - await asyncio.sleep(0.01) - assert len(session.turn_admission_lock._waiters or ()) == 2 - session.turn_admission_lock.release() + await admission_lock.first_waiter.wait() + release_runtime.set() + await admission_lock.second_waiter.wait() + admission_lock.release() - interrupted = await asyncio.wait_for(interrupt_task, timeout=1) - command_response = await asyncio.wait_for(command_task, timeout=1) + interrupted = await interrupt_task + command_response = await command_task + finally: + release_runtime.set() + if admission_lock.owner_task is asyncio.current_task(): + admission_lock.release() + pending_tasks = [task for task in (command_task, interrupt_task) if task is not None] + for task in pending_tasks: + if not task.done(): + task.cancel() + if pending_tasks: + await asyncio.gather(*pending_tasks, return_exceptions=True) + if runtime_started.is_set(): + await runtime.closed_event.wait() assert interrupted.status_code == 200 assert command_response.status_code == 409 assert command_response.json()["canceled"] is True assert session.active_turn_task is None - assert session.turn_admission_lock.locked() is False + assert admission_lock.locked() is False + assert runtime.closed is True @pytest.mark.asyncio @@ -490,44 +535,56 @@ async def test_cancelling_dynamic_command_during_owner_handoff_cleans_reservatio from iac_code.web.app import create_app from iac_code.web.session_manager import WebSessionManager - runtime_started = threading.Event() + loop = asyncio.get_running_loop() + runtime_started = asyncio.Event() release_runtime = threading.Event() + runtime = _TrackedDynamicRuntime() def create_runtime(_options): - runtime_started.set() + loop.call_soon_threadsafe(runtime_started.set) release_runtime.wait() - return _DynamicRuntime() + return runtime monkeypatch.setattr("iac_code.web.runtime.create_agent_runtime", create_runtime) manager = WebSessionManager(projects_dir=tmp_path / "projects", cwd=tmp_path) session = manager.create_session(session_id="dynamic-command-handoff-client-cancel") + admission_lock = _ObservedTurnAdmissionLock() + session.turn_admission_lock = admission_lock app = create_app(session_manager=manager) + command_task = None - async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://testserver") as client: - command_task = asyncio.create_task( - client.post( - f"/api/sessions/{session.session_id}/commands", - json={"command": "/mcp__remote__review details"}, + try: + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://testserver") as client: + command_task = asyncio.create_task( + client.post( + f"/api/sessions/{session.session_id}/commands", + json={"command": "/mcp__remote__review details"}, + ) ) - ) - runtime_did_start = await asyncio.wait_for(asyncio.to_thread(runtime_started.wait, 1), timeout=2) - assert runtime_did_start is True + await runtime_started.wait() - await session.turn_admission_lock.acquire() - release_runtime.set() - for _ in range(100): - if len(session.turn_admission_lock._waiters or ()) >= 1: - break - await asyncio.sleep(0.01) - assert len(session.turn_admission_lock._waiters or ()) == 1 - command_task.cancel() - session.turn_admission_lock.release() + await admission_lock.acquire() + release_runtime.set() + await admission_lock.first_waiter.wait() + command_task.cancel() + admission_lock.release() - with pytest.raises(asyncio.CancelledError): - await command_task + with pytest.raises(asyncio.CancelledError): + await command_task + finally: + release_runtime.set() + if admission_lock.owner_task is asyncio.current_task(): + admission_lock.release() + if command_task is not None: + if not command_task.done(): + command_task.cancel() + await asyncio.gather(command_task, return_exceptions=True) + if runtime_started.is_set(): + await runtime.closed_event.wait() assert session.active_turn_task is None - assert session.turn_admission_lock.locked() is False + assert admission_lock.locked() is False + assert runtime.closed is True @pytest.mark.asyncio diff --git a/uv.lock b/uv.lock index 8d613380..8ca1ca16 100644 --- a/uv.lock +++ b/uv.lock @@ -1641,10 +1641,10 @@ dev = [ [package.metadata] requires-dist = [ - { name = "a2a-sdk", extras = ["http-server", "signing"], marker = "extra == 'a2a'", specifier = ">=1.0.2,<2" }, - { name = "a2a-sdk", extras = ["http-server", "signing"], marker = "extra == 'agui'", specifier = ">=1.0.2,<2" }, - { name = "a2a-sdk", extras = ["http-server", "signing"], marker = "extra == 'http'", specifier = ">=1.0.2,<2" }, - { name = "a2a-sdk", extras = ["signing"], marker = "extra == 'a2a-signing'", specifier = ">=1.0.2,<2" }, + { name = "a2a-sdk", extras = ["http-server", "signing"], marker = "extra == 'a2a'", specifier = "==1.1.0" }, + { name = "a2a-sdk", extras = ["http-server", "signing"], marker = "extra == 'agui'", specifier = "==1.1.0" }, + { name = "a2a-sdk", extras = ["http-server", "signing"], marker = "extra == 'http'", specifier = "==1.1.0" }, + { name = "a2a-sdk", extras = ["signing"], marker = "extra == 'a2a-signing'", specifier = "==1.1.0" }, { name = "ag-ui-protocol", marker = "extra == 'agui'", specifier = "==0.1.20" }, { name = "agent-client-protocol", specifier = ">=0.9.0" }, { name = "aiohttp", specifier = ">=3.10,<4" }, @@ -1697,7 +1697,7 @@ provides-extras = ["http", "a2a", "a2a-signing", "a2a-grpc", "a2a-redis", "agui" [package.metadata.requires-dev] desktop = [ - { name = "a2a-sdk", extras = ["http-server", "signing"], specifier = ">=1.0.2,<2" }, + { name = "a2a-sdk", extras = ["http-server", "signing"], specifier = "==1.1.0" }, { name = "pyinstaller", specifier = "==6.21.0" }, { name = "termaid", marker = "python_full_version >= '3.11'", specifier = ">=0.1" }, { name = "uvicorn", extras = ["standard"], specifier = ">=0.30.0" },