Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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())
Expand Down
6 changes: 1 addition & 5 deletions src/iac_code/a2a/transports/dispatcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
6 changes: 5 additions & 1 deletion src/iac_code/pipeline/engine/show_diagram_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion tests/a2a/test_execution_control_regressions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
2 changes: 1 addition & 1 deletion tests/a2a/test_execution_control_review_regressions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)))
Expand Down
35 changes: 35 additions & 0 deletions tests/a2a/test_transport_dispatcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
4 changes: 2 additions & 2 deletions tests/pipeline/engine/test_show_diagram_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Loading