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
34 changes: 23 additions & 11 deletions src/select_ai/agent/a2a/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
from a2a.server.routes import create_jsonrpc_routes
from a2a.server.tasks import TaskUpdater
from a2a.types import AgentCapabilities, AgentCard, AgentInterface, AgentSkill
from google.protobuf.json_format import ParseDict
from starlette.applications import Starlette
from starlette.responses import JSONResponse
from starlette.routing import Route
Expand All @@ -32,11 +33,9 @@
from select_ai.agent.a2a.task_store import OracleTaskStore
from select_ai.version import __version__

_A2UI_MIME_TYPE = "application/a2ui+json"


def _a2ui_payload(result: str | None) -> dict | None:
"""Return an A2UI response envelope, if ``RUN_TEAM`` returned one."""
def _message_parts(result: str | None):
"""Convert a serialized A2A message into its constituent parts."""
if not result:
return None
try:
Expand All @@ -45,11 +44,24 @@ def _a2ui_payload(result: str | None) -> dict | None:
return None
if not isinstance(payload, dict):
return None
if payload.get("metadata", {}).get("mimeType") != _A2UI_MIME_TYPE:
return None
if not isinstance(payload.get("data"), list):
message_parts = payload.get("parts")
if payload.get("kind") != "message" or not isinstance(message_parts, list):
return None
return payload
parts = []
for part in message_parts:
if not isinstance(part, dict):
return None
if part.get("kind") == "text" and isinstance(part.get("text"), str):
parts.append(new_text_part(part["text"]))
elif part.get("kind") == "data" and "data" in part:
output_part = new_data_part(part["data"])
if isinstance(part.get("metadata"), dict):
ParseDict(part["metadata"], output_part.metadata)
parts.append(output_part)
else:
# Avoid silently discarding an unsupported part type.
return None
return parts or None


class DatabaseTeamExecutor(AgentExecutor):
Expand Down Expand Up @@ -81,11 +93,11 @@ async def execute(self, context, event_queue):
prompt=context.get_user_input(),
params={"conversation_id": conversation_id},
)
a2ui_payload = _a2ui_payload(result)
message_parts = _message_parts(result)
await updater.add_artifact(
parts=(
[new_data_part(a2ui_payload)]
if a2ui_payload is not None
message_parts
if message_parts is not None
else [new_text_part(result or "")]
),
name="database-agent-result",
Expand Down
64 changes: 64 additions & 0 deletions tests/a2a/test_a2ui_parts.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
# -----------------------------------------------------------------------------
# Copyright (c) 2026, Oracle and/or its affiliates.
#
# Licensed under the Universal Permissive License v 1.0 as shown at
# https://oss.oracle.com/licenses/upl.
# -----------------------------------------------------------------------------

import json

import pytest
from google.protobuf.json_format import MessageToDict

pytest.importorskip("a2a")

from select_ai.agent.a2a.server import _message_parts


def test_a2a_message_parts_are_forwarded_with_metadata():
result = json.dumps(
{
"kind": "message",
"parts": [
{"kind": "text", "text": "A2UI visualization ready."},
{
"kind": "data",
"data": {
"version": "v0.9",
"createSurface": {"surfaceId": "smoke-test"},
},
"metadata": {"mimeType": "application/json+a2ui"},
},
],
}
)

parts = _message_parts(result)

assert [
MessageToDict(part, preserving_proto_field_name=True) for part in parts
] == [
{"text": "A2UI visualization ready."},
{
"data": {
"version": "v0.9",
"createSurface": {"surfaceId": "smoke-test"},
},
"metadata": {"mimeType": "application/json+a2ui"},
},
]


def test_a2a_text_message_is_forwarded():
result = json.dumps(
{
"kind": "message",
"parts": [{"kind": "text", "text": "ordinary response"}],
}
)

parts = _message_parts(result)

assert [
MessageToDict(part, preserving_proto_field_name=True) for part in parts
] == [{"text": "ordinary response"}]