diff --git a/src/iac_code/a2a/executor.py b/src/iac_code/a2a/executor.py index ccd73390..b2724f80 100644 --- a/src/iac_code/a2a/executor.py +++ b/src/iac_code/a2a/executor.py @@ -582,6 +582,7 @@ async def _persist_normal_permission_snapshot_event( input_id: str, event_type: str, permission: dict[str, Any], + handoff_proven: bool = False, ) -> None: """Persist handoff Normal permission state in the existing A2A snapshot. @@ -602,7 +603,8 @@ async def _persist_normal_permission_snapshot_event( or normal_handoff.get("targetMode") != "normal" or not _normal_handoff_has_backup_ack(normal_handoff, journal_events) ): - return + if not handoff_proven: + raise RuntimeError(_("Normal permission restore snapshot could not be persisted.")) pipeline_run_id = _string_value(snapshot.get("pipelineRunId")) pipeline_task_id = _string_value(snapshot.get("taskId")) context_id = _string_value(snapshot.get("contextId")) @@ -670,6 +672,7 @@ async def _persist_normal_permission_snapshot_request( cwd: str, session_id: str, pending: Any, + handoff_proven: bool = False, ) -> None: permission = pending.envelope() permission.update( @@ -685,6 +688,7 @@ async def _persist_normal_permission_snapshot_request( input_id=pending.input_id, event_type="permission_requested", permission=permission, + handoff_proven=handoff_proven, ) @@ -694,6 +698,7 @@ async def _persist_normal_permission_snapshot_resolution( session_id: str, response: PermissionResponse, decision: str, + handoff_proven: bool = False, ) -> None: state = _a2a_pipeline_state_for_session(cwd=cwd, session_id=session_id) if state is None: @@ -724,6 +729,7 @@ async def _persist_normal_permission_snapshot_resolution( input_id=response.input_id, event_type="permission_resolved", permission=permission, + handoff_proven=handoff_proven, ) @@ -1949,11 +1955,16 @@ async def consume_normal_stream(target_queue: EventQueue) -> bool: and not self._auto_approve_permissions ) if interactive_permission: + async def persist_request_before_backup(pending_permission: Any) -> None: await _persist_normal_permission_snapshot_request( cwd=cwd, session_id=ctx.session_id, pending=pending_permission, + handoff_proven=self._normal_handoff_has_state_proof( + cwd=cwd, + session_id=ctx.session_id, + ), ) async def persist_resolution_before_backup( @@ -1978,6 +1989,10 @@ async def persist_resolution_before_backup( decision=value, ), decision=value, + handoff_proven=self._normal_handoff_has_state_proof( + cwd=cwd, + session_id=ctx.session_id, + ), ) pending = await publish_interactive_permission_boundary( @@ -2369,6 +2384,10 @@ def audit_claim(value: str) -> bool: session_id=context_record.session_id, response=response, decision=expected_value, + handoff_proven=self._normal_handoff_has_state_proof( + cwd=context_record.cwd, + session_id=context_record.session_id, + ), ) if isinstance(decision, dict) and decision.get("backupStatus") != "committed": claim_id = str(decision.get("claimId") or "") @@ -2443,30 +2462,7 @@ def audit_claim(value: str) -> bool: messages = storage.load(context_record.cwd, context_record.session_id) if not messages: raise InvalidParamsError("permission_resume_invalid: session transcript is unavailable.") - runtime = create_agent_runtime( - AgentFactoryOptions( - model=model, - session_id=context_record.session_id, - cwd=context_record.cwd, - resume_messages=messages, - a2a_safe_mode=_a2a_safe_mode_enabled(), - source="a2a", - ) - ) - configure_runtime_model( - runtime, - model, - from_metadata=self._resolve_model(metadata) is not None, - metadata_api_key=metadata_api_key, - request_policy_override=request_policy_override, - ) - refresh_runtime_cloud_tools(runtime) - task = await self._task_store.get_or_create_task( - task_id=response.task_id, - context_id=response.context_id, - ) - current_assistant_text: list[str] = [] - normal_final_assistant_text = "" + runtime = None try: with a2a_request_context( session_id=context_record.session_id, @@ -2474,6 +2470,30 @@ def audit_claim(value: str) -> bool: aliyun_credential=aliyun_credential, preferred_language=preferred_language, ): + runtime = create_agent_runtime( + AgentFactoryOptions( + model=model, + session_id=context_record.session_id, + cwd=context_record.cwd, + resume_messages=messages, + a2a_safe_mode=_a2a_safe_mode_enabled(), + source="a2a", + ) + ) + configure_runtime_model( + runtime, + model, + from_metadata=self._resolve_model(metadata) is not None, + metadata_api_key=metadata_api_key, + request_policy_override=request_policy_override, + ) + refresh_runtime_cloud_tools(runtime) + task = await self._task_store.get_or_create_task( + task_id=response.task_id, + context_id=response.context_id, + ) + current_assistant_text: list[str] = [] + normal_final_assistant_text = "" async for event in runtime.agent_loop.resume_permission_boundary(record): if isinstance(event, MessageStartEvent): current_assistant_text = [] @@ -2503,7 +2523,8 @@ def audit_claim(value: str) -> bool: if current_assistant_text: normal_final_assistant_text = "".join(current_assistant_text) finally: - await _close_runtime(runtime) + if runtime is not None: + await _close_runtime(runtime) except PermissionWaitSuspended: store.mark_suspended(boundary_id) await self._publish_status( @@ -2597,31 +2618,32 @@ async def _rebuild_normal_permission_audit_event( messages = storage.load(cwd, session_id) if not messages: raise ValueError("permission_resume_invalid: session transcript is unavailable") - runtime = create_agent_runtime( - AgentFactoryOptions( - model=model, - session_id=session_id, - cwd=cwd, - resume_messages=messages, - a2a_safe_mode=_a2a_safe_mode_enabled(), - source="a2a", - ) - ) + runtime = None try: - configure_runtime_model( - runtime, - model, - from_metadata=model_from_metadata, - metadata_api_key=metadata_api_key, - request_policy_override=request_policy_override, - ) - refresh_runtime_cloud_tools(runtime) with a2a_request_context( session_id=session_id, user_id=user_id, aliyun_credential=aliyun_credential, preferred_language=preferred_language, ): + runtime = create_agent_runtime( + AgentFactoryOptions( + model=model, + session_id=session_id, + cwd=cwd, + resume_messages=messages, + a2a_safe_mode=_a2a_safe_mode_enabled(), + source="a2a", + ) + ) + configure_runtime_model( + runtime, + model, + from_metadata=model_from_metadata, + metadata_api_key=metadata_api_key, + request_policy_override=request_policy_override, + ) + refresh_runtime_cloud_tools(runtime) return await runtime.agent_loop.rebuild_permission_audit_event( tool_name=recovered.tool_name, tool_input=recovered.tool_input, @@ -2629,7 +2651,8 @@ async def _rebuild_normal_permission_audit_event( audit_context=recovered.audit_context, ) finally: - await _close_runtime(runtime) + if runtime is not None: + await _close_runtime(runtime) async def _publish_permission_recovery_ack( self, diff --git a/tests/a2a/test_executor.py b/tests/a2a/test_executor.py index 276399b8..ca327fbf 100644 --- a/tests/a2a/test_executor.py +++ b/tests/a2a/test_executor.py @@ -43,7 +43,7 @@ SessionBackupService, SessionReconcileResult, ) -from iac_code.services.session_backup_staging import StagedSessionBackupService +from iac_code.services.session_backup_staging import SessionBackupStagingWorker, StagedSessionBackupService from iac_code.services.session_backup_state import NORMAL_HANDOFF_PROOF_KEY, BackupPublicationProof from iac_code.services.session_storage import SessionStorage from iac_code.skills.frontmatter import SkillFrontmatter @@ -229,7 +229,7 @@ async def test_handoff_normal_permission_is_restored_from_pipeline_snapshot( @pytest.mark.asyncio -async def test_normal_permission_does_not_modify_snapshot_without_committed_normal_handoff( +async def test_normal_permission_is_blocked_without_committed_normal_handoff( monkeypatch: pytest.MonkeyPatch, tmp_path: Path, ) -> None: @@ -270,7 +270,8 @@ async def test_normal_permission_does_not_modify_snapshot_without_committed_norm "toolName": "write_memory", }, ) - await _persist_normal_permission_snapshot_request(cwd=str(cwd), session_id=session_id, pending=pending) + with pytest.raises(RuntimeError, match="Normal permission restore snapshot could not be persisted"): + await _persist_normal_permission_snapshot_request(cwd=str(cwd), session_id=session_id, pending=pending) snapshot = snapshot_store.load() assert snapshot is not None @@ -574,6 +575,111 @@ async def close_stale_runtime() -> None: assert executor._normal_handoff_has_state_proof(cwd=str(cwd), session_id=session_id, state=result.state) is True assert await executor._should_route_pipeline_handoff_to_normal(context_id=context_id, cwd=str(cwd)) is True + from iac_code.a2a.executor import _persist_normal_permission_snapshot_request + + staging_root = tmp_path / "permission-staging" + staged_service = StagedSessionBackupService(staging_root, storage_1, retry_delays=()) + staged_executor = IacCodeA2AExecutor( + task_store=store, + model="qwen3.6-plus", + backup_service=staged_service, + ) + pending = SimpleNamespace( + input_id="permission-cross-sandbox", + envelope=lambda: { + "schemaVersion": 1, + "kind": "permission", + "requestTaskId": "task-normal", + "contextId": context_id, + "inputId": "permission-cross-sandbox", + "toolUseId": "tool-cross-sandbox", + "toolName": "ros_stack", + "title": "Delete ROS stack", + "target": "stack-1", + "isReadOnly": False, + "options": [ + {"id": "allow_once", "label": "Allow once"}, + {"id": "deny", "label": "Deny"}, + ], + }, + ) + await _persist_normal_permission_snapshot_request( + cwd=str(cwd), + session_id=session_id, + pending=pending, + handoff_proven=staged_executor._normal_handoff_has_state_proof( + cwd=str(cwd), + session_id=session_id, + ), + ) + staged = staged_service.backup_session( + str(cwd), + session_id, + reason=BackupReason.INPUT_REQUIRED, + critical=True, + ) + + assert staged.staged_committed is True + assert staged.shared_committed is False + staged_snapshot = A2APipelineSnapshotStore(staged.destination / "a2a" / "pipeline").load() + assert staged_snapshot is not None + assert staged_snapshot["display"]["permissions"][0]["inputId"] == "permission-cross-sandbox" + + assert SessionBackupStagingWorker(staging_root, backup_root).run_once() == 1 + sandbox_3_config = tmp_path / "sandbox-3" + monkeypatch.setenv("IAC_CODE_CONFIG_DIR", str(sandbox_3_config)) + storage_3 = SessionStorage(projects_dir=sandbox_3_config / "projects") + service_3 = SessionBackupService(storage_3, retry_delays=()) + restored = service_3.restore_session(str(cwd), session_id) + restored_snapshot = A2APipelineSnapshotStore( + a2a_pipeline_dir_for_session(cwd=str(cwd), session_id=session_id) + ).load() + + assert restored.restored is True + assert restored_snapshot is not None + assert restored_snapshot["display"]["permissions"][0]["inputId"] == "permission-cross-sandbox" + + from iac_code.a2a.executor import _persist_normal_permission_snapshot_resolution + + response = PermissionResponse( + task_id="task-normal", + context_id=context_id, + request_task_id="task-normal", + input_id="permission-cross-sandbox", + tool_use_id="tool-cross-sandbox", + decision="allow_once", + ) + restored_executor = IacCodeA2AExecutor( + task_store=A2ATaskStore(metrics=NoOpA2AMetrics()), + model="qwen3.6-plus", + backup_service=service_3, + ) + await _persist_normal_permission_snapshot_resolution( + cwd=str(cwd), + session_id=session_id, + response=response, + decision="allow_once", + handoff_proven=restored_executor._normal_handoff_has_state_proof( + cwd=str(cwd), + session_id=session_id, + ), + ) + resolved_snapshot = A2APipelineSnapshotStore( + a2a_pipeline_dir_for_session(cwd=str(cwd), session_id=session_id) + ).load() + + assert resolved_snapshot is not None + resolved_permission = resolved_snapshot["display"]["permissions"][0] + assert resolved_permission["pending"] is False + assert resolved_permission["decision"] == "allow_once" + resolved_backup = service_3.backup_session( + str(cwd), + session_id, + reason=BackupReason.INPUT_REQUIRED, + critical=True, + ) + assert resolved_backup.shared_committed is True + @pytest.mark.asyncio async def test_current_generation_proof_clears_stale_active_task_and_runtime(monkeypatch, tmp_path) -> None: @@ -4578,6 +4684,86 @@ def fake_register_cloud_tools(registry, credentials, services): assert seen_access_key_ids == ["client-id"] +@pytest.mark.asyncio +async def test_restart_audit_runtime_refreshes_cloud_tools_with_resume_request_credential( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + from iac_code.services.providers.aliyun import AliyunCredential, AliyunCredentials + + config_dir = tmp_path / "config" + config_dir.mkdir() + monkeypatch.setenv("IAC_CODE_CONFIG_DIR", str(config_dir)) + cwd = tmp_path / "workspace" + cwd.mkdir() + session_id = "session-rebuild-credential" + SessionStorage().append(str(cwd), session_id, Message(role="user", content="delete the stack")) + seen: list[tuple[str, str | None]] = [] + rebuilt = PermissionRequestEvent( + tool_name="ros_stack", + tool_input={"action": "DeleteStack", "StackId": "stack-1"}, + tool_use_id="tool-1", + response_future=pending_future(), + ) + + class AuditLoop: + async def rebuild_permission_audit_event(self, **_kwargs): + credential = AliyunCredentials.load() + seen.append(("audit", credential.access_key_id if credential else None)) + return rebuilt + + def factory(options): + credential = AliyunCredentials.load() + seen.append(("factory", credential.access_key_id if credential else None)) + return FakeRuntime( + agent_loop=AuditLoop(), + session_id=options.session_id, + tool_registry=object(), + aliyun_services=object(), + ) + + def register_cloud_tools(_registry, credentials, _services): + credential = credentials.get_provider("aliyun") + seen.append(("refresh", credential.access_key_id if credential else None)) + + monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", factory) + monkeypatch.setattr("iac_code.tools.cloud.registry.register_cloud_tools", register_cloud_tools) + monkeypatch.setattr("iac_code.services.providers.aliyun.AliyunCredentials._load_from_iac_code_config", lambda: None) + executor = IacCodeA2AExecutor(task_store=A2ATaskStore(metrics=NoOpA2AMetrics()), model="qwen3.6-plus") + + result = await executor._rebuild_normal_permission_audit_event( + recovered=RecoveredPermissionAuditBoundary( + tool_name="ros_stack", + tool_input={"action": "DeleteStack", "StackId": "stack-1"}, + tool_use_id="tool-1", + audit_context={"session_id": session_id, "cwd": str(cwd)}, + ), + cwd=str(cwd), + session_id=session_id, + model="qwen3.6-plus", + model_from_metadata=False, + metadata_api_key=None, + request_policy_override=None, + user_id="user-1", + aliyun_credential=AliyunCredential( + mode="StsToken", + access_key_id="request-sts-id", + access_key_secret="request-sts-secret", + sts_token="request-sts-token", + sts_expiration=4102444800, + region_id="cn-beijing", + ), + preferred_language="zh", + ) + + assert result is rebuilt + assert seen == [ + ("factory", "request-sts-id"), + ("refresh", "request-sts-id"), + ("audit", "request-sts-id"), + ] + + @pytest.mark.asyncio async def test_suspending_permission_answer_waits_for_owner_then_resumes_once( monkeypatch: pytest.MonkeyPatch, @@ -4921,9 +5107,22 @@ def resolve(self, _boundary_id, **kwargs): task_record = await task_store.get_or_create_task(task_id="task-1", context_id="ctx-1") task_record.state = "input-required" task_store.mirror_task(task_record) - runtime = FakeRuntime(agent_loop=RecoveryLoop(), session_id=context_record.session_id) + runtime = FakeRuntime( + agent_loop=RecoveryLoop(), + session_id=context_record.session_id, + tool_registry=object(), + aliyun_services=object(), + ) + seen_access_key_ids: list[str | None] = [] + + 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) + monkeypatch.setattr("iac_code.a2a.executor.PermissionWaitCheckpointStore", lambda *_args: CheckpointStore()) monkeypatch.setattr("iac_code.a2a.executor.create_agent_runtime", lambda _options: runtime) + monkeypatch.setattr("iac_code.tools.cloud.registry.register_cloud_tools", register_cloud_tools) + monkeypatch.setattr("iac_code.services.providers.aliyun.AliyunCredentials._load_from_iac_code_config", lambda: None) executor = IacCodeA2AExecutor( task_store=task_store, @@ -4941,7 +5140,18 @@ def resolve(self, _boundary_id, **kwargs): queue = FakeEventQueue() assert await executor._resume_persisted_permission( - 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": "resume-sts-id", + "alibaba_cloud_access_key_secret": "resume-sts-secret", + "alibaba_cloud_security_token": "resume-sts-token", + "alibaba_cloud_region_id": "cn-beijing", + } + }, + ), queue, response=response, ) @@ -4974,6 +5184,7 @@ def resolve(self, _boundary_id, **kwargs): assert task_record.state == "input-required" assert checkpoint["phase"] == "RESOLVED" assert len(resolved) == 1 + assert seen_access_key_ids == ["resume-sts-id"] assert backup_service.calls == [(str(tmp_path), context_record.session_id, BackupReason.NORMAL_TURN_END, False)]