Skip to content
Open
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
72 changes: 64 additions & 8 deletions plugins/platforms/a2a/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -532,9 +532,20 @@ def _prepare_task(self, params: dict, peer: str, agent: Optional[dict] = None) -
text = protocol.extract_text(params)
context_id = protocol.extract_context_id(params) or protocol.new_context_id()
task_id = protocol.new_task_id()
message_id = protocol.extract_message_id(params)
scope = self._scope_for_agent(agent)
# messageId idempotency: a retried messageId returns the SAME existing task
# (current state) without starting a second dispatch β€” no duplicate task.
if message_id:
existing = self.tasks.get_by_message_id(message_id, *scope)
if existing:
logger.info("A2A: idempotent retry for message_id=%s -> existing task %s (state=%s); no re-dispatch",
message_id, existing["task_id"], existing["state"])
return protocol.build_task(existing["task_id"], existing["context_id"], existing["state"],
existing.get("reply", ""), created_at=existing.get("created_iso", "")), None
turn = self._turns.track(context_id)
max_turns = protocol.max_pingpong_turns()
rec = self.tasks.create(task_id, context_id, peer, *self._scope_for_agent(agent))
rec = self.tasks.create(task_id, context_id, peer, *scope, message_id=message_id)
if turn > max_turns:
protocol.metrics.anti_loop_triggers += 1
logger.warning("A2A: anti-loop triggered for context %s (turn %d > %d)", context_id, turn, max_turns)
Expand All @@ -548,13 +559,27 @@ def _prepare_task(self, params: dict, peer: str, agent: Optional[dict] = None) -
protocol.metrics.inbound_total += 1
self._register_inline_push(task_id, params, agent=agent)
if not agent.get("local", True):
self._activate_task(task_id)
# Forwarded / cross-profile agents: run the profile subprocess OFF the HTTP
# worker thread so message/send returns WORKING + a stable task id immediately.
# _forward_in_background resolves the pending future; _rpc_message_send's
# background worker finalizes the store (COMPLETED/FAILED) for tasks/get polling.
fut = self._add_pending(task_id, context_id)
try:
reply, state = self._forward_to_profile(agent, peer, context_id, framed)
self._record_outcome(task_id, context_id, peer, state, reply)
return protocol.build_task(task_id, context_id, state, reply, created_at=rec["created_iso"]), None
finally:
if self._loop is not None:
asyncio.run_coroutine_threadsafe(
asyncio.to_thread(self._forward_in_background, agent, peer, context_id, framed, task_id),
self._loop,
)
else:
threading.Thread(target=self._forward_in_background,
args=(agent, peer, context_id, framed, task_id),
name=f"a2a-fwd-{task_id}", daemon=True).start()
except Exception as e:
self._pop_pending(task_id)
return self._end_task(rec, protocol.STATE_FAILED, security.redact_outbound(f"Forward dispatch failed: {e}"))
self.tasks.set_state(task_id, protocol.STATE_WORKING)
return None, {"task_id": task_id, "context_id": context_id, "peer": peer, "future": fut,
"created_iso": rec["created_iso"], "started": time.time()}
if self._loop is None or self._message_handler is None:
return self._end_task(rec, protocol.STATE_FAILED, "Agent gateway not ready to accept A2A tasks.")
fut = self._add_pending(task_id, context_id)
Expand Down Expand Up @@ -609,6 +634,19 @@ def _forward_to_profile(self, agent: dict, peer: str, context_id: str, framed_te
"A2A: could not title forwarded session", commit=True)
return security.redact_outbound((proc.stdout or "").strip()), protocol.STATE_COMPLETED

def _forward_in_background(self, agent: dict, peer: str, context_id: str, framed: str, task_id: str) -> None:
"""Run a forwarded profile task off the HTTP thread and resolve its pending future.

The future is consumed by ``_rpc_message_send``'s background worker, which then
finalizes the TaskStore record (COMPLETED/FAILED) so ``tasks/get`` observes it.
"""
try:
reply, state = self._forward_to_profile(agent, peer, context_id, framed)
except Exception as exc:
reply = security.redact_outbound(f"[forward dispatch error: {exc}]")
state = protocol.STATE_FAILED
self._resolve_task(task_id, state, reply)

