Skip to content

Commit 930e472

Browse files
google-genai-botcopybara-github
authored andcommitted
refactor: add a2a-sdk 0.3.x/1.x compatibility layer for ADK A2A support
Add a shim layer so the ADK A2A integration is compatible with whichever a2a-sdk major version (0.3.x or 1.x) is installed in the user's environment. The shim isolates the differences between the two SDKs - notably protobuf-based 1.x types vs pydantic 0.3.x models, the AgentCard and Part shapes, message/event and agent-card construction, client configuration, and HTTP hosting setup - behind version-agnostic helpers. A2A converters, the agent executor, the remote agent, logging, agent-card building, and the agent registry use these helpers instead of reaching into the SDK directly. The test suite is made version-agnostic so it passes under both SDK majors. PiperOrigin-RevId: 944970218
1 parent 6f66814 commit 930e472

33 files changed

Lines changed: 4438 additions & 2557 deletions

src/google/adk/a2a/_compat.py

Lines changed: 1132 additions & 0 deletions
Large diffs are not rendered by default.

src/google/adk/a2a/agent/config.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,11 +23,11 @@
2323
from typing import Optional
2424
from typing import Union
2525

26-
from a2a.client.middleware import ClientCallContext
2726
from a2a.server.events import Event as A2AEvent
2827
from a2a.types import Message as A2AMessage
2928
from pydantic import BaseModel
3029

30+
from .. import _compat
3131
from ...a2a.converters.part_converter import A2APartToGenAIPartConverter
3232
from ...a2a.converters.part_converter import convert_a2a_part_to_genai_part
3333
from ...a2a.converters.to_adk_event import A2AArtifactUpdateToEventConverter
@@ -46,7 +46,7 @@ class ParametersConfig(BaseModel):
4646
"""Configuration for the parameters passed to the A2A send_message request."""
4747

4848
request_metadata: Optional[dict[str, Any]] = None
49-
client_call_context: Optional[ClientCallContext] = None
49+
client_call_context: Optional[_compat.ClientCallContext] = None
5050
# TODO: Add support for requested_extension and
5151
# message_send_configuration once they are supported by the A2A client.
5252
#

src/google/adk/a2a/agent/interceptors/new_integration_extension.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -17,14 +17,15 @@
1717

1818
from typing import Union
1919

20-
from a2a.client.middleware import ClientCallContext
2120
from a2a.extensions.common import HTTP_EXTENSION_HEADER
2221
from a2a.types import Message as A2AMessage
2322
from google.adk.a2a.agent.config import ParametersConfig
2423
from google.adk.a2a.agent.config import RequestInterceptor
2524
from google.adk.agents.invocation_context import InvocationContext
2625
from google.adk.events.event import Event
2726

27+
from ... import _compat
28+
2829
_NEW_A2A_ADK_INTEGRATION_EXTENSION = (
2930
'https://google.github.io/adk-docs/a2a/a2a-extension/'
3031
)
@@ -37,7 +38,7 @@ async def _before_request(
3738
) -> tuple[Union[A2AMessage, Event], ParametersConfig]:
3839
"""Adds A2A_new_agent_version to client_call_context."""
3940
if params.client_call_context is None:
40-
params.client_call_context = ClientCallContext()
41+
params.client_call_context = _compat.ClientCallContext()
4142

4243
http_kwargs = params.client_call_context.state.get('http_kwargs', {})
4344
headers = http_kwargs.get('headers', {})

src/google/adk/a2a/agent/utils.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -19,12 +19,12 @@
1919
from typing import Optional
2020
from typing import Union
2121

22-
from a2a.client import ClientEvent as A2AClientEvent
23-
from a2a.client.middleware import ClientCallContext
2422
from a2a.types import Message as A2AMessage
2523

24+
from .. import _compat
2625
from ...agents.invocation_context import InvocationContext
2726
from ...events.event import Event
27+
from .._compat import A2AClientEvent
2828
from .config import ParametersConfig
2929
from .config import RequestInterceptor
3030

