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
113 changes: 68 additions & 45 deletions src/iac_code/a2a/executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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"))
Expand Down Expand Up @@ -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(
Expand All @@ -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,
)


Expand All @@ -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:
Expand Down Expand Up @@ -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,
)


Expand Down Expand Up @@ -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(
Expand All @@ -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(
Expand Down Expand Up @@ -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 "")
Expand Down Expand Up @@ -2443,37 +2462,38 @@ 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,
user_id=user_id,
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 = []
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -2597,39 +2618,41 @@ 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,
tool_use_id=recovered.tool_use_id,
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,
Expand Down
Loading
Loading