From 5af2b0518c6e5d9e96ff72060eb6191f6f16150f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=A1=82=E9=A9=AC?= Date: Mon, 31 Aug 2026 15:18:56 +0800 Subject: [PATCH] feat: support active AG-UI pipeline guidance --- src/iac_code/agui/adapter.py | 87 ++++++++++++++++++++++++++++++---- src/iac_code/agui/inputs.py | 1 + tests/agui/test_persistence.py | 74 +++++++++++++++++++++++++++++ 3 files changed, 153 insertions(+), 9 deletions(-) diff --git a/src/iac_code/agui/adapter.py b/src/iac_code/agui/adapter.py index a29047ea..82537d75 100644 --- a/src/iac_code/agui/adapter.py +++ b/src/iac_code/agui/adapter.py @@ -99,6 +99,7 @@ class RunTicket: request_digest: str is_resume: bool preferred_language: str + is_guidance: bool = False local_task: asyncio.Task[Any] | None = None completed: bool = False paused: bool = False @@ -161,6 +162,7 @@ async def admit( await self.start() async with self._lock: props = parse_forwarded_props(run_input.forwarded_props).iac_code + is_guidance = props.active_guidance cwd = resolve_cwd(props.cwd) validate_tools(run_input) try: @@ -174,7 +176,7 @@ async def admit( created = binding is None previous_execution: tuple[str, str | None, str, int, set[str], set[str], bool, bool] | None = None if binding is None: - if run_input.resume is not None: + if run_input.resume is not None or is_guidance: raise AdmissionError("EXECUTION_LOST", "The execution to resume is no longer available.") binding = ThreadBinding( thread_id=run_input.thread_id, @@ -189,14 +191,19 @@ async def admit( if previous_digest is not None: code = "DUPLICATE_RUN_ID" if previous_digest == request_digest else "RUN_ID_CONFLICT" raise AdmissionError(code, "The AG-UI run id has already been used.") - if binding.active_run_id is not None: + if binding.active_run_id is not None and not is_guidance: raise AdmissionError("THREAD_BUSY", "The AG-UI thread already has an active run.") if binding.cwd != cwd or binding.user_id != props.user_id: raise AdmissionError( "THREAD_BINDING_CONFLICT", "The AG-UI thread is already bound to another workspace or caller.", ) - if run_input.resume is not None: + if is_guidance: + if run_input.resume is not None or props.run_mode != "pipeline" or binding.task_id is None: + raise AdmissionError("EXECUTION_LOST", "The execution to resume is no longer available.") + if binding.pending: + raise AdmissionError("RESUME_REQUIRED", "The AG-UI thread is waiting for interrupt responses.") + elif run_input.resume is not None: if props.ros_invocation_id != binding.ros_invocation_id: raise AdmissionError("EXECUTION_LOST", "The resume request does not match the interrupted run.") elif binding.pending: @@ -214,13 +221,15 @@ async def admit( ) self._rotate_execution(binding, ros_invocation_id=props.ros_invocation_id) - binding.active_run_id = run_input.run_id + if not is_guidance: + binding.active_run_id = run_input.run_id binding.run_digests[run_input.run_id] = request_digest try: self._persist_thread(binding) except AguiStateStoreError as exc: binding.run_digests.pop(run_input.run_id, None) - binding.active_run_id = None + if not is_guidance: + binding.active_run_id = None if created: self._threads.pop(binding.thread_id, None) self._executions.pop(binding.execution_id, None) @@ -259,6 +268,7 @@ async def admit( binding=binding, request_digest=request_digest, is_resume=run_input.resume is not None, + is_guidance=is_guidance, preferred_language=normalize_agui_language( props.preferred_language, fallback=normalize_agui_language(preferred_language), @@ -288,11 +298,21 @@ async def stream(self, ticket: RunTicket) -> AsyncIterator[Any]: input=None, timestamp=timestamp_ms(), ) - for reopened in mapper.reopen_pipeline_steps(): - yield reopened ticket.local_task = asyncio.current_task() try: props = parse_forwarded_props(run_input.forwarded_props) + if ticket.is_guidance: + await self._consume_guidance(ticket, props.iac_code) + ticket.completed = True + yield RunFinishedEvent( + thread_id=run_input.thread_id, + run_id=run_input.run_id, + outcome=RunFinishedSuccessOutcome(), + timestamp=timestamp_ms(), + ) + return + for reopened in mapper.reopen_pipeline_steps(): + yield reopened session_emitted = False terminal_state = "" resolved_tools: list[tuple[str, str]] = [] @@ -470,7 +490,8 @@ async def stream(self, ticket: RunTicket) -> AsyncIterator[Any]: raise except Exception: logger.exception("AG-UI A2A adapter run failed", extra={"run_id": run_input.run_id}) - await self._cancel_unrecoverable(ticket) + if not ticket.is_guidance: + await self._cancel_unrecoverable(ticket) ticket.completed = True yield _run_error( run_input, @@ -480,11 +501,55 @@ async def stream(self, ticket: RunTicket) -> AsyncIterator[Any]: usage=aggregate_usage(mapper.usage), ) finally: - if ticket.completed and not ticket.binding.pending: + if ticket.completed and not ticket.is_guidance and not ticket.binding.pending: self._mark_execution_terminal(ticket.binding) self._persist_thread_best_effort(ticket.binding) await self._release_run(ticket) + async def _consume_guidance( + self, + ticket: RunTicket, + props: IacCodeForwardedProps, + ) -> None: + binding = ticket.binding + if binding.task_id is None: + raise AguiError("EXECUTION_LOST", "The A2A task to resume is unavailable.") + user_message = latest_user_message(ticket.run_input) + if user_message is None: + raise AguiError("INVALID_INPUT", "A new run requires a user message.") + message_id, parts = user_message + stream = self.client.stream_message_parts( + self.a2a_url, + parts, + cwd=binding.cwd, + context_id=binding.context_id, + task_id=binding.task_id, + message_id=message_id, + **_a2a_request_options(props, preferred_language=ticket.preferred_language), + ) + try: + async for event in stream: + self._validate_guidance_event(binding, event) + finally: + close_stream = getattr(stream, "aclose", None) + if close_stream is not None: + await close_stream() + + @staticmethod + def _validate_guidance_event(binding: ThreadBinding, event: Any) -> None: + _raise_for_a2a_error(event) + task_id = a2a_task_id(event) + if task_id and task_id != binding.task_id: + raise AguiError("A2A_PROTOCOL_ERROR", "The A2A task identity changed unexpectedly.") + context_id = a2a_context_id(event) + if context_id and context_id != binding.context_id: + raise AguiError("A2A_PROTOCOL_ERROR", "The A2A context identity changed unexpectedly.") + state = a2a_state(event) + if state in _FAILED_STATES: + raise AguiError("A2A_EXECUTION_FAILED", "The A2A execution failed.") + if state == "canceled": + raise AguiError("CANCELLED", "The execution was cancelled.") + def _new_a2a_stream(self, ticket: RunTicket, props: IacCodeForwardedProps) -> Any: user_message = latest_user_message(ticket.run_input) if user_message is None: @@ -1039,6 +1104,10 @@ def _validate_resume(self, ticket: RunTicket) -> list[tuple[PendingInput, Any, A async def disconnect(self, ticket: RunTicket) -> None: if ticket.completed or ticket.paused or ticket.binding.pending: return + if ticket.is_guidance: + ticket.completed = True + await self._release_run(ticket) + return await self._cancel_unrecoverable(ticket) ticket.completed = True await self._release_run(ticket) diff --git a/src/iac_code/agui/inputs.py b/src/iac_code/agui/inputs.py index a147c626..0af3ca88 100644 --- a/src/iac_code/agui/inputs.py +++ b/src/iac_code/agui/inputs.py @@ -57,6 +57,7 @@ class IacCodeForwardedProps(StrictModel): run_mode: Literal["normal", "pipeline"] | None = Field(default=None, alias="runMode") pipeline_name: str | None = Field(default=None, alias="pipelineName") cleanup_only: bool = Field(default=False, alias="cleanupOnly") + active_guidance: bool = Field(default=False, alias="activeGuidance") alibaba_cloud: AlibabaCloudOptions | None = Field(default=None, alias="alibabaCloud", repr=False) diff --git a/tests/agui/test_persistence.py b/tests/agui/test_persistence.py index bc7dfb0b..ad753e61 100644 --- a/tests/agui/test_persistence.py +++ b/tests/agui/test_persistence.py @@ -196,6 +196,80 @@ async def test_disconnect_before_interrupt_is_durable_cancels_the_a2a_task(tmp_p await adapter.aclose() +@pytest.mark.asyncio +async def test_pipeline_guidance_preserves_active_execution_across_adapter_restart( + tmp_path, +) -> None: + state_dir = tmp_path / "state" + fake = FakeA2AClient() + adapter = AguiA2AAdapter( + a2a_url="http://a2a/", + client=fake, + state_dir=state_dir, + ) + initial_payload = _payload(tmp_path) + initial_payload["forwardedProps"]["iacCode"]["runMode"] = "pipeline" + initial_ticket = await adapter.admit(parse_run_input(initial_payload), canonical_digest(initial_payload)) + initial_ticket.binding.task_id = "task-1" + initial_ticket.binding.iac_code_session_id = "session-1" + adapter._persist_thread(initial_ticket.binding) + + guidance_payload = _payload(tmp_path, run_id="run-guidance-1") + guidance_props = guidance_payload["forwardedProps"]["iacCode"] + guidance_props.update( + { + "rosInvocationId": "invocation-guidance-1", + "runMode": "pipeline", + "activeGuidance": True, + } + ) + guidance_ticket = await adapter.admit(parse_run_input(guidance_payload), canonical_digest(guidance_payload)) + guidance_events = [event async for event in adapter.stream(guidance_ticket)] + + assert [event.type for event in guidance_events] == [ + EventType.RUN_STARTED, + EventType.RUN_FINISHED, + ] + assert guidance_ticket.is_guidance is True + assert initial_ticket.binding.active_run_id == "run-1" + assert initial_ticket.binding.task_id == "task-1" + assert fake.stream_options[-1]["task_id"] == "task-1" + assert fake.stream_options[-1]["iac_code_metadata"]["rosInvocationId"] == ("invocation-guidance-1") + assert fake.sent_parts == [{"text": "hello"}] + assert fake.cancelled == [] + await adapter.aclose() + + restarted = AguiA2AAdapter( + a2a_url="http://a2a/", + client=fake, + state_dir=state_dir, + ) + guidance_after_restart = _payload(tmp_path, run_id="run-guidance-2") + restarted_props = guidance_after_restart["forwardedProps"]["iacCode"] + restarted_props.update( + { + "rosInvocationId": "invocation-guidance-2", + "runMode": "pipeline", + "activeGuidance": True, + } + ) + restarted_ticket = await restarted.admit( + parse_run_input(guidance_after_restart), + canonical_digest(guidance_after_restart), + ) + + assert restarted_ticket.binding.active_run_id is None + assert restarted_ticket.binding.task_id == "task-1" + restarted_events = [event async for event in restarted.stream(restarted_ticket)] + assert [event.type for event in restarted_events] == [ + EventType.RUN_STARTED, + EventType.RUN_FINISHED, + ] + assert restarted_ticket.binding.task_id == "task-1" + assert fake.cancelled == [] + await restarted.aclose() + + @pytest.mark.asyncio async def test_restart_restores_interrupt_resume_and_cancel_identity(tmp_path) -> None: state_dir = tmp_path / "state"