@@ -37,7 +37,7 @@ async def execute_before_request_interceptors(
3737
"""Executes registered before_request interceptors."""
3838

3939
params = ParametersConfig(
40-
client_call_context=ClientCallContext(state=ctx.session.state)
40+
client_call_context=_compat.ClientCallContext(state=ctx.session.state)
4141
)
4242
if request_interceptors:
4343
for interceptor in request_interceptors:

src/google/adk/a2a/converters/event_converter.py

Lines changed: 67 additions & 77 deletions
Original file line numberDiff line numberDiff line change
@@ -15,28 +15,21 @@
1515
from __future__ import annotations
1616

1717
from collections.abc import Callable
18-
from datetime import datetime
19-
from datetime import timezone
2018
import logging
2119
from typing import Any
2220
from typing import Dict
2321
from typing import List
2422
from typing import Optional
2523

2624
from a2a.server.events import Event as A2AEvent
27-
from a2a.types import DataPart
2825
from a2a.types import Message
2926
from a2a.types import Part as A2APart
30-
from a2a.types import Role
3127
from a2a.types import Task
32-
from a2a.types import TaskState
33-
from a2a.types import TaskStatus
3428
from a2a.types import TaskStatusUpdateEvent
35-
from a2a.types import TextPart
36-
from google.adk.platform import time as platform_time
3729
from google.adk.platform import uuid as platform_uuid
3830
from google.genai import types as genai_types
3931

32+
from .. import _compat
4033
from ...agents.invocation_context import InvocationContext
4134
from ...events.event import Event
4235
from ...flows.llm_flows.functions import REQUEST_EUC_FUNCTION_CALL_NAME
@@ -184,19 +177,20 @@ def _process_long_running_tool(a2a_part: A2APart, event: Event) -> None:
184177
a2a_part: The A2A part to potentially mark as long-running.
185178
event: The ADK event containing long-running tool information.
186179
"""
180+
meta = _compat.part_metadata(a2a_part)
187181
if (
188-
isinstance(a2a_part.root, DataPart)
182+
_compat.is_data_part(a2a_part)
189183
and event.long_running_tool_ids
190-
and a2a_part.root.metadata
191-
and a2a_part.root.metadata.get(
192-
_get_adk_metadata_key(A2A_DATA_PART_METADATA_TYPE_KEY)
193-
)
184+
and meta
185+
and meta.get(_get_adk_metadata_key(A2A_DATA_PART_METADATA_TYPE_KEY))
194186
== A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL
195-
and a2a_part.root.data.get("id") in event.long_running_tool_ids
196187
):
197-
a2a_part.root.metadata[
198-
_get_adk_metadata_key(A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
199-
] = True
188+
data = _compat.data_part_dict(a2a_part)
189+
if data.get("id") in event.long_running_tool_ids:
190+
meta[
191+
_get_adk_metadata_key(A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
192+
] = True
193+
_compat.set_part_metadata(a2a_part, meta)
200194

201195

202196
def convert_a2a_task_to_event(
@@ -229,7 +223,9 @@ def convert_a2a_task_to_event(
229223
message = None
230224
if a2a_task.artifacts:
231225
message = Message(
232-
message_id="", role=Role.agent, parts=a2a_task.artifacts[-1].parts
226+
message_id="",
227+
role=_compat.ROLE_AGENT,
228+
parts=a2a_task.artifacts[-1].parts,
233229
)
234230
elif (
235231
a2a_task.status
@@ -321,9 +317,10 @@ def convert_a2a_message_to_event(
321317
continue
322318

323319
# Check for long-running tools
320+
pmeta = _compat.part_metadata(a2a_part)
324321
if (
325-
a2a_part.root.metadata
326-
and a2a_part.root.metadata.get(
322+
pmeta
323+
and pmeta.get(
327324
_get_adk_metadata_key(
328325
A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY
329326
)
@@ -372,7 +369,7 @@ def convert_a2a_message_to_event(
372369
def convert_event_to_a2a_message(
373370
event: Event,
374371
invocation_context: InvocationContext | None = None,
375-
role: Role = Role.agent,
372+
role: Any = _compat.ROLE_AGENT,
376373
part_converter: GenAIPartToA2APartConverter = convert_genai_part_to_a2a_part,
377374
) -> Optional[Message]:
378375
"""Converts an ADK event to an A2A message.
@@ -441,27 +438,20 @@ def _create_error_status_event(
441438
if event.error_code:
442439
event_metadata[_get_adk_metadata_key("error_code")] = str(event.error_code)
443440

444-
return TaskStatusUpdateEvent(
441+
err_msg_part = Message(
442+
message_id=platform_uuid.new_uuid(),
443+
role=_compat.ROLE_AGENT,
444+
parts=[_compat.make_text_part(error_message)],
445+
metadata={_get_adk_metadata_key("error_code"): str(event.error_code)}
446+
if event.error_code
447+
else {},
448+
)
449+
return _compat.make_task_status_update_event(
445450
task_id=task_id,
446451
context_id=context_id,
447-
metadata=event_metadata,
448-
status=TaskStatus(
449-
state=TaskState.failed,
450-
message=Message(
451-
message_id=platform_uuid.new_uuid(),
452-
role=Role.agent,
453-
parts=[TextPart(text=error_message)],
454-
metadata={
455-
_get_adk_metadata_key("error_code"): str(event.error_code)
456-
}
457-
if event.error_code
458-
else {},
459-
),
460-
timestamp=datetime.fromtimestamp(
461-
platform_time.get_time(), tz=timezone.utc
462-
).isoformat(),
463-
),
452+
status=_compat.make_task_status(_compat.TS_FAILED, message=err_msg_part),
464453
final=True,
454+
metadata=event_metadata,
465455
)
466456

467457

@@ -484,48 +474,47 @@ def _create_status_update_event(
484474
Returns:
485475
A TaskStatusUpdateEvent with RUNNING state.
486476
"""
487-
status = TaskStatus(
488-
state=TaskState.working,
489-
message=message,
490-
timestamp=datetime.fromtimestamp(
491-
platform_time.get_time(), tz=timezone.utc
492-
).isoformat(),
493-
)
477+
status = _compat.make_task_status(_compat.TS_WORKING, message=message)
478+
479+
def is_euc_call(p: Any) -> bool:
480+
m = _compat.part_metadata(p)
481+
if not m:
482+
return False
483+
data = _compat.data_part_dict(p) if _compat.is_data_part(p) else {}
484+
return (
485+
m.get(_get_adk_metadata_key(A2A_DATA_PART_METADATA_TYPE_KEY))
486+
== A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL
487+
and m.get(
488+
_get_adk_metadata_key(A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
489+
)
490+
is True
491+
and data.get("name") == REQUEST_EUC_FUNCTION_CALL_NAME
492+
)
494493

495-
if any(
496-
part.root.metadata.get(
497-
_get_adk_metadata_key(A2A_DATA_PART_METADATA_TYPE_KEY)
498-
)
499-
== A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL
500-
and part.root.metadata.get(
501-
_get_adk_metadata_key(A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
502-
)
503-
is True
504-
and part.root.data.get("name") == REQUEST_EUC_FUNCTION_CALL_NAME
505-
for part in message.parts
506-
if part.root.metadata
507-
):
508-
status.state = TaskState.auth_required
509-
elif any(
510-
part.root.metadata.get(
511-
_get_adk_metadata_key(A2A_DATA_PART_METADATA_TYPE_KEY)
512-
)
513-
== A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL
514-
and part.root.metadata.get(
515-
_get_adk_metadata_key(A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
516-
)
517-
is True
518-
for part in message.parts
519-
if part.root.metadata
520-
):
521-
status.state = TaskState.input_required
494+
def is_long_running_call(p: Any) -> bool:
495+
m = _compat.part_metadata(p)
496+
if not m:
497+
return False
498+
return (
499+
m.get(_get_adk_metadata_key(A2A_DATA_PART_METADATA_TYPE_KEY))
500+
== A2A_DATA_PART_METADATA_TYPE_FUNCTION_CALL
501+
and m.get(
502+
_get_adk_metadata_key(A2A_DATA_PART_METADATA_IS_LONG_RUNNING_KEY)
503+
)
504+
is True
505+
)
506+
507+
if any(is_euc_call(part) for part in message.parts):
508+
status.state = _compat.TS_AUTH_REQUIRED
509+
elif any(is_long_running_call(part) for part in message.parts):
510+
status.state = _compat.TS_INPUT_REQUIRED
522511

523-
return TaskStatusUpdateEvent(
512+
return _compat.make_task_status_update_event(
524513
task_id=task_id,
525514
context_id=context_id,
526515
status=status,
527-
metadata=_get_context_metadata(event, invocation_context),
528516
final=False,
517+
metadata=_get_context_metadata(event, invocation_context),
529518
)
530519

531520

@@ -560,7 +549,6 @@ def convert_event_to_a2a_events(
560549
a2a_events = []
561550

562551
try:
563-
564552
# Handle error scenarios
565553
if event.error_code:
566554
error_event = _create_error_status_event(
@@ -573,7 +561,9 @@ def convert_event_to_a2a_events(
573561
event,
574562
invocation_context,
575563
part_converter=part_converter,
576-
role=Role.user if event.author == "user" else Role.agent,
564+
role=_compat.ROLE_USER
565+
if event.author == "user"
566+
else _compat.ROLE_AGENT,
577567
)
578568
if message:
579569
running_event = _create_status_update_event(

0 commit comments

Comments
 (0)