def _record_outcome(self, task_id: str, context_id: str, peer: str, state: str, reply: str,
started: Optional[float] = None) -> None:
"""Persist + audit + count a finished task, mark it terminal, and fire its push callback."""
Expand Down Expand Up @@ -663,8 +701,26 @@ def _await_reply(self, pending: dict, keepalive=None) -> tuple[str, str]:
def _rpc_message_send(self, req_id: Any, params: dict, peer: str, agent: Optional[dict] = None, v1_response: bool = False) -> dict:
task, pending = self._prepare_task(params, peer, agent=agent)
if task is None:
state, reply = self._finalize_task(pending, *self._await_reply(pending))
task = protocol.build_task(pending["task_id"], pending["context_id"], state, reply, created_at=pending["created_iso"])
# Async long-task path: return WORKING + the stable task id immediately and finish
# the agent work in a background task. HQ polls tasks/get on the SAME id until
# terminal β€” the original message is never re-sent, so no duplicate execution.
assert pending is not None # (None, pending) when task is None
p: dict = pending
task_id, context_id = p["task_id"], p["context_id"]

def _background(p: dict = p, task_id: str = task_id) -> None:
try:
state, reply = self._await_reply(p)
self._finalize_task(p, state, reply)
except Exception as exc:
logger.exception("A2A: background task %s failed", task_id)
self._finalize_task(p, protocol.STATE_FAILED, f"[async execution error: {exc}]")

if self._loop is not None:
asyncio.run_coroutine_threadsafe(asyncio.to_thread(_background, p), self._loop)
else:
threading.Thread(target=_background, args=(p,), name=f"a2a-bg-{task_id}", daemon=True).start()
task = protocol.build_task(task_id, context_id, protocol.STATE_WORKING, "", created_at=p["created_iso"])
return _ok(req_id, protocol.send_message_response(task) if v1_response else task)

@staticmethod
Expand Down
26 changes: 24 additions & 2 deletions plugins/platforms/a2a/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -182,6 +182,16 @@ def extract_context_id(params: dict) -> str:
return (str(msg.get("contextId") or "") if isinstance(msg, dict) else "") or str(params.get("contextId") or "")


def extract_message_id(params: dict) -> str:
"""v1.0 puts messageId inside the Message; tolerate legacy top-level."""
msg = params.get("message") or {}
if isinstance(msg, dict):
mid = str(msg.get("messageId") or "")
if mid:
return mid
return str(params.get("messageId") or "")


def build_task(task_id: str, context_id: str, state: str, agent_text: str = "", *, created_at: str = "") -> dict:
"""A2A v1.0 Task. ``created_at`` is accepted but NOT serialized: the v1.0 Task proto has no
createdAt and strict ProtoJSON parsers (a2a-sdk) reject unknown fields."""
Expand Down Expand Up @@ -321,9 +331,11 @@ def _push_config_view(rec: dict) -> dict:
return {"configId": rec.get("push_config_id") or "", "taskId": rec["task_id"],
"createdAt": rec.get("created_iso", ""), "pushNotificationConfig": {"url": rec.get("push_url") or ""}}

