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" },