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
141 changes: 125 additions & 16 deletions scripts/provider-update-guide.zh-CN.md

Large diffs are not rendered by default.

121 changes: 115 additions & 6 deletions src/iac_code/a2a/executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,7 @@
configure_runtime_model,
credentials_with_metadata_api_key,
refresh_runtime_cloud_tools,
resolve_a2a_llm_headers,
resolve_a2a_preferred_language,
)
from iac_code.a2a.task_store import A2ATaskStore, _close_runtime
Expand Down Expand Up @@ -116,6 +117,7 @@
mark_cleanup_prompt_message_completed,
)
from iac_code.pipeline.engine.user_input import PipelineUserInput, normalize_pipeline_user_input
from iac_code.providers.request_headers import use_provider_request_headers
from iac_code.providers.request_policy import ProviderRequestPolicy
from iac_code.services.agent_factory import AgentFactoryOptions, create_agent_runtime
from iac_code.services.capabilities.multimodal import is_model_multimodal
Expand Down Expand Up @@ -1267,10 +1269,23 @@ def __init__(
execution_control_service.set_termination_cleanup(self._terminate_detached_execution)
execution_control_service.set_resume_callback(self._task_store.touch_context)

async def resolve_sideband_permission(self, response: PermissionResponse) -> Message | None:
async def resolve_sideband_permission(
self, response: PermissionResponse, *, metadata: Any = None
) -> Message | None:
if not await self._permission_input_registry.is_sideband_response(response):
return None
approved = await self._permission_input_registry.answer(response)
requested_llm_headers = resolve_a2a_llm_headers(metadata)

async def commit_llm_headers() -> None:
await self._task_store.bind_context_llm_headers(
response.context_id,
requested_llm_headers,
)

approved = await self._permission_input_registry.answer(
response,
before_delivery=commit_llm_headers,
)
return permission_ack_message(response, approved=approved)

async def execute(self, context: RequestContext, event_queue: EventQueue) -> None:
Expand All @@ -1295,8 +1310,52 @@ async def execute(self, context: RequestContext, event_queue: EventQueue) -> Non
if existing is not None and existing.owner == owner and current_task is not None:
await existing.attach_task(current_task, mark_working=False)
bind_execution_control(existing)
with a2a_request_context(telemetry_channel=telemetry_channel):
await self._execute(context, event_queue, context_id=context_id)
requested_llm_headers = resolve_a2a_llm_headers(metadata)
with contextlib.ExitStack() as request_scope:
request_scope.enter_context(a2a_request_context(telemetry_channel=telemetry_channel))
llm_headers: Mapping[str, str] | None = None
llm_headers_committed = False
llm_headers_active = False

async def commit_llm_headers() -> None:
"""Update the shared session binding after task/context validation."""

nonlocal llm_headers, llm_headers_committed
if llm_headers_committed:
return
llm_headers = await self._task_store.bind_context_llm_headers(
context_id,
requested_llm_headers,
)
llm_headers_committed = True

async def activate_llm_headers() -> None:
"""Activate committed headers in the current request's async context."""

await commit_llm_headers()
await activate_bound_llm_headers()

async def activate_bound_llm_headers() -> None:
"""Activate the session binding without replacing it from this request."""

nonlocal llm_headers
nonlocal llm_headers_active
if llm_headers_active:
return
if llm_headers is None:
llm_headers = await self._task_store.bind_context_llm_headers(context_id, None)
assert llm_headers is not None
request_scope.enter_context(use_provider_request_headers(llm_headers, live=True))
llm_headers_active = True

await self._execute(
context,
event_queue,
context_id=context_id,
commit_llm_headers=commit_llm_headers,
activate_bound_llm_headers=activate_bound_llm_headers,
activate_llm_headers=activate_llm_headers,
)
finally:
control = current_execution_control()
if control is not None:
Expand All @@ -1312,7 +1371,31 @@ async def execute(self, context: RequestContext, event_queue: EventQueue) -> Non
reset_execution_control(execution_scope)
reset_execution_participants(participant_scope)

async def _execute(self, context: RequestContext, event_queue: EventQueue, *, context_id: str) -> None:
async def _execute(
self,
context: RequestContext,
event_queue: EventQueue,
*,
context_id: str,
commit_llm_headers: Callable[[], Awaitable[None]] | None = None,
activate_bound_llm_headers: Callable[[], Awaitable[None]] | None = None,
activate_llm_headers: Callable[[], Awaitable[None]] | None = None,
) -> None:
if commit_llm_headers is None:

async def commit_llm_headers() -> None:
return None

if activate_bound_llm_headers is None:

async def activate_bound_llm_headers() -> None:
return None

if activate_llm_headers is None:

async def activate_llm_headers() -> None:
return None

requested_task_id = context.task_id or None
task_id = requested_task_id or "task-" + uuid.uuid4().hex[:12]
permission_response = parse_permission_response(getattr(context, "message", None))
Expand All @@ -1324,8 +1407,19 @@ async def _execute(self, context: RequestContext, event_queue: EventQueue, *, co
pending = None
try:
pending = await self._permission_input_registry.pending_for_response(permission_response)
owner = self._task_store.owner_for_context(getattr(context, "call_context", None))
if owner:
permission_task = await self._task_store.get_task_record(permission_response.task_id)
if permission_task.context_id != permission_response.context_id:
raise PermissionIdentityValidationError("permission_task_context_changed", retryable=False)
if permission_task.owner and permission_task.owner != owner:
raise PermissionIdentityValidationError("permission_task_owner_changed", retryable=False)
with a2a_request_context(aliyun_credential=response_credential):
approved = await self._permission_input_registry.answer(permission_response)
approved = await self._permission_input_registry.answer(
permission_response,
before_delivery=commit_llm_headers,
)
await activate_bound_llm_headers()
except PermissionIdentityValidationError as exc:
await self._publish_permission_identity_error(
event_queue,
Expand All @@ -1342,6 +1436,7 @@ async def _execute(self, context: RequestContext, event_queue: EventQueue, *, co
context,
event_queue,
response=permission_response,
activate_llm_headers=activate_llm_headers,
):
return
raise
Expand Down Expand Up @@ -1383,6 +1478,7 @@ async def _execute(self, context: RequestContext, event_queue: EventQueue, *, co
context,
event_queue,
response=permission_response,
activate_llm_headers=activate_llm_headers,
):
return
raise InvalidParamsError("permission_resume_invalid: suspended permission is unavailable.")
Expand Down Expand Up @@ -1612,6 +1708,7 @@ async def release_context_execution() -> None:
request_policy_override=request_policy_override,
backup_service=self._backup_service,
pipeline_name=requested_pipeline_name,
context_ready_callback=activate_llm_headers,
)
try:
pipeline_result = await pipeline_executor.execute(
Expand Down Expand Up @@ -1703,6 +1800,7 @@ def runtime_factory(session_id: str) -> Any:
cwd=cwd,
runtime_factory=runtime_factory,
)
await activate_llm_headers()
control = current_execution_control()
if control is not None:
control.bind_session(ctx.session_id)
Expand Down Expand Up @@ -2410,6 +2508,7 @@ async def _resume_persisted_permission(
event_queue: EventQueue,
*,
response: PermissionResponse,
activate_llm_headers: Callable[[], Awaitable[None]] | None = None,
) -> bool:
"""Claim and resume a permission whose process-local registry was lost."""

Expand All @@ -2425,6 +2524,9 @@ async def _resume_persisted_permission(
return False
if task_record.context_id != response.context_id:
raise InvalidParamsError("input_response_mismatch: permission task context changed.")
owner = self._task_store.owner_for_context(getattr(context, "call_context", None))
if owner and task_record.owner and task_record.owner != owner:
raise InvalidParamsError("Task belongs to a different owner")
minimum_generation = getattr(task_record, "expected_permission_backup_generation", None)
try:
reconcile_result = await self._reconcile_session_before_route(
Expand Down Expand Up @@ -2468,6 +2570,8 @@ async def _resume_persisted_permission(
decision = record.get("decision")
if not isinstance(decision, dict) or decision.get("value") != expected_value:
raise InvalidParamsError("permission_resume_invalid: permission decision conflicts with receipt.")
if activate_llm_headers is not None:
await activate_llm_headers()
await self._publish_permission_recovery_ack(
event_queue,
response=response,
Expand All @@ -2482,6 +2586,8 @@ async def _resume_persisted_permission(
raise InvalidParamsError(
"permission_resume_invalid: permission decision conflicts with active recovery."
)
if activate_llm_headers is not None:
await activate_llm_headers()
await self._publish_permission_recovery_ack(
event_queue,
response=response,
Expand Down Expand Up @@ -2518,6 +2624,7 @@ def make_pipeline_executor() -> IacCodeA2APipelineExecutor:
metadata_api_key=metadata_api_key,
request_policy_override=request_policy_override,
backup_service=self._backup_service,
context_ready_callback=activate_llm_headers,
)

persisted_decision = record.get("decision")
Expand Down Expand Up @@ -2669,6 +2776,8 @@ def audit_claim(value: str) -> bool:
task_record = await self._task_store.get_or_create_task(
task_id=response.task_id, context_id=response.context_id, restore_interrupted=False
)
if activate_llm_headers is not None:
await activate_llm_headers()
if not await self._permission_wait_coordinator.acquire_restore(boundary_id):
await self._publish_permission_recovery_ack(
event_queue,
Expand Down
26 changes: 23 additions & 3 deletions src/iac_code/a2a/input_required.py
Original file line number Diff line number Diff line change
Expand Up @@ -149,7 +149,13 @@ def envelope(self) -> dict[str, Any]:


class PermissionResolutionOwner(Protocol):
async def resolve_permission(self, pending: PendingPermission, response: PermissionResponse) -> bool: ...
async def resolve_permission(
self,
pending: PendingPermission,
response: PermissionResponse,
*,
before_delivery: Callable[[], Awaitable[None]] | None = None,
) -> bool: ...

async def fail_permission(self, pending: PendingPermission) -> None: ...

Expand Down Expand Up @@ -1161,10 +1167,21 @@ async def register(
request.resolution_owner_managed = True
return pending

async def answer(self, response: PermissionResponse) -> bool:
async def answer(
self,
response: PermissionResponse,
*,
before_delivery: Callable[[], Awaitable[None]] | None = None,
) -> bool:
pending = await self._lookup(response)
if pending.resolution_owner is not None:
return await pending.resolution_owner.resolve_permission(pending, response)
if before_delivery is None:
return await pending.resolution_owner.resolve_permission(pending, response)
return await pending.resolution_owner.resolve_permission(
pending,
response,
before_delivery=before_delivery,
)

coordinator = self._permission_wait_coordinator
if coordinator is not None and pending.boundary_id is not None:
Expand All @@ -1188,6 +1205,7 @@ def audit_new_claim(value: str) -> bool:
source="user",
on_new_claim=audit_new_claim,
before_delivery=lambda record: self._backup_claim_before_delivery(pending, record),
before_release=before_delivery,
)
except (LookupError, ValueError) as exc:
raise InvalidParamsError(f"permission_resume_invalid: {exc}") from exc
Expand All @@ -1213,6 +1231,8 @@ def audit_new_claim(value: str) -> bool:
)
if approved and not audit_ok:
approved = False
if before_delivery is not None:
await before_delivery()
future.set_result(approved)
return approved

Expand Down
37 changes: 34 additions & 3 deletions src/iac_code/a2a/metadata_redaction.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@


class A2AMetadataEchoRedactor:
"""Compatibility wrapper that now preserves canonical A2A message data."""
"""Preserve canonical A2A message data while withholding provider headers."""

def redact_message_echo(
self,
Expand All @@ -18,8 +18,10 @@ def redact_message_echo(
public_path_roots: Iterable[Mapping[str, str]] | None = None,
) -> Message:
del public_path_roots
payload = MessageToDict(message, preserving_proto_field_name=False)
_remove_llm_headers(payload.get("metadata"))
result = Message()
ParseDict(MessageToDict(message, preserving_proto_field_name=False), result)
ParseDict(payload, result)
return result

def redact(
Expand All @@ -29,4 +31,33 @@ def redact(
public_path_roots: Iterable[Mapping[str, str]] | None = None,
) -> Any:
del public_path_roots
return copy.deepcopy(value)
result = copy.deepcopy(value)
if isinstance(result, dict):
metadata = result.get("metadata")
_remove_llm_headers(metadata if isinstance(metadata, dict) else result)
return result


def _remove_llm_headers(metadata: Any) -> None:
if not isinstance(metadata, dict):
return
iac_code = metadata.get("iac_code")
if isinstance(iac_code, dict):
iac_code.pop("llm_headers", None)


def strip_llm_headers_from_metadata(metadata: Any) -> None:
"""Remove provider headers from dict or protobuf Struct metadata in place."""

if isinstance(metadata, dict):
_remove_llm_headers(metadata)
return
fields = getattr(metadata, "fields", None)
if fields is None:
return
iac_code = fields.get("iac_code")
if iac_code is None or iac_code.WhichOneof("kind") != "struct_value":
return
iac_fields = iac_code.struct_value.fields
if "llm_headers" in iac_fields:
del iac_fields["llm_headers"]
4 changes: 4 additions & 0 deletions src/iac_code/a2a/pipeline_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -417,6 +417,7 @@ def __init__(
backup_service: Any | None = None,
aliyun_delegated_executor_factory: Any | None = None,
pipeline_name: str | None = None,
context_ready_callback: Callable[[], Awaitable[None]] | None = None,
) -> None:
self._task_store = task_store
self._model = model
Expand All @@ -443,6 +444,7 @@ def __init__(
self._backup_service = backup_service or SessionBackupService()
self._aliyun_delegated_executor_factory = aliyun_delegated_executor_factory
self._pipeline_name_override = pipeline_name or None
self._context_ready_callback = context_ready_callback

def _resolve_pipeline_name(self) -> str:
"""Pipeline this executor must run.
Expand Down Expand Up @@ -617,6 +619,8 @@ def runtime_factory(session_id: str) -> Any:
cwd=cwd,
runtime_factory=runtime_factory,
)
if self._context_ready_callback is not None:
await self._context_ready_callback()
control = current_execution_control()
if control is not None:
control.bind_session(ctx.session_id)
Expand Down
10 changes: 9 additions & 1 deletion src/iac_code/a2a/pipeline_stream.py
Original file line number Diff line number Diff line change
Expand Up @@ -1473,7 +1473,13 @@ class _PipelinePermissionResolutionOwner:
def __init__(self, publisher: PipelineA2AEventPublisher) -> None:
self.publisher = publisher

async def resolve_permission(self, pending: PendingPermission, response: PermissionResponse) -> bool:
async def resolve_permission(
self,
pending: PendingPermission,
response: PermissionResponse,
*,
before_delivery: Callable[[], Awaitable[None]] | None = None,
) -> bool:
registry = self.publisher.permission_input_registry
if registry is None:
raise PipelineA2APersistenceError("Sub Pipeline permission registry is unavailable")
Expand Down Expand Up @@ -1503,6 +1509,8 @@ async def resolve_permission(self, pending: PendingPermission, response: Permiss
if future is None or future.done():
await registry.complete(pending)
raise PipelineA2APersistenceError("Sub Pipeline permission wait point is unavailable")
if before_delivery is not None:
await before_delivery()
future.set_result(approved)
await registry.complete(pending)
return approved
Expand Down
Loading
Loading