def create(self, task_id: str, context_id: str, peer: str, agent_slug: str = "", tenant: str = "") -> dict:
def create(self, task_id: str, context_id: str, peer: str, agent_slug: str = "", tenant: str = "",
*, message_id: str = "") -> dict:
rec = {"task_id": task_id, "context_id": context_id, "peer": peer, "agent_slug": agent_slug or "", "tenant": tenant or "",
"state": STATE_SUBMITTED, "reply": "", "created_at": time.time(), "created_iso": now_iso(), "push_url": "", "push_config_id": ""}
"message_id": message_id or "", "state": STATE_SUBMITTED, "reply": "", "created_at": time.time(),
"created_iso": now_iso(), "push_url": "", "push_config_id": ""}
with self._lock:
self._tasks[task_id] = rec
return dict(rec)
Expand Down Expand Up @@ -367,6 +379,16 @@ def get(self, task_id: str, agent_slug: str = "", tenant: str = "") -> Optional[
with self._lock:
return dict(rec) if (rec := self._scoped(task_id, agent_slug, tenant)) else None

def get_by_message_id(self, message_id: str, agent_slug: str = "", tenant: str = "") -> Optional[dict]:
"""Return the task already created for a caller messageId, or None (idempotency)."""
if not message_id:
return None
with self._lock:
for rec in self._tasks.values():
if rec.get("message_id") == message_id and self._in_scope(rec, agent_slug, tenant):
return dict(rec)
return None

def complete(self, task_id: str, state: str, reply: str = "") -> Optional[dict]:
"""Transition a task to a terminal state. Idempotent."""
with self._lock:
Expand Down
4 changes: 2 additions & 2 deletions plugins/platforms/a2a/security.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,8 +31,8 @@ def _startup_env(name: str) -> str:


def _parse_peer_tokens(raw: str) -> dict[str, str]:
""""alice:tok1,bob:tok2" -> {token: peer_name}."""
pairs = [tuple(s.strip() for s in pair.split(":", 1)) for pair in raw.split(",") if ":" in pair]
"""``alice=tok1,bob=tok2`` -> {token: peer_name}."""
pairs = [tuple(s.strip() for s in pair.split("=", 1)) for pair in raw.split(",") if "=" in pair]
return {token: name for name, token in pairs if name and token}


Expand Down
12 changes: 10 additions & 2 deletions tests/plugins/test_a2a_phase23.py
Original file line number Diff line number Diff line change
Expand Up @@ -526,8 +526,16 @@ def fake_forward(*_args):
"peer", agent=agent,
)

assert pending is None
assert adapter.tasks.get(terminal["id"])["state"] == protocol.STATE_COMPLETED
# Forwarded tasks are async: _prepare_task returns a WORKING pending that the
# background thread resolves; the caller finalizes the store afterwards.
assert terminal is None
assert pending is not None
pend: dict = pending
state, reply = pend["future"].result(timeout=5)
assert state == protocol.STATE_COMPLETED
assert reply == "forwarded reply"
adapter._finalize_task(pend, state, reply)
assert adapter.tasks.get(pend["task_id"])["state"] == protocol.STATE_COMPLETED
assert adapter.tasks.get("t-live")["state"] == protocol.STATE_WORKING
assert adapter.tasks.get("t-within-reply-window")["state"] == protocol.STATE_WORKING

Expand Down
55 changes: 37 additions & 18 deletions tests/plugins/test_a2a_plugin.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ def test_host_widens_with_shared_token(self, monkeypatch):

def test_host_widens_with_peer_tokens(self, monkeypatch):
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.setenv("A2A_PEER_TOKENS", "alice:tok1")
monkeypatch.setenv("A2A_PEER_TOKENS", "alice=tok1")
monkeypatch.setenv("A2A_HOST", "0.0.0.0")
assert security.localhost_only() is False
assert security.A2ASecurityContext.capture().resolve_bind_host() == "0.0.0.0"
Expand All @@ -86,12 +86,12 @@ def test_no_tokens_identity_is_client_ip(self, monkeypatch):

def test_peer_token_maps_to_name(self, monkeypatch):
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.setenv("A2A_PEER_TOKENS", "alice:tok-a, bob:tok-b")
monkeypatch.setenv("A2A_PEER_TOKENS", "alice=tok-a, bob=tok-b")
assert security.A2ASecurityContext.capture().authenticate("Bearer tok-a", "1.2.3.4") == "alice"
assert security.A2ASecurityContext.capture().authenticate("Bearer tok-b", "1.2.3.4") == "bob"

def test_wrong_or_missing_token_rejected(self, monkeypatch):
monkeypatch.setenv("A2A_PEER_TOKENS", "alice:tok-a")
monkeypatch.setenv("A2A_PEER_TOKENS", "alice=tok-a")
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
assert security.A2ASecurityContext.capture().authenticate("Bearer nope", "1.2.3.4") is None
assert security.A2ASecurityContext.capture().authenticate(None, "1.2.3.4") is None
Expand All @@ -105,7 +105,7 @@ def test_shared_token_identity_is_ip(self, monkeypatch):

def test_peer_tokens_beat_shared(self, monkeypatch):
monkeypatch.setenv("A2A_BEARER_TOKEN", "shared-tok")
monkeypatch.setenv("A2A_PEER_TOKENS", "carol:tok-c")
monkeypatch.setenv("A2A_PEER_TOKENS", "carol=tok-c")
assert security.A2ASecurityContext.capture().authenticate("Bearer tok-c", "1.1.1.1") == "carol"
assert security.A2ASecurityContext.capture().authenticate("Bearer shared-tok", "1.1.1.1") == "ip:1.1.1.1"

