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 a138cd42b..64b4d03b5 100644 --- a/scripts/a2a/e2e/execution_control/run_execution_control_scenarios.py +++ b/scripts/a2a/e2e/execution_control/run_execution_control_scenarios.py @@ -539,7 +539,7 @@ def _assert_other_context_progress(self) -> None: other = self._stream(_message_payload(self.workspace), name="other-context") started = time.monotonic() other.start() - other.join(min(self.timeout, 2.0)) + other.join(self.timeout) other_context = _first(other.snapshot(), "contextId") assert other_context and other_context != self.context_id assert "ISOLATION_FIXTURE_FINAL" in json.dumps(other.snapshot()) diff --git a/src/iac_code/a2a/transports/dispatcher.py b/src/iac_code/a2a/transports/dispatcher.py index a485d4734..808e0acb1 100644 --- a/src/iac_code/a2a/transports/dispatcher.py +++ b/src/iac_code/a2a/transports/dispatcher.py @@ -1043,15 +1043,11 @@ 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 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 INTERRUPTED_TASK_STATES: + if event_state in TERMINAL_TASK_STATES or event_state in INTERRUPTED_TASK_STATES: return - if terminal_state_seen: - 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 diff --git a/src/iac_code/pipeline/engine/show_diagram_tool.py b/src/iac_code/pipeline/engine/show_diagram_tool.py index 8ae3ef5cc..acbe4cd8b 100644 --- a/src/iac_code/pipeline/engine/show_diagram_tool.py +++ b/src/iac_code/pipeline/engine/show_diagram_tool.py @@ -155,7 +155,11 @@ async def execute(self, *, tool_input: dict[str, Any], context: ToolContext) -> ) except asyncio.CancelledError: if context.event_queue is not None: - await context.event_queue.put( + # Cancellation may be requested more than once while the tool executor + # is draining this task. The queue is intentionally unbounded, so use + # the synchronous API to guarantee the fallback event is visible before + # propagating cancellation. + context.event_queue.put_nowait( _diagram_event_from_render_result( candidate_name=candidate_name, template_content=template_content, diff --git a/tests/a2a/test_execution_control_regressions.py b/tests/a2a/test_execution_control_regressions.py index bcba0c416..eb77aa04d 100644 --- a/tests/a2a/test_execution_control_regressions.py +++ b/tests/a2a/test_execution_control_regressions.py @@ -543,7 +543,7 @@ def gated_write(snapshot): assert closed.is_set() assert control.phase == "terminating" and not control.release_ready await asyncio.wait_for( - store.get_or_create_context(context_id="ctx-2", cwd=str(tmp_path), runtime_factory=lambda _: object()), 1 + store.get_or_create_context(context_id="ctx-2", cwd=str(tmp_path), runtime_factory=lambda _: object()), 3 ) assert writes == [(False, False)] release.set() diff --git a/tests/a2a/test_execution_control_review_regressions.py b/tests/a2a/test_execution_control_review_regressions.py index 8e111a85f..0a4adac0c 100644 --- a/tests/a2a/test_execution_control_review_regressions.py +++ b/tests/a2a/test_execution_control_review_regressions.py @@ -293,7 +293,7 @@ def gated_save(snapshot): assert writes == [(False, False)] assert not control.release_ready await asyncio.wait_for( - store.get_or_create_context(context_id="ctx-2", cwd=str(tmp_path), runtime_factory=lambda _: object()), 1 + store.get_or_create_context(context_id="ctx-2", cwd=str(tmp_path), runtime_factory=lambda _: object()), 3 ) # A previously queued SDK event may arrive while the final write is in flight. await store.save(Task(id="task-1", context_id="ctx-1", status=TaskStatus(state=TaskState.TASK_STATE_WORKING))) diff --git a/tests/a2a/test_transport_dispatcher.py b/tests/a2a/test_transport_dispatcher.py index 25d875a50..ea063d187 100644 --- a/tests/a2a/test_transport_dispatcher.py +++ b/tests/a2a/test_transport_dispatcher.py @@ -1241,6 +1241,41 @@ async def hanging_sdk_subscription(self, params, context): ] +@pytest.mark.asyncio +async def test_subscribe_to_task_stops_after_terminal_status(monkeypatch) -> None: + task = Task( + id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_WORKING), + ) + store = A2ATaskStore() + call_context = ServerCallContext() + await store.save(task, call_context) + record = await store.get_or_create_task(task_id="task-1", context_id="ctx-1") + record.active_task = asyncio.current_task() + + async def hanging_sdk_subscription(self, params, context): + yield TaskStatusUpdateEvent( + task_id="task-1", + context_id="ctx-1", + status=TaskStatus(state=TaskState.TASK_STATE_COMPLETED), + ) + await asyncio.Event().wait() + + monkeypatch.setattr(DefaultRequestHandler, "on_subscribe_to_task", hanging_sdk_subscription) + handler = IacCodeRequestHandler.__new__(IacCodeRequestHandler) + handler.task_store = store + handler._validate_extensions = lambda context: None + + events = await asyncio.wait_for( + _collect_async(handler.on_subscribe_to_task(SubscribeToTaskRequest(id="task-1"), call_context)), + timeout=0.5, + ) + + assert len(events) == 1 + assert events[0].status.state == TaskState.TASK_STATE_COMPLETED + + @pytest.mark.asyncio async def test_subscribe_to_task_recovers_terminal_snapshot_when_sdk_stream_ends_early(monkeypatch) -> None: task = Task( diff --git a/tests/pipeline/engine/test_show_diagram_tool.py b/tests/pipeline/engine/test_show_diagram_tool.py index 80cce2d4d..c4e775437 100644 --- a/tests/pipeline/engine/test_show_diagram_tool.py +++ b/tests/pipeline/engine/test_show_diagram_tool.py @@ -578,14 +578,14 @@ async def test_facts_mode_emits_fallback_optimized_event_when_executor_cancels_l class ShortTimeoutShowArchitectureDiagramTool(ShowArchitectureDiagramTool): @property def timeout(self) -> float | None: - return 0.05 + return 1.0 template = tmp_path / "template.yml" template.write_text(SLB_TEMPLATE, encoding="utf-8") queue: asyncio.Queue = asyncio.Queue() registry = ToolRegistry() registry.register(ShortTimeoutShowArchitectureDiagramTool()) - executor = ToolExecutor(registry, tool_timeout=0.05) + executor = ToolExecutor(registry, tool_timeout=1.0) async def slow_create_semantic_plan_for_architecture_with_llm( architecture_context: dict,