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
50 changes: 46 additions & 4 deletions src/agent_env/a2a_agent/object_transfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,9 @@
import math
import re
from collections.abc import AsyncIterator, Callable, Collection, Mapping
from contextlib import asynccontextmanager
from contextlib import asynccontextmanager, nullcontext
from dataclasses import dataclass
from typing import Any, Literal, TypeVar
from typing import TYPE_CHECKING, Any, Literal, TypeVar

import httpx
from agentenv_protocol.a2a_agent import (
Expand Down Expand Up @@ -37,6 +37,7 @@
from agentenv_protocol.transfers import ReadObject, WriteNamespaceGrant, WriteObject
from pydantic import BaseModel, ValidationError

from agent_env.a2a_agent import protocol
from agent_env.a2a_agent.protocol import raise_for_extension_status
from agent_env.a2a_agent.staging import StagedObjectStore, staged_store
from agent_env.config import get_config
Expand All @@ -45,6 +46,9 @@
from agent_env.store.object_store.object_store import readable_url
from agent_env.store.object_store.local.grant_server import unreachable_hint

if TYPE_CHECKING:
from agent_env.task_step.context import DeployedAgent

logger = logging.getLogger(__name__)

# "objects": the call carries grants and the agent moves the bytes. "legacy": the inline forms,
Expand Down Expand Up @@ -308,11 +312,13 @@ async def readable_parts(
card: Mapping[str, Any] | None,
sandbox_type: str | None,
expires_in: int,
shareable: Collection[str] | None = None,
) -> AsyncIterator[list[dict]]:
"""``parts`` as the agent at ``a2a_url`` can read them: a file part naming an object a configured store
owns names an HTTPS URL for it instead, one ``readable_url`` gives for at least ``expires_in`` seconds,
or else a copy staged on the agent for the length of the block. Other parts, and file parts naming
anything else, are sent as they are. Raises when an owned object can be given no URL the agent can read."""
anything else, are sent as they are; so is an owned object ``shareable`` (when given) does not name.
Raises when an owned object can be given no URL the agent can read."""
config = get_config()
readable = list(parts)
staged: dict[int, StagedObjectStore] = {} # by identity: a store need not be hashable
Expand All @@ -322,7 +328,7 @@ async def readable_parts(
if not isinstance(uri, str):
continue
store = config.get_object_store_at(uri)
if not store.owns(uri):
if not store.owns(uri) or (shareable is not None and uri not in shareable):
continue
url = await asyncio.to_thread(readable_url, store, uri, sandbox_type=sandbox_type, expires_in=expires_in)
if url is None:
Expand All @@ -346,6 +352,42 @@ async def readable_parts(
await staging.release()


async def send_and_wait(
a2a_url: str,
parts: list[dict],
*,
agent: DeployedAgent | None,
message_id: str,
context_id: str | None,
timeout_seconds: int,
poll_interval_seconds: int,
before_send: Callable[[], None] | None = None,
shareable: Collection[str] | None = None,
) -> tuple[str, dict]:
"""Send ``parts`` to the A2A peer at ``a2a_url`` and wait for its task to end: the task's id and its
final state. A peer that is an ``agent`` agent-env deployed is sent each file part a configured store
owns, of those ``shareable`` names when given, as an HTTPS URL it can read (``readable_parts``); any
other peer, such as a human's hub, reads the store itself and is sent the parts as they are.
``before_send`` runs once the parts are ready, just before the message goes out. An ``agent``'s sandbox is
watched while it works, so one that dies is given up on (``poll_a2a_task``)."""
sending = (
readable_parts(
parts, a2a_url=a2a_url, card=agent.a2a_card, sandbox_type=agent.sandbox_type,
expires_in=timeout_seconds, shareable=shareable,
)
if agent is not None
else nullcontext(parts)
)
async with sending as sent:
if before_send is not None:
before_send()
task_id, _ = await protocol.send_a2a_message(a2a_url, sent, message_id, context_id, timeout_seconds)
return task_id, await protocol.poll_a2a_task(
a2a_url, task_id, timeout_seconds, poll_interval_seconds,
sandbox_id=agent.sandbox_id if agent is not None else None,
)


def skill_bundle_request(
store: ObjectStore, *, name: str, description: str, object_url: str
) -> BundleSkillRequest:
Expand Down
50 changes: 27 additions & 23 deletions src/agent_env/task_step/task_steps/prompt_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
from agent_env.a2a_agent.object_transfer import (
TrajectoryUpload,
fetch_trajectory,
readable_parts,
send_and_wait,
trajectory_mode,
)
from agent_env.a2a_agent.staging import draining, staged_changelogs, transfer_store
Expand Down Expand Up @@ -91,6 +91,11 @@ def _duplicates_prompt_text(parts: list[dict], prompt_text: Optional[str]) -> bo
return parts == [{"kind": "text", "text": prompt_text}]


def _file_uris(parts: list[dict]) -> list[str]:
files = (part.get("file") for part in parts if part.get("kind") == "file")
return [file["uri"] for file in files if isinstance(file, dict) and isinstance(file.get("uri"), str)]


_DEFAULT_USER_SIM_OUTPUT_FORMAT: dict[str, Any] = {
"type": "json_schema",
"schema": {
Expand Down Expand Up @@ -474,18 +479,20 @@ async def _execute_conversation(
traj_ext_cached = A2AAgent.find_extension(card, A2AAgent.EXT_TRAJECTORY)
final_state: str = TaskState.completed.value
trajectory_s3_uri: Optional[str] = None
# What the target has been sent: of the files its replies name, the only ones a user-sim is made able
# to read, so a reply naming any other object a store owns can't read it out through the user-sim.
sent_to_target: set[str] = set()

for turn in range(self.max_conversation_turns):
# `target_a2a_task_id` is the client A2A message id sent to the target
# agent and recorded on the conversation as `a2a_task_id`.
# This id is also used as part of the key name for the trajectory S3 object.
target_a2a_task_id = uuid.uuid4().hex
# Only the sent copy names readable URLs; the turn records the objects' own URLs, and only
# once the copy is ready, so an object the agent can't be sent leaves no turn waiting.
async with readable_parts(
current_user_parts, a2a_url=target_url, card=card,
sandbox_type=agent.sandbox_type, expires_in=self.timeout_seconds,
) as sent_parts:

# The turn records the objects' own URLs, and only once the parts are ready to send, so an
# object the agent can't be sent leaves no turn waiting.
def record_turn() -> None:
sent_to_target.update(_file_uris(current_user_parts))
conversation_store.add_a2a_task(
conversation_id=conversation_id,
parts=current_user_parts,
Expand All @@ -502,14 +509,11 @@ async def _execute_conversation(
else list(current_user_parts)
)

sent_task_id, _ = await protocol.send_a2a_message(
target_url, sent_parts, target_a2a_task_id,
solver_context_id, self.timeout_seconds,
)
result = await protocol.poll_a2a_task(
target_url, sent_task_id, self.timeout_seconds, self.poll_interval_seconds,
sandbox_id=agent.sandbox_id,
)
sent_task_id, result = await send_and_wait(
target_url, current_user_parts, agent=agent, message_id=target_a2a_task_id,
context_id=solver_context_id, timeout_seconds=self.timeout_seconds,
poll_interval_seconds=self.poll_interval_seconds, before_send=record_turn,
)
target_state = result["status"]["state"]
status_msg = (result.get("status") or {}).get("message") or {}
final_terminal = protocol.TerminalResponse.from_message(status_msg)
Expand Down Expand Up @@ -573,14 +577,14 @@ async def _execute_conversation(

user_a2a_task_id = uuid.uuid4().hex
try:
sent_user_task_id, _ = await protocol.send_a2a_message(
user_url, agent_response_parts, user_a2a_task_id,
conversation_id, self.user_agent_timeout_seconds,
)
user_result = await protocol.poll_a2a_task(
user_url, sent_user_task_id,
self.user_agent_timeout_seconds, self.poll_interval_seconds,
sandbox_id=user_sim.sandbox_id if is_user_sim else None,
# A user-sim runs in a sandbox agent-env deployed; a human peer, registered or named by
# user_a2a_url, has none and reads the store itself.
_, user_result = await send_and_wait(
user_url, agent_response_parts,
agent=user_sim if is_user_sim and user_sim.sandbox_id else None, shareable=sent_to_target,
message_id=user_a2a_task_id, context_id=conversation_id,
timeout_seconds=self.user_agent_timeout_seconds,
poll_interval_seconds=self.poll_interval_seconds,
)
except TimeoutError as e:
logger.warning(f"user_a2a_url timeout for conversation {conversation_id} ({e}); marking abandoned")
Expand Down
32 changes: 31 additions & 1 deletion tst/data/a2a_agent/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,17 +3,25 @@
Deterministic: it echoes the prompt and records a native trajectory. Agent-config, mcp-config,
trajectory and triggers come from the agentenv-protocol framework. ``model_params`` is a write-only
config field so the agent-config negotiation and its redaction on read-back can be observed.

It also moves files both ways. After the echo it adds a ``read <name> over <scheme>: <text>`` line for
each file part it is sent, fetching only inline bytes and HTTP(S) URLs (``could not read ...`` otherwise),
and it answers each ``send-file <uri>`` line of its prompt with a file part naming that URI.
"""

import base64
from typing import Any
from urllib.parse import urlsplit

import httpx
from agentenv_protocol.a2a_agent import (
MCP_CONFIG_V1,
TRAJECTORY_V1,
TRIGGERS_V1,
AgentConfig,
AgentEnvAgent,
AgentIdentity,
FilePart,
TaskRequest,
TaskResult,
TextPart,
Expand All @@ -22,6 +30,9 @@
a2a_agent,
)

SEND_FILE = "send-file "
READ_TIMEOUT_SECONDS = 120


class EchoAgentConfig(AgentConfig):
model: str | None = None
Expand All @@ -30,6 +41,22 @@ class EchoAgentConfig(AgentConfig):
model_params: WriteOnly[dict[str, Any] | None] = None


async def _read(part: FilePart) -> str:
"""What a file part holds, inline or at an HTTP(S) URL."""
scheme = "bytes" if part.bytes is not None else urlsplit(part.uri).scheme
try:
if part.bytes is not None:
data = base64.b64decode(part.bytes)
else:
async with httpx.AsyncClient(timeout=READ_TIMEOUT_SECONDS) as client:
response = await client.get(part.uri)
response.raise_for_status()
data = response.content
except Exception as exc: # noqa: BLE001 -- said in the reply, so a test sees why
return f"could not read {part.name} over {scheme}: {type(exc).__name__}"
return f"read {part.name} over {scheme}: {data.decode(errors='replace').strip()}"


@a2a_agent(
identity=AgentIdentity(
name="agentenv-echo-agent",
Expand All @@ -43,10 +70,13 @@ class EchoAgent(AgentEnvAgent):
async def run(self, request: TaskRequest[EchoAgentConfig]) -> TaskResult:
prompt = "\n".join(part.text for part in request.parts if isinstance(part, TextPart))
reply = f"Echo: {prompt}"
reply = "\n".join([reply, *[await _read(part) for part in request.parts if isinstance(part, FilePart)]])
sends = [line.removeprefix(SEND_FILE).strip() for line in prompt.splitlines() if line.startswith(SEND_FILE)]
files = [FilePart(uri=uri, name=uri.rsplit("/", 1)[-1], mime_type="text/plain") for uri in sends]
return (
TaskResult.builder()
.succeeded()
.add_text(reply)
.parts([TextPart(text=reply), *files])
.usage(Usage(tool_call_count=0, input_tokens=len(prompt), output_tokens=len(reply)))
.native_trajectory(
format="agentenv-echo-agent/v1",
Expand Down
Loading
Loading