Expand Down Expand Up @@ -1170,7 +1170,7 @@ def _post_unauth():
def test_peer_token_identity_used_for_framing(self, monkeypatch):
"""The authenticated peer-token name (not anything in the body) is the
identity the agent sees in the privacy frame."""
monkeypatch.setenv("A2A_PEER_TOKENS", "alice:tok-alice")
monkeypatch.setenv("A2A_PEER_TOKENS", "alice=tok-alice")
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.setenv("A2A_HOST", "127.0.0.1")

Expand Down Expand Up @@ -1206,13 +1206,13 @@ def test_multiplex_adapter_keeps_profile_scoped_peer_tokens(self, monkeypatch):
set_secret_scope,
)

monkeypatch.setenv("A2A_PEER_TOKENS", "default:default-token")
monkeypatch.setenv("A2A_PEER_TOKENS", "default=default-token")
monkeypatch.delenv("A2A_BEARER_TOKEN", raising=False)
monkeypatch.setenv("A2A_HOST", "127.0.0.1")

set_multiplex_active(True)
scope_token = set_secret_scope(
{"A2A_PEER_TOKENS": "secondary:secondary-token"}
{"A2A_PEER_TOKENS": "secondary=secondary-token"}
)
try:
adapter, base = _make_live_adapter(monkeypatch)
Expand Down Expand Up @@ -1400,10 +1400,17 @@ def fake_forward(agent_arg, peer, context_id, framed_text):
"peer-x",
agent=agent,
)
assert pending is None
assert terminal["status"]["state"] == protocol.STATE_COMPLETED
assert protocol.extract_text(terminal["artifacts"][0]) == "dev reply"
assert adapter.tasks.get(terminal["id"])["state"] == protocol.STATE_COMPLETED
# Forwarded tasks are async: _prepare_task returns a WORKING pending task that resolves
# off the HTTP thread; the caller (_rpc_message_send's background worker) finalizes the store.
assert terminal is None
assert pending is not None
pend: dict = pending
assert adapter.tasks.get(pend["task_id"])["state"] == protocol.STATE_WORKING
state, reply = pend["future"].result(timeout=5)
assert state == protocol.STATE_COMPLETED
assert reply == "dev reply"
adapter._finalize_task(pend, state, reply)
assert adapter.tasks.get(pend["task_id"])["state"] == protocol.STATE_COMPLETED


class TestClientTenantAndDiscovery:
Expand Down Expand Up @@ -1466,13 +1473,25 @@ async def run():
assert resp["id"] == "1"
assert set(resp["result"].keys()) == {"task"}
task = resp["result"]["task"]
assert task["status"]["state"] == protocol.STATE_COMPLETED
assert "hello v1" in protocol.extract_text(task["artifacts"][0])
get_resp = await asyncio.to_thread(_post_json, base + "/", {
"jsonrpc": "2.0", "id": "2", "method": "GetTask",
"params": {"id": task["id"]},
}, {"A2A-Version": "1.0"})
assert get_resp["result"]["id"] == task["id"]
assert task["status"]["state"] == protocol.STATE_WORKING
task_id = task["id"]
# Async lifecycle: message/send returns WORKING immediately; the agent finishes in
# the background and tasks/get observes the SAME task id reaching COMPLETED.
async def poll_get():
r = await asyncio.to_thread(_post_json, base + "/", {
"jsonrpc": "2.0", "id": "2", "method": "GetTask",
"params": {"id": task_id},
}, {"A2A-Version": "1.0"})
return r["result"]
got = await poll_get()
for _ in range(100):
if got["status"]["state"] == protocol.STATE_COMPLETED:
break
await asyncio.sleep(0.05)
got = await poll_get()
assert got["status"]["state"] == protocol.STATE_COMPLETED
assert got["id"] == task_id
assert "hello v1" in protocol.extract_text(got["artifacts"][0])
list_resp = await asyncio.to_thread(_post_json, base + "/", {
"jsonrpc": "2.0", "id": "3", "method": "ListTasks",
"params": {"contextId": task["contextId"], "pageSize": 10},
Expand Down