Skip to content

Commit c0f190d

Browse files
committed
fix(a2a): preserve normal permission recovery
1 parent 69a39f9 commit c0f190d

2 files changed

Lines changed: 284 additions & 50 deletions

File tree

‎src/iac_code/a2a/executor.py‎

Lines changed: 68 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -582,6 +582,7 @@ async def _persist_normal_permission_snapshot_event(
582582
input_id: str,
583583
event_type: str,
584584
permission: dict[str, Any],
585+
handoff_proven: bool = False,
585586
) -> None:
586587
"""Persist handoff Normal permission state in the existing A2A snapshot.
587588
@@ -602,7 +603,8 @@ async def _persist_normal_permission_snapshot_event(
602603
or normal_handoff.get("targetMode") != "normal"
603604
or not _normal_handoff_has_backup_ack(normal_handoff, journal_events)
604605
):
605-
return
606+
if not handoff_proven:
607+
raise RuntimeError(_("Normal permission restore snapshot could not be persisted."))
606608
pipeline_run_id = _string_value(snapshot.get("pipelineRunId"))
607609
pipeline_task_id = _string_value(snapshot.get("taskId"))
608610
context_id = _string_value(snapshot.get("contextId"))
@@ -670,6 +672,7 @@ async def _persist_normal_permission_snapshot_request(
670672
cwd: str,
671673
session_id: str,
672674
pending: Any,
675+
handoff_proven: bool = False,
673676
) -> None:
674677
permission = pending.envelope()
675678
permission.update(
@@ -685,6 +688,7 @@ async def _persist_normal_permission_snapshot_request(
685688
input_id=pending.input_id,
686689
event_type="permission_requested",
687690
permission=permission,
691+
handoff_proven=handoff_proven,
688692
)
689693

690694

@@ -694,6 +698,7 @@ async def _persist_normal_permission_snapshot_resolution(
694698
session_id: str,
695699
response: PermissionResponse,
696700
decision: str,
701+
handoff_proven: bool = False,
697702
) -> None:
698703
state = _a2a_pipeline_state_for_session(cwd=cwd, session_id=session_id)
699704
if state is None:
@@ -724,6 +729,7 @@ async def _persist_normal_permission_snapshot_resolution(
724729
input_id=response.input_id,
725730
event_type="permission_resolved",
726731
permission=permission,
732+
handoff_proven=handoff_proven,
727733
)
728734

729735

@@ -1949,11 +1955,16 @@ async def consume_normal_stream(target_queue: EventQueue) -> bool:
19491955
and not self._auto_approve_permissions
19501956
)
19511957
if interactive_permission:
1958+
19521959
async def persist_request_before_backup(pending_permission: Any) -> None:
19531960
await _persist_normal_permission_snapshot_request(
19541961
cwd=cwd,
19551962
session_id=ctx.session_id,
19561963
pending=pending_permission,
1964+
handoff_proven=self._normal_handoff_has_state_proof(
1965+
cwd=cwd,
1966+
session_id=ctx.session_id,
1967+
),
19571968
)
19581969

19591970
async def persist_resolution_before_backup(
@@ -1978,6 +1989,10 @@ async def persist_resolution_before_backup(
19781989
decision=value,
19791990
),
19801991
decision=value,
1992+
handoff_proven=self._normal_handoff_has_state_proof(
1993+
cwd=cwd,
1994+
session_id=ctx.session_id,
1995+
),
19811996
)
19821997

19831998
pending = await publish_interactive_permission_boundary(
@@ -2369,6 +2384,10 @@ def audit_claim(value: str) -> bool:
23692384
session_id=context_record.session_id,
23702385
response=response,
23712386
decision=expected_value,
2387+
handoff_proven=self._normal_handoff_has_state_proof(
2388+
cwd=context_record.cwd,
2389+
session_id=context_record.session_id,
2390+
),
23722391
)
23732392
if isinstance(decision, dict) and decision.get("backupStatus") != "committed":
23742393
claim_id = str(decision.get("claimId") or "")
@@ -2443,37 +2462,38 @@ def audit_claim(value: str) -> bool:
24432462
messages = storage.load(context_record.cwd, context_record.session_id)
24442463
if not messages:
24452464
raise InvalidParamsError("permission_resume_invalid: session transcript is unavailable.")
2446-
runtime = create_agent_runtime(
2447-
AgentFactoryOptions(
2448-
model=model,
2449-
session_id=context_record.session_id,
2450-
cwd=context_record.cwd,
2451-
resume_messages=messages,
2452-
a2a_safe_mode=_a2a_safe_mode_enabled(),
2453-
source="a2a",
2454-
)
2455-
)
2456-
configure_runtime_model(
2457-
runtime,
2458-
model,
2459-
from_metadata=self._resolve_model(metadata) is not None,
2460-
metadata_api_key=metadata_api_key,
2461-
request_policy_override=request_policy_override,
2462-
)
2463-
refresh_runtime_cloud_tools(runtime)
2464-
task = await self._task_store.get_or_create_task(
2465-
task_id=response.task_id,
2466-
context_id=response.context_id,
2467-
)
2468-
current_assistant_text: list[str] = []
2469-
normal_final_assistant_text = ""
2465+
runtime = None
24702466
try:
24712467
with a2a_request_context(
24722468
session_id=context_record.session_id,
24732469
user_id=user_id,
24742470
aliyun_credential=aliyun_credential,
24752471
preferred_language=preferred_language,
24762472
):
2473+
runtime = create_agent_runtime(
2474+
AgentFactoryOptions(
2475+
model=model,
2476+
session_id=context_record.session_id,
2477+
cwd=context_record.cwd,
2478+
resume_messages=messages,
2479+
a2a_safe_mode=_a2a_safe_mode_enabled(),
2480+
source="a2a",
2481+
)
2482+
)
2483+
configure_runtime_model(
2484+
runtime,
2485+
model,
2486+
from_metadata=self._resolve_model(metadata) is not None,
2487+
metadata_api_key=metadata_api_key,
2488+
request_policy_override=request_policy_override,
2489+
)
2490+
refresh_runtime_cloud_tools(runtime)
2491+
task = await self._task_store.get_or_create_task(
2492+
task_id=response.task_id,
2493+
context_id=response.context_id,
2494+
)
2495+
current_assistant_text: list[str] = []
2496+
normal_final_assistant_text = ""
24772497
async for event in runtime.agent_loop.resume_permission_boundary(record):
24782498
if isinstance(event, MessageStartEvent):
24792499
current_assistant_text = []
@@ -2503,7 +2523,8 @@ def audit_claim(value: str) -> bool:
25032523
if current_assistant_text:
25042524
normal_final_assistant_text = "".join(current_assistant_text)
25052525
finally:
2506-
await _close_runtime(runtime)
2526+
if runtime is not None:
2527+
await _close_runtime(runtime)
25072528
except PermissionWaitSuspended:
25082529
store.mark_suspended(boundary_id)
25092530
await self._publish_status(
@@ -2597,39 +2618,41 @@ async def _rebuild_normal_permission_audit_event(
25972618
messages = storage.load(cwd, session_id)
25982619
if not messages:
25992620
raise ValueError("permission_resume_invalid: session transcript is unavailable")
2600-
runtime = create_agent_runtime(
2601-
AgentFactoryOptions(
2602-
model=model,
2603-
session_id=session_id,
2604-
cwd=cwd,
2605-
resume_messages=messages,
2606-
a2a_safe_mode=_a2a_safe_mode_enabled(),
2607-
source="a2a",
2608-
)
2609-
)
2621+
runtime = None
26102622
try:
2611-
configure_runtime_model(
2612-
runtime,
2613-
model,
2614-
from_metadata=model_from_metadata,
2615-
metadata_api_key=metadata_api_key,
2616-
request_policy_override=request_policy_override,
2617-
)
2618-
refresh_runtime_cloud_tools(runtime)
26192623
with a2a_request_context(
26202624
session_id=session_id,
26212625
user_id=user_id,
26222626
aliyun_credential=aliyun_credential,
26232627
preferred_language=preferred_language,
26242628
):
2629+
runtime = create_agent_runtime(
2630+
AgentFactoryOptions(
2631+
model=model,
2632+
session_id=session_id,
2633+
cwd=cwd,
2634+
resume_messages=messages,
2635+
a2a_safe_mode=_a2a_safe_mode_enabled(),
2636+
source="a2a",
2637+
)
2638+
)
2639+
configure_runtime_model(
2640+
runtime,
2641+
model,
2642+
from_metadata=model_from_metadata,
2643+
metadata_api_key=metadata_api_key,
2644+
request_policy_override=request_policy_override,
2645+
)
2646+
refresh_runtime_cloud_tools(runtime)
26252647
return await runtime.agent_loop.rebuild_permission_audit_event(
26262648
tool_name=recovered.tool_name,
26272649
tool_input=recovered.tool_input,
26282650
tool_use_id=recovered.tool_use_id,
26292651
audit_context=recovered.audit_context,
26302652
)
26312653
finally:
2632-
await _close_runtime(runtime)
2654+
if runtime is not None:
2655+
await _close_runtime(runtime)
26332656

26342657
async def _publish_permission_recovery_ack(
26352658
self,

0 commit comments

Comments
 (0)