From 4c5bfc757050f2ef9bb45c2cd709eb2e9435445b Mon Sep 17 00:00:00 2001 From: Kamil Tomaszek Date: Tue, 30 Jun 2026 04:27:00 -0700 Subject: [PATCH] feat(a2a): add an A2A v1.0 client with Android support and deprecate the v0.3 JvmA2AAgent PiperOrigin-RevId: 940379674 --- a2a/build.gradle.kts | 38 +- .../adk/kt/a2a/android/AndroidA2AAgent.kt | 92 ++ .../a2a/android/JsonRpcHttpClientTransport.kt | 230 ++++ .../JsonRpcHttpClientTransportConfig.kt | 24 + .../JsonRpcHttpClientTransportProvider.kt | 47 + ...ient.transport.spi.ClientTransportProvider | 1 + .../adk/kt/a2a/agent/A2AAgentAndroidTest.kt | 208 ++++ .../a2a/agent/AgentCardResolverAndroidTest.kt | 141 +++ .../android/JsonRpcHttpClientTransportTest.kt | 321 +++++ .../google/adk/kt/a2a/agent/A2AAgentImpl.kt | 225 ++++ .../adk/kt/a2a/agent/AgentCardResolver.kt | 66 + .../adk/kt/a2a/converters/A2aConverters.kt | 454 +++++++ .../adk/kt/a2a/agent/A2AAgentImplTest.kt | 770 ++++++++++++ .../kt/a2a/converters/A2aConvertersTest.kt | 1085 +++++++++++++++++ .../google/adk/kt/a2a/agent/LegacyA2AAgent.kt | 4 + .../kt/a2a/converters/LegacyA2aConverters.kt | 5 - .../com/google/adk/kt/a2a/jvm/A2AAgent.kt | 81 ++ .../adk/kt/a2a/agent/LegacyA2AAgentTest.kt | 1 + .../com/google/adk/kt/a2a/jvm/A2AAgentTest.kt | 79 ++ .../adk/kt/examples/a2a/A2AAgentDemo.kt | 44 + .../adk/kt/examples/a2a/JvmA2AAgentDemo.kt | 77 -- gradle/libs.versions.toml | 13 +- 22 files changed, 3911 insertions(+), 95 deletions(-) create mode 100644 a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/AndroidA2AAgent.kt create mode 100644 a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransport.kt create mode 100644 a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportConfig.kt create mode 100644 a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportProvider.kt create mode 100644 a2a/src/androidMain/resources/META-INF/services/org.a2aproject.sdk.client.transport.spi.ClientTransportProvider create mode 100644 a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/agent/A2AAgentAndroidTest.kt create mode 100644 a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/agent/AgentCardResolverAndroidTest.kt create mode 100644 a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportTest.kt create mode 100644 a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/agent/A2AAgentImpl.kt create mode 100644 a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/agent/AgentCardResolver.kt create mode 100644 a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/converters/A2aConverters.kt create mode 100644 a2a/src/commonJvmAndroidTest/kotlin/com/google/adk/kt/a2a/agent/A2AAgentImplTest.kt create mode 100644 a2a/src/commonJvmAndroidTest/kotlin/com/google/adk/kt/a2a/converters/A2aConvertersTest.kt create mode 100644 a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/jvm/A2AAgent.kt create mode 100644 a2a/src/jvmTest/kotlin/com/google/adk/kt/a2a/jvm/A2AAgentTest.kt create mode 100644 examples/src/main/kotlin/com/google/adk/kt/examples/a2a/A2AAgentDemo.kt delete mode 100644 examples/src/main/kotlin/com/google/adk/kt/examples/a2a/JvmA2AAgentDemo.kt diff --git a/a2a/build.gradle.kts b/a2a/build.gradle.kts index 5fe089c5..65466789 100644 --- a/a2a/build.gradle.kts +++ b/a2a/build.gradle.kts @@ -16,10 +16,19 @@ plugins { kotlin("multiplatform") + id("com.android.kotlin.multiplatform.library") id("maven-publish") } kotlin { + // AGP 9 KMP Android library target (replaces com.android.library + androidTarget). + android { + namespace = "com.google.adk.a2a" + compileSdk = rootProject.extra["androidCompileSdk"] as Int + // A2A requires API 35: the SDK's ClientBuilder.build() calls ServiceLoader.stream() (API 35+). + // Core stays at the shared androidMinSdk; only A2A gates higher. + minSdk = 35 + } jvm() sourceSets { @@ -38,14 +47,22 @@ kotlin { val commonJvmAndroidMain by creating { dependsOn(commonMain) dependencies { - implementation(libs.jackson.databind) - implementation(libs.jackson.datatype.jsr310) + implementation(libs.kotlinx.serialization) + implementation(libs.a2a.sdk.client) + implementation(libs.a2a.sdk.common) + implementation(libs.a2a.sdk.spec) } } - // jvmMain: deprecated v0.3 (`io.a2a.*`) path, JVM-only. + // jvmMain hosts the deprecated v0.3 path (JVM-only); androidMain stays v1.0-only. val jvmMain by getting { dependsOn(commonJvmAndroidMain) dependencies { + // JVM v1.0 factory uses the SDK's proto-based JSON-RPC transport; kept off the Android + // path. + implementation(libs.a2a.sdk.transport.jsonrpc) + // Jackson is JVM-only (deprecated v0.3 converters); kept off the Android artifact. + implementation(libs.jackson.databind) + implementation(libs.jackson.datatype.jsr310) implementation(libs.jackson.module.kotlin) implementation(libs.a2a.legacy.sdk.client) implementation(libs.a2a.legacy.sdk.common) @@ -58,19 +75,26 @@ kotlin { implementation(libs.google.truth) implementation(libs.mockito.kotlin) implementation(libs.kotlinx.coroutines.test) + implementation(libs.okhttp.mockwebserver) implementation(libs.a2a.legacy.sdk.client) implementation(libs.a2a.legacy.sdk.spec) } } + val androidMain by getting { + dependsOn(commonJvmAndroidMain) + dependencies { implementation(libs.a2a.sdk.http.client.android) } + } } } // Coordinates the Kotlin Multiplatform plugin uses for the publications it // auto-creates: -// - `kotlinMultiplatform` -> google-adk-kotlin-a2a (root metadata) -// - `jvm` -> google-adk-kotlin-a2a-jvm (KMP target) -// POM metadata, Dokka javadoc, and GPG signing are configured in the root -// build.gradle.kts. +// - `kotlinMultiplatform` -> google-adk-kotlin-a2a (root metadata) +// - `jvm` -> google-adk-kotlin-a2a-jvm (KMP target) +// - `androidRelease` -> google-adk-kotlin-a2a-android (KMP target) +// Per-target suffixes (`-jvm`, `-android`) are appended by the KMP plugin +// automatically. POM metadata, Dokka javadoc, and GPG signing are configured in +// the root build.gradle.kts. publishing { publications.withType().configureEach { if (name == "kotlinMultiplatform") { diff --git a/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/AndroidA2AAgent.kt b/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/AndroidA2AAgent.kt new file mode 100644 index 00000000..b281f758 --- /dev/null +++ b/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/AndroidA2AAgent.kt @@ -0,0 +1,92 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.android + +import com.google.adk.kt.a2a.agent.A2AAgentImpl +import com.google.adk.kt.a2a.agent.BaseRemoteA2AAgent +import com.google.adk.kt.a2a.agent.resolveAgentCard +import com.google.adk.kt.agents.BaseAgent +import com.google.adk.kt.callbacks.AfterAgentCallback +import com.google.adk.kt.callbacks.BeforeAgentCallback +import org.a2aproject.sdk.client.Client +import org.a2aproject.sdk.client.config.ClientConfig +import org.a2aproject.sdk.client.http.A2AHttpClient +import org.a2aproject.sdk.client.http.AndroidA2AHttpClient +import org.a2aproject.sdk.spec.AgentCard + +/** + * Builds a framework-internal Android A2A [Client] backed by the proto-free, non-streaming + * [JsonRpcHttpClientTransport] (which uses [httpClient], an [AndroidA2AHttpClient] by default). + */ +internal fun androidA2AClient( + agentCard: AgentCard, + httpClient: A2AHttpClient = AndroidA2AHttpClient(), +): Client = + Client.builder(agentCard) + .clientConfig(ClientConfig.Builder().setStreaming(false).build()) + .withTransport( + JsonRpcHttpClientTransport::class.java, + JsonRpcHttpClientTransportConfig(httpClient), + ) + .build() + +/** + * Builds an Android [A2AAgent] from an already-resolved [agentCard], wiring up the Android client + * so the caller never supplies a client and card separately. + * + * The Android proto-free transport supports only non-streaming `message/send`, so the agent always + * runs in non-streaming mode regardless of the remote card's streaming capability. + */ +fun AndroidA2AAgent( + name: String, + agentCard: AgentCard, + httpClient: A2AHttpClient = AndroidA2AHttpClient(), + subAgents: List = emptyList(), + beforeAgentCallbacks: List = emptyList(), + afterAgentCallbacks: List = emptyList(), +): BaseRemoteA2AAgent = + A2AAgentImpl( + name = name, + a2aClient = androidA2AClient(agentCard, httpClient), + agentCard = agentCard, + streaming = false, + subAgents = subAgents, + beforeAgentCallbacks = beforeAgentCallbacks, + afterAgentCallbacks = afterAgentCallbacks, + ) + +/** + * Builds an Android [A2AAgent] from [agentCardUrl], auto-fetching the [AgentCard] from the remote + * agent's `/.well-known/agent-card.json` (like ADK Python/Go). Suspends on the network fetch, so + * call it off the main thread. + */ +suspend fun AndroidA2AAgent( + name: String, + agentCardUrl: String, + httpClient: A2AHttpClient = AndroidA2AHttpClient(), + subAgents: List = emptyList(), + beforeAgentCallbacks: List = emptyList(), + afterAgentCallbacks: List = emptyList(), +): BaseRemoteA2AAgent = + AndroidA2AAgent( + name = name, + agentCard = resolveAgentCard(httpClient, agentCardUrl), + httpClient = httpClient, + subAgents = subAgents, + beforeAgentCallbacks = beforeAgentCallbacks, + afterAgentCallbacks = afterAgentCallbacks, + ) diff --git a/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransport.kt b/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransport.kt new file mode 100644 index 00000000..f1aac66e --- /dev/null +++ b/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransport.kt @@ -0,0 +1,230 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.android + +import com.google.gson.JsonObject +import com.google.gson.JsonParser +import java.util.function.Consumer +import org.a2aproject.sdk.client.http.A2AHttpClient +import org.a2aproject.sdk.client.transport.spi.ClientTransport +import org.a2aproject.sdk.client.transport.spi.interceptors.ClientCallContext +import org.a2aproject.sdk.client.transport.spi.interceptors.ClientCallInterceptor +import org.a2aproject.sdk.client.transport.spi.interceptors.PayloadAndHeaders +import org.a2aproject.sdk.jsonrpc.common.json.JsonProcessingException +import org.a2aproject.sdk.jsonrpc.common.json.JsonUtil +import org.a2aproject.sdk.jsonrpc.common.wrappers.ListTasksResult +import org.a2aproject.sdk.jsonrpc.common.wrappers.SendMessageRequest +import org.a2aproject.sdk.spec.A2AClientException +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.spec.CancelTaskParams +import org.a2aproject.sdk.spec.DeleteTaskPushNotificationConfigParams +import org.a2aproject.sdk.spec.EventKind +import org.a2aproject.sdk.spec.GetExtendedAgentCardParams +import org.a2aproject.sdk.spec.GetTaskPushNotificationConfigParams +import org.a2aproject.sdk.spec.ListTaskPushNotificationConfigsParams +import org.a2aproject.sdk.spec.ListTaskPushNotificationConfigsResult +import org.a2aproject.sdk.spec.ListTasksParams +import org.a2aproject.sdk.spec.MessageSendParams +import org.a2aproject.sdk.spec.StreamingEventKind +import org.a2aproject.sdk.spec.Task +import org.a2aproject.sdk.spec.TaskIdParams +import org.a2aproject.sdk.spec.TaskPushNotificationConfig +import org.a2aproject.sdk.spec.TaskQueryParams + +/** + * A proto-free [ClientTransport] that runs a real non-streaming `message/send` JSON-RPC round-trip + * over an injected [A2AHttpClient]. + * + * The SDK's own JSON-RPC transport can't run on Android: it serializes with protobuf's + * `JsonFormat`, which needs full protobuf, not the proto-lite runtime Android uses. So this does + * the round-trip by hand with the SDK's proto-free `jsonrpccommon` [JsonUtil]. The remaining + * [ClientTransport] methods throw [UnsupportedOperationException]. + * + * Any [ClientCallInterceptor]s configured on the SDK `Client` are applied to `message/send`: their + * headers (auth, logging, tracing) are attached to the outgoing request. See + * [applyInterceptorHeaders]. + * + * Internal to this module; wired up by `androidA2AClient(...)`. + */ +internal class JsonRpcHttpClientTransport( + private val httpClient: A2AHttpClient, + private val url: String, + private val agentCard: AgentCard? = null, + private val interceptors: List = emptyList(), +) : ClientTransport { + + override fun sendMessage(request: MessageSendParams, context: ClientCallContext?): EventKind { + val body = serializeSendMessage(request) + val interceptorHeaders = applyInterceptorHeaders(SEND_MESSAGE_METHOD, request, context) + + try { + val response = + httpClient + .createPost() + .url(url) + .addHeader(A2AHttpClient.CONTENT_TYPE, A2AHttpClient.APPLICATION_JSON) + .addHeader(A2A_VERSION_HEADER, A2A_VERSION) + .addHeaders(interceptorHeaders) + .body(body) + .post() + if (response.status() < 200 || response.status() >= 300) { + throw A2AClientException("Unexpected HTTP status: ${response.status()}") + } + return parseSendMessageResponse(response.body()) + } catch (e: A2AClientException) { + throw e + } catch (e: InterruptedException) { + Thread.currentThread().interrupt() + throw A2AClientException("Android A2A HTTP round-trip interrupted", e) + } catch (e: Exception) { + throw A2AClientException("Android A2A HTTP round-trip failed", e) + } + } + + /** + * Serializes a `message/send` request to its JSON-RPC wire form. + * + * [JsonUtil] serializes a `Message` (a `StreamingEventKind`) with a `{"": {...}}` wrapper, + * which double-wraps `params.message`; a request needs the bare message, so we strip one layer: + * ``` + * JsonUtil: "params": { "message": { "message": {...} } } + * wire form: "params": { "message": {...} } + * ``` + */ + private fun serializeSendMessage(request: MessageSendParams): String = + try { + val envelope: JsonObject = + JsonParser.parseString( + JsonUtil.toJson(SendMessageRequest(JSONRPC_VERSION, REQUEST_ID, request)) + ) + .asJsonObject + val params = envelope.getAsJsonObject("params") + params.add("message", params.getAsJsonObject("message").get("message")) + envelope.toString() + } catch (e: JsonProcessingException) { + throw A2AClientException("Failed to serialize A2A request", e) + } + + /** + * Runs the configured [interceptors] and returns the HTTP headers to attach to the request. + * + * Mirrors the SDK's JSON-RPC transport: it seeds the headers from the call [context], then lets + * each interceptor add to them in order. Interceptors can also rewrite the payload, but that path + * is proto-based and doesn't apply to this proto-free transport, so only the headers are used. + */ + private fun applyInterceptorHeaders( + methodName: String, + payload: Any?, + context: ClientCallContext?, + ): Map { + var payloadAndHeaders = PayloadAndHeaders(payload, context?.headers) + for (interceptor in interceptors) { + payloadAndHeaders = + interceptor.intercept( + methodName, + payloadAndHeaders.payload, + payloadAndHeaders.headers, + agentCard, + context, + ) + } + return payloadAndHeaders.headers + } + + // --- Unused operations ------------------------------------------------------------------------- + + override fun sendMessageStreaming( + request: MessageSendParams, + eventConsumer: Consumer, + errorConsumer: Consumer, + context: ClientCallContext?, + ): Unit = + throw UnsupportedOperationException("streaming not supported by JsonRpcHttpClientTransport") + + override fun getTask(request: TaskQueryParams, context: ClientCallContext?): Task = + throw UnsupportedOperationException() + + override fun cancelTask(request: CancelTaskParams, context: ClientCallContext?): Task = + throw UnsupportedOperationException() + + override fun listTasks(request: ListTasksParams, context: ClientCallContext?): ListTasksResult = + throw UnsupportedOperationException() + + override fun createTaskPushNotificationConfiguration( + request: TaskPushNotificationConfig, + context: ClientCallContext?, + ): TaskPushNotificationConfig = throw UnsupportedOperationException() + + override fun getTaskPushNotificationConfiguration( + request: GetTaskPushNotificationConfigParams, + context: ClientCallContext?, + ): TaskPushNotificationConfig = throw UnsupportedOperationException() + + override fun listTaskPushNotificationConfigurations( + request: ListTaskPushNotificationConfigsParams, + context: ClientCallContext?, + ): ListTaskPushNotificationConfigsResult = throw UnsupportedOperationException() + + override fun deleteTaskPushNotificationConfigurations( + request: DeleteTaskPushNotificationConfigParams, + context: ClientCallContext?, + ): Unit = throw UnsupportedOperationException() + + override fun subscribeToTask( + request: TaskIdParams, + eventConsumer: Consumer, + errorConsumer: Consumer, + context: ClientCallContext?, + ): Unit = throw UnsupportedOperationException() + + override fun getExtendedAgentCard( + params: GetExtendedAgentCardParams, + context: ClientCallContext?, + ): AgentCard = throw UnsupportedOperationException() + + override fun close() {} + + private companion object { + const val JSONRPC_VERSION = "2.0" + const val REQUEST_ID = "1" + const val A2A_VERSION_HEADER = "A2A-Version" + const val A2A_VERSION = "1.0" + const val SEND_MESSAGE_METHOD = "message/send" + + /** Parses a JSON-RPC `message/send` response body into its result [EventKind]. */ + fun parseSendMessageResponse(responseBody: String): EventKind { + val envelope = JsonParser.parseString(responseBody).asJsonObject + + val errorNode = envelope.get("error") + if (errorNode != null && !errorNode.isJsonNull) { + throw A2AClientException("A2A JSON-RPC error: $errorNode") + } + + val resultNode = envelope.get("result") + if (resultNode == null || !resultNode.isJsonObject) { + throw A2AClientException("A2A JSON-RPC response missing 'result' object") + } + + return try { + // Result is a Task/Message; the SDK's StreamingEventKind adapter picks the concrete type. + JsonUtil.fromJson(resultNode.toString(), StreamingEventKind::class.java) as EventKind + } catch (e: JsonProcessingException) { + throw A2AClientException("Failed to parse A2A response result", e) + } + } + } +} diff --git a/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportConfig.kt b/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportConfig.kt new file mode 100644 index 00000000..88856a14 --- /dev/null +++ b/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportConfig.kt @@ -0,0 +1,24 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.android + +import org.a2aproject.sdk.client.http.A2AHttpClient +import org.a2aproject.sdk.client.transport.spi.ClientTransportConfig + +/** Config carrying the [A2AHttpClient] for [JsonRpcHttpClientTransport]. */ +internal class JsonRpcHttpClientTransportConfig(val httpClient: A2AHttpClient) : + ClientTransportConfig() diff --git a/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportProvider.kt b/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportProvider.kt new file mode 100644 index 00000000..782a0865 --- /dev/null +++ b/a2a/src/androidMain/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportProvider.kt @@ -0,0 +1,47 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.android + +import org.a2aproject.sdk.client.transport.spi.ClientTransportProvider +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.spec.AgentInterface +import org.a2aproject.sdk.spec.TransportProtocol + +/** + * SPI provider that lets `Client.builder(card).withTransport(...)` build a + * [JsonRpcHttpClientTransport] via the SDK's public builder. Discovered through `ServiceLoader`. + */ +internal class JsonRpcHttpClientTransportProvider : + ClientTransportProvider { + + override fun create( + config: JsonRpcHttpClientTransportConfig, + agentCard: AgentCard, + agentInterface: AgentInterface, + ): JsonRpcHttpClientTransport = + JsonRpcHttpClientTransport( + config.httpClient, + agentInterface.url(), + agentCard, + config.interceptors, + ) + + override fun getTransportProtocol(): String = TransportProtocol.JSONRPC.asString() + + override fun getTransportProtocolClass(): Class = + JsonRpcHttpClientTransport::class.java +} diff --git a/a2a/src/androidMain/resources/META-INF/services/org.a2aproject.sdk.client.transport.spi.ClientTransportProvider b/a2a/src/androidMain/resources/META-INF/services/org.a2aproject.sdk.client.transport.spi.ClientTransportProvider new file mode 100644 index 00000000..80dd4a21 --- /dev/null +++ b/a2a/src/androidMain/resources/META-INF/services/org.a2aproject.sdk.client.transport.spi.ClientTransportProvider @@ -0,0 +1 @@ +com.google.adk.kt.a2a.android.JsonRpcHttpClientTransportProvider diff --git a/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/agent/A2AAgentAndroidTest.kt b/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/agent/A2AAgentAndroidTest.kt new file mode 100644 index 00000000..97292228 --- /dev/null +++ b/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/agent/A2AAgentAndroidTest.kt @@ -0,0 +1,208 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.agent + +import androidx.test.ext.junit.runners.AndroidJUnit4 +import com.google.adk.kt.a2a.android.AndroidA2AAgent +import com.google.adk.kt.agents.InvocationContext +import com.google.adk.kt.events.Event +import com.google.adk.kt.sessions.Session +import com.google.adk.kt.sessions.SessionKey +import com.google.adk.kt.testing.DummyAgent +import com.google.adk.kt.testing.userMessage +import com.google.common.truth.Truth.assertThat +import kotlinx.coroutines.flow.toList +import kotlinx.coroutines.test.runTest +import okhttp3.mockwebserver.MockResponse +import okhttp3.mockwebserver.MockWebServer +import okhttp3.mockwebserver.RecordedRequest +import org.a2aproject.sdk.jsonrpc.common.json.JsonUtil +import org.a2aproject.sdk.jsonrpc.common.wrappers.SendMessageResponse +import org.a2aproject.sdk.spec.AgentCapabilities +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.spec.AgentInterface +import org.a2aproject.sdk.spec.Message +import org.a2aproject.sdk.spec.Task +import org.a2aproject.sdk.spec.TaskState +import org.a2aproject.sdk.spec.TaskStatus +import org.a2aproject.sdk.spec.TextPart +import org.a2aproject.sdk.spec.TransportProtocol +import org.junit.After +import org.junit.Before +import org.junit.Test +import org.junit.runner.RunWith + +/** + * Builds an Android A2A agent via [AndroidA2AAgent] and runs a round-trip against a [MockWebServer] + * on the Robolectric runtime, exercising the proto-free Android transport. + * + * The MockWebServer returns a real JSON-RPC response built with the SDK's own + * [JsonUtil] + [SendMessageResponse], and the transport parses it back with [JsonUtil], so this + * exercises the actual proto-free A2A serialization on the Android runtime in both directions. + */ +@RunWith(AndroidJUnit4::class) +class A2AAgentAndroidTest { + + private lateinit var server: MockWebServer + + @Before + fun setUp() { + server = MockWebServer() + server.start() + } + + @After + fun tearDown() { + server.shutdown() + } + + @Test + fun runAsync_androidHttpClient_realRoundTrip_emitsAgentEventAndSendsRequest() = runTest { + val agentReply = "Hello from the Android A2A agent!" + + // Build a real JSON-RPC `message/send` response with the SDK's own proto-free serialization: + // a completed Task whose status message carries the agent text. Using a completed Task (rather + // than a bare Message) lets the agent's non-streaming flow recognise the turn as terminal. + val agentMessage = + Message.builder() + .messageId("agent-message-1") + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart(agentReply))) + .build() + val responseTask = + Task.builder() + .id("android-task-1") + .contextId("android-context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED, agentMessage, null)) + .build() + val responseBody = JsonUtil.toJson(SendMessageResponse("2.0", "1", responseTask, null)) + server.enqueue(MockResponse().setBody(responseBody)) + val serverUrl = server.url("/a2a").toString() + + val agentCard = + AgentCard.builder() + .name("remote-agent") + .description("Remote Agent") + .url(serverUrl) + .version("1.0.0") + .defaultInputModes(listOf("text")) + .defaultOutputModes(listOf("text")) + .skills(listOf()) + .supportedInterfaces( + listOf(AgentInterface(TransportProtocol.JSONRPC.asString(), serverUrl)) + ) + .capabilities(AgentCapabilities.builder().streaming(false).build()) + .build() + + val agent = AndroidA2AAgent(name = "remote-agent", agentCard = agentCard) + + val session = + Session( + key = SessionKey(appName = "demo", userId = "user", id = "session-1"), + events = + mutableListOf( + Event(invocationId = "invocation-0", author = "user", content = userMessage("hello")) + ), + ) + val context = InvocationContext(agent = DummyAgent(), session = session, runConfig = null) + + val events = agent.runAsync(context).toList() + + // The agent emitted an ADK Event carrying the agent text returned over the real HTTP + // round-trip. + val emittedTexts = events.mapNotNull { it.content?.parts?.firstOrNull()?.text } + assertThat(emittedTexts).contains(agentReply) + + // The user message actually traversed AndroidA2AHttpClient and reached the server. + val recorded = server.takeRecordedRequestOrFail() + assertThat(recorded.path).isEqualTo("/a2a") + assertThat(recorded.method).isEqualTo("POST") + assertThat(recorded.getHeader("A2A-Version")).isEqualTo("1.0") + val body = recorded.body.readUtf8() + assertThat(body).contains("hello") + // params.message must be the flat message, not double-wrapped by the wrapper adapter. + assertThat(body).doesNotContain("\"message\":{\"message\"") + } + + @Test + fun runAsync_streamingCapableCard_usesNonStreamingWithoutCrashing() = runTest { + // Regression: a card advertising streaming used to route the SDK Client to + // `sendMessageStreaming`, which the proto-free Android transport does not implement (it threw + // UnsupportedOperationException). The Android agent now forces non-streaming, so `message/send` + // is used and the round-trip completes. + val agentReply = "Hello from a streaming-capable agent!" + val agentMessage = + Message.builder() + .messageId("agent-message-1") + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart(agentReply))) + .build() + val responseTask = + Task.builder() + .id("android-task-1") + .contextId("android-context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED, agentMessage, null)) + .build() + val responseBody = JsonUtil.toJson(SendMessageResponse("2.0", "1", responseTask, null)) + server.enqueue(MockResponse().setBody(responseBody)) + val serverUrl = server.url("/a2a").toString() + + val agentCard = + AgentCard.builder() + .name("remote-agent") + .description("Remote Agent") + .url(serverUrl) + .version("1.0.0") + .defaultInputModes(listOf("text")) + .defaultOutputModes(listOf("text")) + .skills(listOf()) + .supportedInterfaces( + listOf(AgentInterface(TransportProtocol.JSONRPC.asString(), serverUrl)) + ) + // Streaming-capable card: the case that previously crashed the Android transport. + .capabilities(AgentCapabilities.builder().streaming(true).build()) + .build() + + val agent = AndroidA2AAgent(name = "remote-agent", agentCard = agentCard) + assertThat(agent.isStreamingEnabled).isFalse() + + val session = + Session( + key = SessionKey(appName = "demo", userId = "user", id = "session-1"), + events = + mutableListOf( + Event(invocationId = "invocation-0", author = "user", content = userMessage("hello")) + ), + ) + val context = InvocationContext(agent = DummyAgent(), session = session, runConfig = null) + + // Would throw UnsupportedOperationException before the fix; now completes over `message/send`. + val events = agent.runAsync(context).toList() + + val emittedTexts = events.mapNotNull { it.content?.parts?.firstOrNull()?.text } + assertThat(emittedTexts).contains(agentReply) + assertThat(server.takeRecordedRequestOrFail().path).isEqualTo("/a2a") + } +} + +private fun MockWebServer.takeRecordedRequestOrFail(): RecordedRequest = + try { + takeRequest() + } catch (e: InterruptedException) { + Thread.currentThread().interrupt() + throw AssertionError("Interrupted while waiting for the recorded request", e) + } diff --git a/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/agent/AgentCardResolverAndroidTest.kt b/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/agent/AgentCardResolverAndroidTest.kt new file mode 100644 index 00000000..3276e7de --- /dev/null +++ b/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/agent/AgentCardResolverAndroidTest.kt @@ -0,0 +1,141 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.agent + +import androidx.test.ext.junit.runners.AndroidJUnit4 +import com.google.adk.kt.a2a.android.AndroidA2AAgent +import com.google.common.truth.Truth.assertThat +import kotlinx.coroutines.test.runTest +import okhttp3.mockwebserver.MockResponse +import okhttp3.mockwebserver.MockWebServer +import org.a2aproject.sdk.client.http.AndroidA2AHttpClient +import org.a2aproject.sdk.jsonrpc.common.json.JsonUtil +import org.a2aproject.sdk.spec.AgentCapabilities +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.spec.AgentInterface +import org.a2aproject.sdk.spec.TransportProtocol +import org.junit.After +import org.junit.Before +import org.junit.Test +import org.junit.runner.RunWith + +/** + * Verifies the proto-free Android agent-card auto-fetch: [resolveAgentCard] and [AndroidA2AAgent] + * fetch and parse a card from the standard `/.well-known/agent-card.json` endpoint (served here by + * [MockWebServer] on the Robolectric runtime). + */ +@RunWith(AndroidJUnit4::class) +class AgentCardResolverAndroidTest { + + private lateinit var server: MockWebServer + + @Before + fun setUp() { + server = MockWebServer() + server.start() + } + + @After + fun tearDown() { + server.shutdown() + } + + private fun agentCard(url: String): AgentCard = + AgentCard.builder() + .name("remote-agent") + .description("Remote Agent") + .url(url) + .version("1.0.0") + .defaultInputModes(listOf("text")) + .defaultOutputModes(listOf("text")) + .skills(listOf()) + .supportedInterfaces(listOf(AgentInterface(TransportProtocol.JSONRPC.asString(), url))) + .capabilities(AgentCapabilities.builder().streaming(false).build()) + .build() + + @Test + fun resolveAgentCard_fetchesAndParsesFromWellKnownEndpoint() = runTest { + val baseUrl = server.url("/").toString() + server.enqueue(MockResponse().setBody(JsonUtil.toJson(agentCard(baseUrl)))) + + val resolved = resolveAgentCard(AndroidA2AHttpClient(), baseUrl) + + assertThat(resolved.name()).isEqualTo("remote-agent") + assertThat(resolved.description()).isEqualTo("Remote Agent") + assertThat(recordedPath()).isEqualTo("/.well-known/agent-card.json") + } + + @Test + fun androidA2AAgent_autoFetchesCard_andPopulatesDescription() = runTest { + val baseUrl = server.url("/").toString() + server.enqueue(MockResponse().setBody(JsonUtil.toJson(agentCard(baseUrl)))) + + val agent = AndroidA2AAgent(name = "remote-agent", agentCardUrl = baseUrl) + + assertThat(agent.description).isEqualTo("Remote Agent") + } + + @Test + fun resolveAgentCard_httpError_throwsResolutionError() = runTest { + server.enqueue(MockResponse().setResponseCode(500)) + val baseUrl = server.url("/").toString() + + val e = runCatching { resolveAgentCard(AndroidA2AHttpClient(), baseUrl) }.exceptionOrNull() + + assertThat(e).isInstanceOf(BaseRemoteA2AAgent.AgentCardResolutionError::class.java) + assertThat(e).hasMessageThat().contains("Failed to fetch agent card") + } + + @Test + fun resolveAgentCard_malformedBody_throwsResolutionError() = runTest { + server.enqueue(MockResponse().setBody("not json")) + val baseUrl = server.url("/").toString() + + val e = runCatching { resolveAgentCard(AndroidA2AHttpClient(), baseUrl) }.exceptionOrNull() + + assertThat(e).isInstanceOf(BaseRemoteA2AAgent.AgentCardResolutionError::class.java) + assertThat(e).hasMessageThat().contains("Failed to parse agent card") + } + + @Test + fun resolveAgentCard_emptyBody_throwsResolutionError() = runTest { + // A 200 with an empty body parses to null; must surface as AgentCardResolutionError, not an + // NPE. + server.enqueue(MockResponse().setBody("")) + val baseUrl = server.url("/").toString() + + val e = runCatching { resolveAgentCard(AndroidA2AHttpClient(), baseUrl) }.exceptionOrNull() + + assertThat(e).isInstanceOf(BaseRemoteA2AAgent.AgentCardResolutionError::class.java) + assertThat(e).hasMessageThat().contains("Empty agent card response") + } + + @Test + fun resolveAgentCard_networkFailure_throwsResolutionError() = runTest { + // Point at a server that is immediately shut down, so the GET fails with an IOException. + val deadServer = MockWebServer().apply { start() } + val baseUrl = deadServer.url("/").toString() + deadServer.shutdown() + + val e = runCatching { resolveAgentCard(AndroidA2AHttpClient(), baseUrl) }.exceptionOrNull() + + assertThat(e).isInstanceOf(BaseRemoteA2AAgent.AgentCardResolutionError::class.java) + assertThat(e).hasMessageThat().contains("Failed to fetch agent card") + } + + private fun recordedPath(): String? = server.takeRequest().path +} diff --git a/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportTest.kt b/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportTest.kt new file mode 100644 index 00000000..ef968742 --- /dev/null +++ b/a2a/src/androidRobolectricTest/kotlin/com/google/adk/kt/a2a/android/JsonRpcHttpClientTransportTest.kt @@ -0,0 +1,321 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.android + +import androidx.test.ext.junit.runners.AndroidJUnit4 +import com.google.common.truth.Truth.assertThat +import okhttp3.mockwebserver.MockResponse +import okhttp3.mockwebserver.MockWebServer +import org.a2aproject.sdk.client.http.A2AHttpClient +import org.a2aproject.sdk.client.http.AndroidA2AHttpClient +import org.a2aproject.sdk.client.transport.spi.interceptors.ClientCallContext +import org.a2aproject.sdk.client.transport.spi.interceptors.ClientCallInterceptor +import org.a2aproject.sdk.client.transport.spi.interceptors.PayloadAndHeaders +import org.a2aproject.sdk.jsonrpc.common.json.JsonUtil +import org.a2aproject.sdk.jsonrpc.common.wrappers.SendMessageResponse +import org.a2aproject.sdk.spec.A2AClientException +import org.a2aproject.sdk.spec.AgentCapabilities +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.spec.AgentInterface +import org.a2aproject.sdk.spec.Message +import org.a2aproject.sdk.spec.MessageSendParams +import org.a2aproject.sdk.spec.Task +import org.a2aproject.sdk.spec.TaskState +import org.a2aproject.sdk.spec.TaskStatus +import org.a2aproject.sdk.spec.TextPart +import org.junit.After +import org.junit.Assert.assertThrows +import org.junit.Before +import org.junit.Test +import org.junit.runner.RunWith + +/** + * Directly exercises [JsonRpcHttpClientTransport] over a [MockWebServer] on the Robolectric + * runtime. + */ +@RunWith(AndroidJUnit4::class) +class JsonRpcHttpClientTransportTest { + + private lateinit var server: MockWebServer + + @Before + fun setUp() { + server = MockWebServer() + server.start() + } + + @After + fun tearDown() { + server.shutdown() + } + + private fun transport(): JsonRpcHttpClientTransport = + JsonRpcHttpClientTransport(AndroidA2AHttpClient(), server.url("/a2a").toString()) + + private fun sendMessageParams(): MessageSendParams { + val message = + Message.builder() + .messageId("req-1") + .role(Message.Role.ROLE_USER) + .parts(listOf(TextPart("hello"))) + .build() + return MessageSendParams.builder().message(message).build() + } + + @Test + fun sendMessage_taskResponse_returnsTask() { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .build() + server.enqueue( + MockResponse().setBody(JsonUtil.toJson(SendMessageResponse("2.0", "1", task, null))) + ) + + val result = transport().sendMessage(sendMessageParams(), null) + + assertThat(result).isInstanceOf(Task::class.java) + assertThat((result as Task).id).isEqualTo("task-1") + } + + @Test + fun sendMessage_sendsVersionHeaderAndFlatMessage() { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .build() + server.enqueue( + MockResponse().setBody(JsonUtil.toJson(SendMessageResponse("2.0", "1", task, null))) + ) + + assertThat(transport().sendMessage(sendMessageParams(), null)).isNotNull() + + val recorded = server.takeRequest() + assertThat(recorded.getHeader("A2A-Version")).isEqualTo("1.0") + // params.message is the flat message, not double-wrapped by the wrapper adapter. + assertThat(recorded.body.readUtf8()).doesNotContain("\"message\":{\"message\"") + } + + @Test + fun sendMessage_messageResponse_returnsMessage() { + val message = + Message.builder() + .messageId("msg-1") + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart("hi"))) + .build() + server.enqueue( + MockResponse().setBody(JsonUtil.toJson(SendMessageResponse("2.0", "1", message, null))) + ) + + val result = transport().sendMessage(sendMessageParams(), null) + + assertThat(result).isInstanceOf(Message::class.java) + assertThat((result as Message).messageId).isEqualTo("msg-1") + } + + @Test + fun sendMessage_non2xxStatus_throws() { + server.enqueue(MockResponse().setResponseCode(500).setBody("{}")) + + assertThrows(A2AClientException::class.java) { + transport().sendMessage(sendMessageParams(), null) + } + } + + @Test + fun sendMessage_jsonRpcError_throws() { + server.enqueue( + MockResponse() + .setBody( + "{\"jsonrpc\":\"2.0\",\"id\":\"1\",\"error\":{\"code\":-32000,\"message\":\"boom\"}}" + ) + ) + + val e = + assertThrows(A2AClientException::class.java) { + transport().sendMessage(sendMessageParams(), null) + } + assertThat(e).hasMessageThat().contains("A2A JSON-RPC error") + } + + @Test + fun sendMessage_missingResult_throws() { + server.enqueue(MockResponse().setBody("{\"jsonrpc\":\"2.0\",\"id\":\"1\"}")) + + val e = + assertThrows(A2AClientException::class.java) { + transport().sendMessage(sendMessageParams(), null) + } + assertThat(e).hasMessageThat().contains("missing 'result'") + } + + @Test + fun sendMessage_nonObjectResult_throws() { + server.enqueue(MockResponse().setBody("{\"jsonrpc\":\"2.0\",\"id\":\"1\",\"result\":\"oops\"}")) + + val e = + assertThrows(A2AClientException::class.java) { + transport().sendMessage(sendMessageParams(), null) + } + assertThat(e).hasMessageThat().contains("missing 'result'") + } + + @Test + fun sendMessage_malformedJson_throws() { + server.enqueue(MockResponse().setBody("not json")) + + assertThrows(A2AClientException::class.java) { + transport().sendMessage(sendMessageParams(), null) + } + } + + private fun agentInterface(): AgentInterface = + AgentInterface("JSONRPC", server.url("/a2a").toString()) + + private fun agentCard(): AgentCard = + AgentCard.builder() + .name("remote-agent") + .description("Remote Agent") + .url(server.url("/a2a").toString()) + .version("1.0.0") + .defaultInputModes(listOf("text")) + .defaultOutputModes(listOf("text")) + .skills(listOf()) + .supportedInterfaces(listOf(agentInterface())) + .capabilities(AgentCapabilities.builder().streaming(false).build()) + .build() + + @Test + fun sendMessage_usesInjectedHttpClient() { + val marker = RuntimeException("injected client used") + val injected = + object : A2AHttpClient { + override fun createGet(): A2AHttpClient.GetBuilder = throw marker + + override fun createPost(): A2AHttpClient.PostBuilder = throw marker + + override fun createDelete(): A2AHttpClient.DeleteBuilder = throw marker + } + + val transport = JsonRpcHttpClientTransport(injected, server.url("/a2a").toString()) + val thrown = + assertThrows(A2AClientException::class.java) { + transport.sendMessage(sendMessageParams(), null) + } + + assertThat(thrown).hasCauseThat().isSameInstanceAs(marker) + } + + @Test + fun sendMessage_interrupted_restoresFlagAndWraps() { + val interrupting = + object : A2AHttpClient { + override fun createGet(): A2AHttpClient.GetBuilder = throw UnsupportedOperationException() + + override fun createPost(): A2AHttpClient.PostBuilder = throw InterruptedException("boom") + + override fun createDelete(): A2AHttpClient.DeleteBuilder = + throw UnsupportedOperationException() + } + val transport = JsonRpcHttpClientTransport(interrupting, server.url("/a2a").toString()) + + val e = + assertThrows(A2AClientException::class.java) { + transport.sendMessage(sendMessageParams(), null) + } + + assertThat(e).hasMessageThat().contains("interrupted") + assertThat(Thread.interrupted()).isTrue() + } + + @Test + fun sendMessage_withInterceptor_addsReturnedHeadersToRequest() { + enqueueTaskResponse() + val interceptor = RecordingInterceptor("secret-token") + val transport = + JsonRpcHttpClientTransport( + AndroidA2AHttpClient(), + server.url("/a2a").toString(), + agentCard(), + listOf(interceptor), + ) + + assertThat(transport.sendMessage(sendMessageParams(), null)).isNotNull() + + val recorded = server.takeRequest() + assertThat(recorded.getHeader("Authorization")).isEqualTo("Bearer secret-token") + assertThat(interceptor.capturedMethod).isEqualTo("message/send") + } + + @Test + fun sendMessage_withContextHeaders_seedsInterceptorAndRequest() { + enqueueTaskResponse() + val interceptor = RecordingInterceptor("token") + val transport = + JsonRpcHttpClientTransport( + AndroidA2AHttpClient(), + server.url("/a2a").toString(), + agentCard(), + listOf(interceptor), + ) + val context = ClientCallContext(emptyMap(), mapOf("X-From-Context" to "ctx-value")) + + assertThat(transport.sendMessage(sendMessageParams(), context)).isNotNull() + + // The interceptor receives the context headers as its seed input... + assertThat(interceptor.capturedHeaders).containsEntry("X-From-Context", "ctx-value") + // ...and both the seed and interceptor-added headers reach the wire. + val recorded = server.takeRequest() + assertThat(recorded.getHeader("X-From-Context")).isEqualTo("ctx-value") + assertThat(recorded.getHeader("Authorization")).isEqualTo("Bearer token") + } + + private fun enqueueTaskResponse() { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .build() + server.enqueue( + MockResponse().setBody(JsonUtil.toJson(SendMessageResponse("2.0", "1", task, null))) + ) + } + + /** Test interceptor that records its inputs and appends an `Authorization` header. */ + private class RecordingInterceptor(private val token: String) : ClientCallInterceptor() { + var capturedMethod: String? = null + var capturedHeaders: Map = emptyMap() + + override fun intercept( + methodName: String, + payload: Any?, + headers: Map, + agentCard: AgentCard?, + clientCallContext: ClientCallContext?, + ): PayloadAndHeaders { + capturedMethod = methodName + capturedHeaders = headers + return PayloadAndHeaders(payload, headers + ("Authorization" to "Bearer $token")) + } + } +} diff --git a/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/agent/A2AAgentImpl.kt b/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/agent/A2AAgentImpl.kt new file mode 100644 index 00000000..3a817e09 --- /dev/null +++ b/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/agent/A2AAgentImpl.kt @@ -0,0 +1,225 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.agent + +import com.google.adk.kt.a2a.converters.isCompleted +import com.google.adk.kt.a2a.converters.isLastChunk +import com.google.adk.kt.a2a.converters.shouldBuffer +import com.google.adk.kt.a2a.converters.shouldResetBuffer +import com.google.adk.kt.a2a.converters.toA2aMessage +import com.google.adk.kt.a2a.converters.toAdkEvent +import com.google.adk.kt.agents.BaseAgent +import com.google.adk.kt.agents.InvocationContext +import com.google.adk.kt.callbacks.AfterAgentCallback +import com.google.adk.kt.callbacks.BeforeAgentCallback +import com.google.adk.kt.events.Event +import com.google.adk.kt.logging.LoggerFactory +import java.util.concurrent.atomic.AtomicReference +import java.util.function.BiConsumer +import java.util.function.Consumer +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.channels.Channel +import kotlinx.coroutines.channels.awaitClose +import kotlinx.coroutines.flow.Flow +import kotlinx.coroutines.flow.callbackFlow +import kotlinx.coroutines.launch +import org.a2aproject.sdk.client.Client +import org.a2aproject.sdk.client.ClientEvent +import org.a2aproject.sdk.client.MessageEvent +import org.a2aproject.sdk.client.TaskEvent +import org.a2aproject.sdk.client.TaskUpdateEvent +import org.a2aproject.sdk.jsonrpc.common.json.JsonUtil +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.spec.CancelTaskParams +import org.a2aproject.sdk.spec.Message +import org.a2aproject.sdk.spec.TaskState + +/** Agent that communicates with a remote A2A agent via an A2A client. */ +internal class A2AAgentImpl( + name: String, + private val userDescription: String? = null, + private val a2aClient: Client, + private val agentCard: AgentCard, + private val streaming: Boolean = true, + subAgents: List = emptyList(), + beforeAgentCallbacks: List = emptyList(), + afterAgentCallbacks: List = emptyList(), +) : + BaseRemoteA2AAgent( + name = name, + description = userDescription ?: "", + subAgents = subAgents, + beforeAgentCallbacks = beforeAgentCallbacks, + afterAgentCallbacks = afterAgentCallbacks, + ) { + private val logger = LoggerFactory.getLogger(A2AAgentImpl::class) + + override val description: String + get() = userDescription ?: agentCard.description() ?: "" + + override val isStreamingEnabled: Boolean by lazy { + streaming && agentCard.capabilities().streaming() + } + + override fun createA2aCallbackFlow( + context: InvocationContext, + outboundEvent: Event, + ): Flow = callbackFlow { + val activeTaskId = AtomicReference() + val isTaskTerminal = AtomicReference(false) + + // Bridges the non-suspendable Java client handler to suspendable coroutines. + // Note: Channel integration asynchronously defers processing, which removes natural + // backpressure. + // We use an UNLIMITED channel because local processing should outpace remote A2A network + // responses, which are bounded by LLM limits. + // Alternatively, block the thread using: runBlocking { eventChannel.send(responseEvent) } + val eventChannel = Channel(Channel.UNLIMITED) + val message = outboundEvent.toA2aMessage() + + // Suppress because processing is serialized via eventChannel, avoiding concurrent execution. + @Suppress("UnsafeCoroutineCrossing") + val processorJob = launch { + val a2aAggregator = A2AStreamingResponseAggregator(context.invocationId, name) + val debugRequest = serializeMessageToJson(message) + + for (responseEvent in eventChannel) { + val result = processClientEvent(responseEvent, context, debugRequest, a2aAggregator) + + result.taskId?.let { activeTaskId.set(it) } + result.isTerminal?.let { isTaskTerminal.set(it) } + + for (event in result.eventsToEmit) { + send(event) + } + + if (result.shouldClose) { + break + } + } + close() + } + + // This handler may be called multiple times for each response from A2A. + val handler = + BiConsumer { responseEvent, _ -> + // UNLIMITED channel: trySend only fails once the flow is closed; dropping a late event then + // is fine. + eventChannel.trySend(responseEvent).getOrNull() + } + + val errorHandler = + Consumer { ex -> + val e: Throwable? = ex + if (e != null) close(e) else eventChannel.close() + } + + a2aClient.sendMessage(message, listOf(handler), errorHandler, null) + + awaitClose { + eventChannel.close() + processorJob.cancel() + if (!isTaskTerminal.get()) { + activeTaskId.get()?.let { id -> cancelTask(id) } + } + } + } + + private fun processClientEvent( + responseEvent: ClientEvent, + context: InvocationContext, + debugRequest: Result, + a2aAggregator: A2AStreamingResponseAggregator, + ): EventProcessResult { + var taskId: String? = null + var isTerminal: Boolean? = null + + when (responseEvent) { + is TaskEvent -> { + taskId = responseEvent.task.id + isTerminal = isTerminal(responseEvent.task.status.state()) + } + is TaskUpdateEvent -> { + taskId = responseEvent.task.id + isTerminal = isTerminal(responseEvent.task.status.state()) + } + else -> {} + } + + val events = mutableListOf() + val adkEvent = responseEvent.toAdkEvent(context) + if (adkEvent != null) { + events.addAll( + a2aAggregator.processEvent( + adkEvent = addMetadata(adkEvent, responseEvent, debugRequest), + isCompleted = responseEvent.isCompleted(), + shouldBuffer = responseEvent.shouldBuffer(), + shouldResetBuffer = responseEvent.shouldResetBuffer(), + isLastChunk = responseEvent.isLastChunk(), + ) + ) + } + + return EventProcessResult( + eventsToEmit = events, + taskId = taskId, + isTerminal = isTerminal, + shouldClose = responseEvent.isCompleted(), + ) + } + + private fun cancelTask(taskId: String) { + @Suppress("GlobalCoroutineDispatchers", "UnsafeCoroutineCrossing") + val unusedJob = + CoroutineScope(Dispatchers.IO).launch { + try { + a2aClient.cancelTask(CancelTaskParams(taskId)) + } catch (e: Exception) { + logger.warn(e) { "Failed to cancel task $taskId" } + } + } + } + + // Debug metadata via the SDK's reflection-free JsonUtil so it also works under Android R8. + private fun serializeMessageToJson(message: Message): Result = + runCatching { JsonUtil.toJson(message) } + .onFailure { e -> logger.warn(e) { "Failed to serialize request" } } + + private fun addMetadata( + event: Event, + responseEvent: ClientEvent?, + debugRequest: Result, + ): Event { + val debugResponse = responseEvent?.let { + runCatching { serializeClientEvent(it) } + .onFailure { e -> logger.warn(e) { "Failed to serialize response metadata" } } + } + return addA2AMetadata(event = event, debugRequest = debugRequest, debugResponse = debugResponse) + } + + // Unwrap to the underlying spec type, which JsonUtil can serialize (unlike the client wrapper). + private fun serializeClientEvent(event: ClientEvent): String = + when (event) { + is MessageEvent -> JsonUtil.toJson(event.message) + is TaskEvent -> JsonUtil.toJson(event.task) + is TaskUpdateEvent -> JsonUtil.toJson(event.task) + } + + private fun isTerminal(state: TaskState): Boolean = + state.isFinal || state == TaskState.TASK_STATE_INPUT_REQUIRED +} diff --git a/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/agent/AgentCardResolver.kt b/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/agent/AgentCardResolver.kt new file mode 100644 index 00000000..512dd816 --- /dev/null +++ b/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/agent/AgentCardResolver.kt @@ -0,0 +1,66 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.agent + +import com.google.adk.kt.a2a.agent.BaseRemoteA2AAgent.AgentCardResolutionError +import java.io.IOException +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.withContext +import org.a2aproject.sdk.client.http.A2AHttpClient +import org.a2aproject.sdk.client.http.A2AHttpResponse +import org.a2aproject.sdk.jsonrpc.common.json.JsonProcessingException +import org.a2aproject.sdk.jsonrpc.common.json.JsonUtil +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.util.Utils + +/** + * Fetches an [AgentCard] from [agentCardUrl] (a base agent URL or a full card URL) over + * [httpClient], reading the standard `/.well-known/agent-card.json` endpoint. Proto-free: parses + * with the SDK's reflection-free [JsonUtil], so it also works on Android, where the SDK's own + * proto-based `A2ACardResolver` isn't available. + * + * Framework-internal building block behind the `A2AAgent(...)`/`androidA2AAgent(...)` factories. + */ +internal suspend fun resolveAgentCard(httpClient: A2AHttpClient, agentCardUrl: String): AgentCard { + val cardUrl = + Utils.buildCardUrl(Utils.stripWellKnownSuffix(agentCardUrl), Utils.DEFAULT_AGENT_CARD_PATH) + val response = + try { + withContext(Dispatchers.IO) { httpGetCard(httpClient, cardUrl) } + } catch (e: IOException) { + throw AgentCardResolutionError("Failed to fetch agent card from $cardUrl", e) + } catch (e: InterruptedException) { + Thread.currentThread().interrupt() + throw AgentCardResolutionError("Interrupted while fetching agent card from $cardUrl", e) + } + if (!response.success()) { + throw AgentCardResolutionError( + "Failed to fetch agent card from $cardUrl: HTTP ${response.status()}" + ) + } + val card = + try { + JsonUtil.fromJson(response.body(), AgentCard::class.java) + } catch (e: JsonProcessingException) { + throw AgentCardResolutionError("Failed to parse agent card from $cardUrl", e) + } + return card ?: throw AgentCardResolutionError("Empty agent card response from $cardUrl") +} + +// Blocking GET, kept out of the suspend body so the SuspendBlocks lint is satisfied. +private fun httpGetCard(httpClient: A2AHttpClient, cardUrl: String): A2AHttpResponse = + httpClient.createGet().url(cardUrl).addHeader("Accept", A2AHttpClient.APPLICATION_JSON).get() diff --git a/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/converters/A2aConverters.kt b/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/converters/A2aConverters.kt new file mode 100644 index 00000000..2d03c200 --- /dev/null +++ b/a2a/src/commonJvmAndroidMain/kotlin/com/google/adk/kt/a2a/converters/A2aConverters.kt @@ -0,0 +1,454 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +@file:OptIn(FrameworkInternalApi::class) + +package com.google.adk.kt.a2a.converters + +import com.google.adk.kt.agents.InvocationContext +import com.google.adk.kt.annotations.FrameworkInternalApi +import com.google.adk.kt.events.Event +import com.google.adk.kt.ids.Uuid +import com.google.adk.kt.serialization.adkJson +import com.google.adk.kt.serialization.anyToJsonElement +import com.google.adk.kt.serialization.jsonElementToAny +import com.google.adk.kt.types.Blob +import com.google.adk.kt.types.Content +import com.google.adk.kt.types.FileData +import com.google.adk.kt.types.FunctionCall +import com.google.adk.kt.types.FunctionResponse +import com.google.adk.kt.types.GroundingMetadata +import com.google.adk.kt.types.Part +import com.google.adk.kt.types.Role +import com.google.adk.kt.types.UsageMetadata +import java.util.Base64 +import kotlin.reflect.KClass +import kotlinx.serialization.KSerializer +import kotlinx.serialization.json.JsonElement +import org.a2aproject.sdk.client.ClientEvent +import org.a2aproject.sdk.client.MessageEvent +import org.a2aproject.sdk.client.TaskEvent +import org.a2aproject.sdk.client.TaskUpdateEvent +import org.a2aproject.sdk.spec.Artifact +import org.a2aproject.sdk.spec.DataPart +import org.a2aproject.sdk.spec.FileContent +import org.a2aproject.sdk.spec.FilePart +import org.a2aproject.sdk.spec.FileWithBytes +import org.a2aproject.sdk.spec.FileWithUri +import org.a2aproject.sdk.spec.Message +import org.a2aproject.sdk.spec.Part as A2APart +import org.a2aproject.sdk.spec.Task +import org.a2aproject.sdk.spec.TaskArtifactUpdateEvent +import org.a2aproject.sdk.spec.TaskState +import org.a2aproject.sdk.spec.TaskStatusUpdateEvent +import org.a2aproject.sdk.spec.TextPart +import org.slf4j.LoggerFactory + +private val logger = LoggerFactory.getLogger("A2aConverters") + +private val metadataParser = + object : A2AMetadataParser { + override fun parse(metadata: Any?, clazz: KClass): T? { + if (metadata == null) return null + val serializer = serializerFor(clazz) + if (serializer == null) { + logger.warn("No serializer for metadata type ${clazz.simpleName}; skipping") + return null + } + return try { + @Suppress("UNCHECKED_CAST") + adkJson.decodeFromJsonElement(serializer, metadata.toMetadataJsonElement()) as T + } catch (e: Exception) { + logger.warn("Failed to parse metadata of type ${clazz.simpleName}", e) + null + } + } + } + +// A2A metadata is only ever parsed into these two ADK types; see [updateEventMetadata]. +// Visible for testing the defensive null fallback. +internal fun serializerFor(clazz: KClass<*>): KSerializer<*>? = + when (clazz) { + GroundingMetadata::class -> GroundingMetadata.serializer() + UsageMetadata::class -> UsageMetadata.serializer() + else -> null + } + +// A2A metadata values arrive either as a JSON string or as an already-decoded Map/primitive tree. +private fun Any.toMetadataJsonElement(): JsonElement = + if (this is String) adkJson.parseToJsonElement(this) else anyToJsonElement(this) + +private val PENDING_STATES = setOf(TaskState.TASK_STATE_WORKING, TaskState.TASK_STATE_SUBMITTED) + +// DataPart types +internal const val TYPE_FUNCTION_CALL = "function_call" +internal const val TYPE_FUNCTION_RESPONSE = "function_response" +internal const val DEFAULT_ERROR_MESSAGE = "A2A task failed" + +/** Converts a A2A [ClientEvent] to an ADK [Event]. */ +internal fun ClientEvent.toAdkEvent(invocationContext: InvocationContext): Event? { + return when (this) { + is MessageEvent -> message.toAdkEvent(invocationContext) + is TaskEvent -> task.toAdkEvent(invocationContext) + is TaskUpdateEvent -> toAdkEvent(invocationContext) + } +} + +/** Returns true if the event should be buffered for streaming. */ +internal fun ClientEvent.shouldBuffer(): Boolean { + if (this is TaskUpdateEvent) { + return this.updateEvent !is TaskStatusUpdateEvent + } + if (this is TaskEvent) { + return this.task.artifacts.orEmpty().isNotEmpty() + } + return true +} + +/** Returns true if the buffer should be reset for this event. */ +internal fun ClientEvent.shouldResetBuffer(): Boolean { + if (this is TaskEvent) { + return true + } + if (this is TaskUpdateEvent) { + val innerEvent = this.updateEvent + if (innerEvent is TaskArtifactUpdateEvent) { + return innerEvent.append == false && innerEvent.lastChunk == false + } + } + return false +} + +internal fun ClientEvent.isLastChunk(): Boolean { + if (this is TaskUpdateEvent) { + val innerEvent = this.updateEvent + if (innerEvent is TaskArtifactUpdateEvent) { + return innerEvent.lastChunk == true + } + } + return false +} + +/** Returns true if the event indicates task completion. */ +internal fun ClientEvent.isCompleted(): Boolean { + val state = + when (this) { + is TaskEvent -> this.task.status.state() + is TaskUpdateEvent -> this.task.status.state() + else -> TaskState.UNRECOGNIZED + } + return state == TaskState.TASK_STATE_COMPLETED +} + +/** Converts an artifact to an ADK event. */ +internal fun Artifact.toAdkEvent(invocationContext: InvocationContext): Event { + val adkParts = parts().toAdk() + return remoteAgentEvent(invocationContext) + .copy( + content = Content(role = Role.MODEL, parts = adkParts), + longRunningToolIds = longRunningToolIds(parts(), adkParts), + ) +} + +/** Converts an A2A message back to ADK events. */ +internal fun Message.toAdkEvent(invocationContext: InvocationContext): Event { + val adkParts = parts.toAdk() + val event = + remoteAgentEvent(invocationContext).copy(content = Content(role = Role.MODEL, parts = adkParts)) + return event.updateEventMetadata(metadata, taskId, contextId, metadataParser) +} + +/** Converts an A2A message back to ADK events with thought marking. */ +internal fun Message.toAdkEvent(invocationContext: InvocationContext, isPending: Boolean): Event { + val adkParts = parts.toAdk().map { it.copy(thought = isPending) } + return remoteAgentEvent(invocationContext) + .copy(content = Content(role = Role.MODEL, parts = adkParts)) +} + +/** Converts an A2A [Task] to an ADK [Event]. */ +internal fun Task.toAdkEvent(invocationContext: InvocationContext): Event { + val adkParts = mutableListOf() + val longRunningToolIds = mutableSetOf() + + for (artifact in artifacts.orEmpty()) { + val converted = artifact.parts().toAdk() + longRunningToolIds.addAll(longRunningToolIds(artifact.parts(), converted)) + adkParts.addAll(converted) + } + + var errorMessage: String? = null + status.message()?.let { msg -> + val msgParts = msg.parts.toAdk() + longRunningToolIds.addAll(longRunningToolIds(msg.parts, msgParts)) + if ( + status.state() == TaskState.TASK_STATE_FAILED && + msgParts.size == 1 && + msgParts[0].text != null + ) { + errorMessage = msgParts[0].text + } else { + adkParts.addAll(msgParts) + } + } + + errorMessage = + errorMessage ?: DEFAULT_ERROR_MESSAGE.takeIf { status.state() == TaskState.TASK_STATE_FAILED } + val isFinal = status.state().isFinal || status.state() == TaskState.TASK_STATE_INPUT_REQUIRED + + if (adkParts.isEmpty() && !isFinal) { + return emptyEvent(invocationContext) + } + + val event = + remoteAgentEvent(invocationContext) + .copy( + content = if (adkParts.isNotEmpty()) Content(role = Role.MODEL, parts = adkParts) else null, + longRunningToolIds = + if (status.state() == TaskState.TASK_STATE_INPUT_REQUIRED) longRunningToolIds + else emptySet(), + turnComplete = isFinal, + errorMessage = errorMessage, + ) + + return event.updateEventMetadata(metadata, id, contextId, metadataParser) +} + +/** Converts an A2A Part to an ADK Part. */ +internal fun A2APart<*>.toAdk(): Part { + return when (this) { + is TextPart -> Part(text = text) + is FilePart -> { + val fileContent = file as FileContent + when (fileContent) { + is FileWithUri -> + Part(fileData = FileData(mimeType = fileContent.mimeType(), fileUri = fileContent.uri())) + is FileWithBytes -> + Part( + inlineData = + Blob( + fileContent.mimeType(), + fileContent.name(), + Base64.getDecoder().decode(fileContent.bytes()), + ) + ) + } + } + is DataPart -> toAdk() + else -> throw IllegalArgumentException("Unsupported A2A Part type: ${this::class.simpleName}") + } +} + +/** Converts a list of A2A Parts to a list of ADK Parts. */ +internal fun List>.toAdk(): List = map { it.toAdk() } + +private fun TaskUpdateEvent.toAdkEvent(context: InvocationContext): Event? { + return when (val update = updateEvent) { + is TaskArtifactUpdateEvent -> { + val isAppend = update.append == true + val isLastChunk = update.lastChunk == true + + if (isLastChunk && update.metadata.isPartial()) { + return null + } + + val eventPart = update.artifact.toAdkEvent(context) + if (eventPart.content?.parts.isNullOrEmpty()) { + return null + } + + eventPart + .copy(partial = isAppend || !isLastChunk) + .updateEventMetadata(update.metadata, update.taskId, update.contextId, metadataParser) + } + is TaskStatusUpdateEvent -> { + val status = update.status + val taskState = task.status.state() + + val messageEvent = + status.message()?.let { msg -> + if (taskState == TaskState.TASK_STATE_FAILED) { + remoteAgentEvent(context) + .copy(errorMessage = msg.parts.filterIsInstance().firstOrNull()?.text) + } else { + msg.toAdkEvent(context, PENDING_STATES.contains(taskState)) + } + } + + val finalEvent = + if (update.isFinal) { + val baseEvent = messageEvent ?: remoteAgentEvent(context) + baseEvent.copy( + turnComplete = true, + partial = false, + errorMessage = + baseEvent.errorMessage + ?: DEFAULT_ERROR_MESSAGE.takeIf { taskState == TaskState.TASK_STATE_FAILED }, + ) + } else { + messageEvent + } + + finalEvent?.updateEventMetadata( + update.metadata, + update.taskId, + update.contextId, + metadataParser, + ) + } + } +} + +private fun longRunningToolIds( + a2aParts: List>, + adkParts: List, +): Set { + return a2aParts + .zip(adkParts) + .filter { (a2aPart, _) -> + a2aPart is DataPart && a2aPart.metadata?.get(MetadataKeys.IS_LONG_RUNNING) == true + } + .mapNotNull { (_, adkPart) -> adkPart.functionCall?.id } + .toSet() +} + +// Converts a DataPart to an ADK Part. +// Note: We use coerceToMap for arguments and response to handle cases where the data +// is received as a string or non-map type, matching Java behavior. +private fun DataPart.toAdk(): Part { + val type = metadata?.get(MetadataKeys.TYPE) as? String + // In A2A v1.0, DataPart.data is typed as Object (it may hold any JSON value). For ADK function + // call/response parts the payload is always a JSON object, so coerce it to a map here. + @Suppress("UNCHECKED_CAST") val dataMap = data as Map + return when (type) { + TYPE_FUNCTION_CALL -> { + val coercedData = dataMap.toMutableMap() + coercedData["args"] = coerceToMap(dataMap["args"]) + val fc = + adkJson.decodeFromJsonElement(FunctionCall.serializer(), anyToJsonElement(coercedData)) + Part(functionCall = fc) + } + TYPE_FUNCTION_RESPONSE -> { + val coercedData = dataMap.toMutableMap() + coercedData["response"] = coerceToMap(dataMap["response"]) + val fr = + adkJson.decodeFromJsonElement(FunctionResponse.serializer(), anyToJsonElement(coercedData)) + Part(functionResponse = fr) + } + else -> throw IllegalArgumentException("Unsupported A2A DataPart type: $type") + } +} + +private fun Map?.isPartial() = this?.get(MetadataKeys.PARTIAL) == true + +/** Converts a GenAI Content object to a list of A2A Parts. */ +internal fun Content.toA2aParts(isPartial: Boolean): List> { + return parts.map { it.toA2A(isPartial) } +} + +/** Converts an ADK Part to an A2A Part. */ +internal fun Part.toA2A(isPartial: Boolean = false): A2APart<*> { + text?.let { + return TextPart(it) + } + fileData?.let { fd -> + return FilePart(FileWithUri(fd.mimeType ?: "application/octet-stream", "", fd.fileUri ?: "")) + } + inlineData?.let { blob -> + val bytesStr = blob.data?.let { Base64.getEncoder().encodeToString(it) } ?: "" + return FilePart( + FileWithBytes(blob.mimeType ?: "application/octet-stream", blob.displayName ?: "", bytesStr) + ) + } + functionCall?.let { + return it.toA2A(isPartial) + } + functionResponse?.let { + return it.toA2A() + } + throw IllegalArgumentException("Unsupported ADK Part content") +} + +/** Converts an ADK Event to an A2A Message. */ +internal fun Event.toA2aMessage(): Message { + return Message.builder() + .messageId(id.ifEmpty { Uuid.random() }) + .role(author.takeIf { it == "user" }?.let { Message.Role.ROLE_USER } ?: Message.Role.ROLE_AGENT) + .parts(content?.parts?.map { it.toA2A() } ?: emptyList()) + .apply { + if (taskId.isNotEmpty()) taskId(taskId) + if (contextId.isNotEmpty()) contextId(contextId) + if (author.isNotEmpty()) metadata(mapOf(MetadataKeys.AUTHOR to author)) + } + .build() +} + +/** Returns the parts from the context events that should be sent to the agent. */ +internal fun InvocationContext.extractA2aParts(): List> { + val preprocessedEvents = extractPreprocessedEvents() + if (preprocessedEvents.isEmpty()) { + return emptyList() + } + + val lastResponseIndex = session.events.indexOfLast { it.author == agent.name } + + return preprocessedEvents.flatMapIndexed { index, event -> + val actualIndex = lastResponseIndex + 1 + index + val eventParts = event.content?.toA2aParts(event.partial) ?: emptyList() + logger.debug( + "Event index=$actualIndex author=${event.author} extracted parts=${eventParts.size}" + ) + eventParts + } +} + +private fun FunctionCall.toA2A(isPartial: Boolean): DataPart { + val metadata = mutableMapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL) + if (isPartial) { + metadata[MetadataKeys.PARTIAL] = true + } + @Suppress("UNCHECKED_CAST") + val dataMap = + jsonElementToAny(adkJson.encodeToJsonElement(FunctionCall.serializer(), this)) + as Map + return DataPart(dataMap, metadata) +} + +private fun FunctionResponse.toA2A(): DataPart { + @Suppress("UNCHECKED_CAST") + val dataMap = + jsonElementToAny(adkJson.encodeToJsonElement(FunctionResponse.serializer(), this)) + as Map + return DataPart(dataMap, mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_RESPONSE)) +} + +private fun coerceToMap(value: Any?): Map = + when (value) { + null -> emptyMap() + is Map<*, *> -> value.entries.associate { it.key.toString() to it.value } + is String -> + if (value.isEmpty()) { + emptyMap() + } else { + try { + @Suppress("UNCHECKED_CAST") + (jsonElementToAny(adkJson.parseToJsonElement(value)) as? Map) + ?: mapOf("value" to value) + } catch (e: Exception) { + logger.warn("Failed to parse map from string payload", e) + mapOf("value" to value) + } + } + else -> mapOf("value" to value) + } diff --git a/a2a/src/commonJvmAndroidTest/kotlin/com/google/adk/kt/a2a/agent/A2AAgentImplTest.kt b/a2a/src/commonJvmAndroidTest/kotlin/com/google/adk/kt/a2a/agent/A2AAgentImplTest.kt new file mode 100644 index 00000000..3ef739cf --- /dev/null +++ b/a2a/src/commonJvmAndroidTest/kotlin/com/google/adk/kt/a2a/agent/A2AAgentImplTest.kt @@ -0,0 +1,770 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.agent + +import com.google.adk.kt.agents.InvocationContext +import com.google.adk.kt.callbacks.AfterAgentCallback +import com.google.adk.kt.callbacks.BeforeAgentCallback +import com.google.adk.kt.callbacks.CallbackChoice +import com.google.adk.kt.events.Event +import com.google.adk.kt.events.EventActions +import com.google.adk.kt.sessions.Session +import com.google.adk.kt.sessions.SessionKey +import com.google.adk.kt.testing.DummyAgent +import com.google.adk.kt.testing.modelMessage +import com.google.adk.kt.testing.userMessage +import com.google.adk.kt.types.Content +import com.google.adk.kt.types.FunctionResponse +import com.google.adk.kt.types.Part +import com.google.common.truth.Truth.assertThat +import java.util.function.BiConsumer +import java.util.function.Consumer +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.delay +import kotlinx.coroutines.flow.first +import kotlinx.coroutines.flow.toList +import kotlinx.coroutines.launch +import kotlinx.coroutines.runBlocking +import kotlinx.coroutines.test.runTest +import org.a2aproject.sdk.client.Client +import org.a2aproject.sdk.client.ClientEvent +import org.a2aproject.sdk.client.TaskEvent +import org.a2aproject.sdk.client.TaskUpdateEvent +import org.a2aproject.sdk.spec.AgentCapabilities +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.spec.Artifact +import org.a2aproject.sdk.spec.DataPart +import org.a2aproject.sdk.spec.FilePart +import org.a2aproject.sdk.spec.FileWithUri +import org.a2aproject.sdk.spec.Message +import org.a2aproject.sdk.spec.Part as A2APart +import org.a2aproject.sdk.spec.Task +import org.a2aproject.sdk.spec.TaskArtifactUpdateEvent +import org.a2aproject.sdk.spec.TaskState +import org.a2aproject.sdk.spec.TaskStatus +import org.a2aproject.sdk.spec.TaskStatusUpdateEvent +import org.a2aproject.sdk.spec.TextPart +import org.junit.Before +import org.junit.Test +import org.junit.runner.RunWith +import org.junit.runners.JUnit4 +import org.mockito.kotlin.after +import org.mockito.kotlin.any +import org.mockito.kotlin.argumentCaptor +import org.mockito.kotlin.doAnswer +import org.mockito.kotlin.isNull +import org.mockito.kotlin.mock +import org.mockito.kotlin.never +import org.mockito.kotlin.timeout +import org.mockito.kotlin.verify +import org.mockito.kotlin.verifyNoInteractions +import org.mockito.kotlin.whenever + +@RunWith(JUnit4::class) +class A2AAgentImplTest { + + private lateinit var mockClient: Client + private lateinit var agentCard: AgentCard + private lateinit var invocationContext: InvocationContext + + @Before + fun setUp() { + mockClient = mock() + agentCard = + AgentCard.builder() + .name("remote-agent") + .description("Remote Agent") + .url("http://example.com") + .version("1.0.0") + .defaultInputModes(listOf("text")) + .defaultOutputModes(listOf("text")) + .skills(listOf()) + .supportedInterfaces(listOf()) + .capabilities(AgentCapabilities.builder().streaming(true).build()) + .build() + + whenever(mockClient.cancelTask(any())).thenReturn(null) + + val mockAgent = DummyAgent() + + val session = + Session( + key = SessionKey(appName = "demo", userId = "user", id = "session-1"), + events = + mutableListOf( + Event(invocationId = "invocation-0", author = "user", content = userMessage("hello")) + ), + ) + + invocationContext = InvocationContext(agent = mockAgent, session = session, runConfig = null) + } + + @Test + fun createAgent_streaming_false_returnsNonStreamingAgent() { + val agent = createTestAgent(streaming = false) + assertThat(agent.isStreamingEnabled).isFalse() + } + + @Test + fun createAgent_streaming_true_returnsStreamingAgent() { + val agent = createTestAgent() + assertThat(agent.isStreamingEnabled).isTrue() + } + + @Test + fun runAsync_emptySession_shortCircuits() = runTest { + val agent = createTestAgent() + invocationContext.session.events.clear() + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(1) + assertThat(events[0].turnComplete).isTrue() + assertThat(events[0].content).isNull() + + val messageCaptor = argumentCaptor() + verify(mockClient, never()) + .sendMessage( + messageCaptor.capture(), + any>>(), + any>(), + isNull(), + ) + } + + @Test + fun description_userDescriptionProvided_returnsUserDescription() { + val agent = + A2AAgentImpl( + name = "test-agent", + userDescription = "Custom User Description", + a2aClient = mockClient, + agentCard = agentCard, + ) + assertThat(agent.description).isEqualTo("Custom User Description") + } + + @Test + fun description_userDescriptionNull_agentCardProvided_returnsAgentCardDescription() { + val card = AgentCard.builder(agentCard).description("Card Description").build() + val agent = + A2AAgentImpl( + name = "test-agent", + userDescription = null, + a2aClient = mockClient, + agentCard = card, + ) + assertThat(agent.description).isEqualTo("Card Description") + } + + @Test + fun runAsync_streamingDisabled_emitsEventsWithoutFinalAggregation() = runTest { + val agent = createTestAgent(streaming = false) + + mockStreamResponse(this) { consumer -> + consumer.accept(createTaskEvent(TaskState.TASK_STATE_COMPLETED, "Done"), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + // Should contain the event from the stream, but no final aggregated event + assertThat(events).hasSize(1) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("Done") + assertThat(events[0].partial).isFalse() + assertThat(events[0].turnComplete).isTrue() + } + + @Test + fun runAsync_streamingEnabled_singleCompletedEvent_skipsAggregation() = runTest { + val agent = createTestAgent() + + mockStreamResponse(this) { consumer -> + consumer.accept(createTaskEvent(TaskState.TASK_STATE_COMPLETED, "Done"), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + // Should contain the event from the stream + assertThat(events).hasSize(1) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("Done") + assertThat(events[0].turnComplete).isTrue() + } + + @Test + fun runAsync_streamingEnabled_aggregatesPartialEvents() = runTest { + val agent = createTestAgent() + + mockStreamResponse(this) { consumer -> + consumer.accept(createPartialEvent("Hello ", true, false), agentCard) + consumer.accept(createPartialEvent("world", true, true), agentCard) + consumer.accept(createTaskEvent(TaskState.TASK_STATE_COMPLETED, "Final"), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(4) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("Hello ") + assertThat(events[0].partial).isTrue() + assertThat(events[0].turnComplete).isFalse() + assertThat(events[1].content?.parts?.firstOrNull()?.text).isEqualTo("world") + assertThat(events[1].partial).isTrue() + assertThat(events[1].turnComplete).isFalse() + assertThat(events[2].content?.parts?.firstOrNull()?.text).isEqualTo("Hello world") + assertThat(events[2].partial).isFalse() + assertThat(events[2].turnComplete).isFalse() + assertThat(events[3].content?.parts?.firstOrNull()?.text).isEqualTo("Final") + assertThat(events[3].partial).isFalse() + assertThat(events[3].turnComplete).isTrue() + } + + @Test + fun runAsync_aggregatesInterleavedFunctionCalls() = runTest { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept(createPartialEvent("Hello ", true, false), agentCard) + consumer.accept(createPartialFunctionCallEvent("get_weather", "call_1"), agentCard) + consumer.accept(createPartialEvent("World!", true, false), agentCard) + consumer.accept(createFinalEvent("Final"), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(5) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("Hello ") + assertThat(events[1].content?.parts?.firstOrNull()?.text).isEqualTo("Hello ") // Aggregated + assertThat(events[2].content?.parts?.firstOrNull()?.functionCall?.name).isEqualTo("get_weather") + assertThat(events[3].content?.parts?.firstOrNull()?.text).isEqualTo("World!") + assertThat(events[4].content?.parts?.firstOrNull()?.text).isEqualTo("Final") + } + + @Test + fun runAsync_aggregatesFiles() = runTest { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept(createPartialEvent("Here is a file: ", true, false), agentCard) + consumer.accept( + createPartialFileEvent("http://example.com/file.txt", "text/plain"), + agentCard, + ) + consumer.accept(createFinalEvent("Done"), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(4) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("Here is a file: ") + assertThat(events[1].content?.parts?.firstOrNull()?.text) + .isEqualTo("Here is a file: ") // Aggregated + assertThat(events[2].content?.parts?.firstOrNull()?.fileData?.fileUri) + .isEqualTo("http://example.com/file.txt") + assertThat(events[3].content?.parts?.firstOrNull()?.text).isEqualTo("Done") + } + + @Test + fun runAsync_taskEventSnapshotResetsBuffer_emptyFinalEvent() = runTest { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept(createPartialEvent("Hello ", true, false), agentCard) + consumer.accept(createPartialEvent("World!", true, false), agentCard) + + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .build() + consumer.accept(TaskEvent(task), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(3) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("Hello ") + assertThat(events[1].content?.parts?.firstOrNull()?.text).isEqualTo("World!") + assertThat(events[2].content).isNull() + assertThat(events[2].turnComplete).isTrue() + } + + @Test + fun runAsync_taskStatusUpdateEventFlushesBuffer() = runTest { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept(createPartialEvent("Hello ", true, false), agentCard) + consumer.accept(createPartialEvent("World!", true, false), agentCard) + + val status = TaskStatus(TaskState.TASK_STATE_COMPLETED) + val update = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + consumer.accept(TaskUpdateEvent(task, update), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(3) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("Hello ") + assertThat(events[1].content?.parts?.firstOrNull()?.text).isEqualTo("World!") + assertThat(events[2].content?.parts?.firstOrNull()?.text).isEqualTo("Hello World!") + assertThat(events[2].turnComplete).isTrue() + } + + @Test + fun runAsync_taskEventSnapshotResetsBuffer_withContent() = runTest { + val agent = createTestAgent() + + mockStreamResponse(this) { consumer -> + consumer.accept(createPartialEvent("1", true, false), agentCard) + consumer.accept(createPartialEvent("2", true, false), agentCard) + consumer.accept(createPartialEvent("3", false, false), agentCard) + consumer.accept(createPartialEvent("4", true, false), agentCard) + consumer.accept(createFinalEvent("5"), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(5) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("1") + assertThat(events[1].content?.parts?.firstOrNull()?.text).isEqualTo("2") + assertThat(events[2].content?.parts?.firstOrNull()?.text).isEqualTo("3") + assertThat(events[3].content?.parts?.firstOrNull()?.text).isEqualTo("4") + assertThat(events[4].content?.parts?.firstOrNull()?.text).isEqualTo("5") + } + + @Test + fun runAsync_taskStatusUpdateEventFlushesBuffer_withContent() = runTest { + val agent = createTestAgent() + + mockStreamResponse(this) { consumer -> + consumer.accept(createPartialEvent("1", true, false), agentCard) + consumer.accept(createPartialEvent("2", true, false), agentCard) + consumer.accept(createPartialEvent("3", false, false), agentCard) + consumer.accept(createPartialEvent("4", true, false), agentCard) + + val status = TaskStatus(TaskState.TASK_STATE_COMPLETED) + val update = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + consumer.accept(TaskUpdateEvent(task, update), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(5) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("1") + assertThat(events[1].content?.parts?.firstOrNull()?.text).isEqualTo("2") + assertThat(events[2].content?.parts?.firstOrNull()?.text).isEqualTo("3") + assertThat(events[3].content?.parts?.firstOrNull()?.text).isEqualTo("4") + assertThat(events[4].content?.parts?.firstOrNull()?.text).isEqualTo("34") + assertThat(events[4].turnComplete).isTrue() + } + + @Test + fun runAsync_handlesTasksWithStatusMessage() = runTest { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept(createTaskEvent(TaskState.TASK_STATE_COMPLETED, "hello"), agentCard) + } + val events = agent.runAsync(invocationContext).toList() + assertThat(events).hasSize(1) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("hello") + } + + @Test + fun runAsync_handlesTasksWithMultipartArtifact() = runTest { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept( + createTestEvent( + listOf(TextPart("hello"), TextPart("world")), + TaskState.TASK_STATE_COMPLETED, + append = false, + lastChunk = false, + ), + agentCard, + ) + } + val events = agent.runAsync(invocationContext).toList() + assertThat(events).hasSize(1) + val parts = events[0].content?.parts + assertThat(parts).hasSize(2) + assertThat(parts?.get(0)?.text).isEqualTo("hello") + assertThat(parts?.get(1)?.text).isEqualTo("world") + } + + @Test + fun runAsync_handlesNonFinalStatusUpdatesAsThoughts() = runTest { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept( + createStatusUpdateEvent(TaskState.TASK_STATE_SUBMITTED, "submitted..."), + agentCard, + ) + consumer.accept( + createStatusUpdateEvent(TaskState.TASK_STATE_WORKING, "working..."), + agentCard, + ) + consumer.accept(createFinalEvent("done"), agentCard) + } + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(3) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("submitted...") + assertThat(events[0].content?.parts?.firstOrNull()?.thought).isEqualTo(true) + assertThat(events[1].content?.parts?.firstOrNull()?.text).isEqualTo("working...") + assertThat(events[1].content?.parts?.firstOrNull()?.thought).isEqualTo(true) + assertThat(events[2].content?.parts?.firstOrNull()?.text).isEqualTo("done") + assertThat(events[2].content?.parts?.firstOrNull()?.thought).isNotEqualTo(true) + } + + @Test + fun runAsync_constructsRequestWithHistory() = runTest { + val agent = createTestAgent() + val historySession = + Session( + key = SessionKey(appName = "demo", userId = "user", id = "session-2"), + events = + mutableListOf( + Event(invocationId = "invocation-1", author = "user", content = userMessage("hello")), + Event(invocationId = "invocation-1", author = "model", content = modelMessage("hi")), + Event( + invocationId = "invocation-1", + author = "user", + content = userMessage("how are you?"), + ), + ), + ) + val context = invocationContext.copy(session = historySession) + mockStreamResponse(this) { consumer -> consumer.accept(createFinalEvent("fine"), agentCard) } + + agent.runAsync(context).toList() + + val messageCaptor = argumentCaptor() + verify(mockClient) + .sendMessage( + messageCaptor.capture(), + any>>(), + any>(), + isNull(), + ) + + val message = messageCaptor.firstValue + assertThat(message.role).isEqualTo(Message.Role.ROLE_USER) + assertThat(message.parts).hasSize(4) + assertThat((message.parts[0] as TextPart).text).isEqualTo("hello") + assertThat((message.parts[1] as TextPart).text).isEqualTo("For context:") + assertThat((message.parts[2] as TextPart).text).isEqualTo("[model] said: hi") + assertThat((message.parts[3] as TextPart).text).isEqualTo("how are you?") + } + + @Test + fun runAsync_constructsRequestWithFunctionResponse() = runTest { + val agent = createTestAgent() + val sessionWithFR = + Session( + key = SessionKey(appName = "demo", userId = "user", id = "session-3"), + events = + mutableListOf( + Event( + invocationId = "invocation-1", + author = "user", + content = + Content( + role = "user", + parts = + listOf( + Part( + functionResponse = + FunctionResponse( + name = "fn", + id = "call-1", + response = mapOf("status" to "ok"), + ) + ) + ), + ), + ) + ), + ) + val context = invocationContext.copy(session = sessionWithFR) + mockStreamResponse(this) { consumer -> consumer.accept(createFinalEvent("ok"), agentCard) } + + agent.runAsync(context).toList() + + val messageCaptor = argumentCaptor() + verify(mockClient) + .sendMessage( + messageCaptor.capture(), + any>>(), + any>(), + isNull(), + ) + + val message = messageCaptor.firstValue + assertThat(message.parts).hasSize(1) + val part = message.parts[0] + assertThat(part).isInstanceOf(DataPart::class.java) + val dataPart = part as DataPart + val data = dataPart.data as Map<*, *> + assertThat(data["name"]).isEqualTo("fn") + assertThat(data["id"]).isEqualTo("call-1") + assertThat(dataPart.metadata?.get("adk_type")).isEqualTo("function_response") + } + + @Test + fun runAsync_handlesClientError() = runTest { + val agent = createTestAgent() + + val error = RuntimeException("Connection failed") + mockStreamError(this, error) + + val result = runCatching { agent.runAsync(invocationContext).toList() } + assertThat(result.isFailure).isTrue() + assertThat(result.exceptionOrNull()?.message).contains("Connection failed") + } + + @Test + fun runAsync_invokesBeforeAndAfterCallbacks() = runTest { + var beforeCalled = false + var afterCalled = false + val agent = + createTestAgent( + beforeCallbacks = + listOf( + BeforeAgentCallback { _ -> + beforeCalled = true + CallbackChoice.Continue(EventActions()) + } + ), + afterCallbacks = + listOf( + AfterAgentCallback { _ -> + afterCalled = true + CallbackChoice.Continue(Unit) + } + ), + ) + mockStreamResponse(this) { consumer -> consumer.accept(createFinalEvent("done"), agentCard) } + + agent.runAsync(invocationContext).toList() + + assertThat(beforeCalled).isTrue() + assertThat(afterCalled).isTrue() + } + + @Test + fun runAsync_beforeCallbackCanShortCircuit() = runTest { + val shortCircuitContent = modelMessage("short circuit") + val agent = + createTestAgent( + beforeCallbacks = + listOf(BeforeAgentCallback { _ -> CallbackChoice.Break(shortCircuitContent) }) + ) + + val events = agent.runAsync(invocationContext).toList() + + assertThat(events).hasSize(1) + assertThat(events[0].content?.parts?.firstOrNull()?.text).isEqualTo("short circuit") + + verifyNoInteractions(mockClient) + } + + // Uses runBlocking (real time) because task cancellation runs on Dispatchers.IO. + @Test + fun runAsync_nonTerminalTaskAbandoned_cancelsRemoteTask() = runBlocking { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept( + createStatusUpdateEvent(TaskState.TASK_STATE_WORKING, "working..."), + agentCard, + ) + } + + // Collecting a single event abandons the still-open task, triggering cleanup. + agent.runAsync(invocationContext).first() + + verify(mockClient, timeout(5000)).cancelTask(any()) + Unit + } + + @Test + fun runAsync_terminalInputRequiredTaskAbandoned_doesNotCancelRemoteTask() = runBlocking { + val agent = createTestAgent() + mockStreamResponse(this) { consumer -> + consumer.accept( + createStatusUpdateEvent(TaskState.TASK_STATE_INPUT_REQUIRED, "need input"), + agentCard, + ) + } + + // An input-required task is terminal, so abandoning the flow must not cancel it. + agent.runAsync(invocationContext).first() + + verify(mockClient, after(500).never()).cancelTask(any()) + Unit + } + + private fun createTestAgent( + streaming: Boolean = true, + beforeCallbacks: List = emptyList(), + afterCallbacks: List = emptyList(), + ): BaseRemoteA2AAgent { + return A2AAgentImpl( + name = "remote-agent", + a2aClient = mockClient, + agentCard = agentCard, + streaming = streaming, + beforeAgentCallbacks = beforeCallbacks, + afterAgentCallbacks = afterCallbacks, + ) + } + + private fun assertEventText(event: Event?, expectedText: String) { + assertThat(event?.content?.parts?.firstOrNull()?.text).isEqualTo(expectedText) + } + + private fun mockStreamError(scope: CoroutineScope, error: Throwable) { + doAnswer { invocation -> + val errorConsumer = invocation.getArgument>(2) + scope.launch { + delay(10) + errorConsumer.accept(error) + } + null + } + .whenever(mockClient) + .sendMessage( + any(), + any>>(), + any>(), + isNull(), + ) + } + + private fun mockStreamResponse( + scope: CoroutineScope, + responseProducer: (BiConsumer) -> Unit, + ) { + doAnswer { invocation -> + val consumers = invocation.getArgument>>(1) + val consumer = consumers[0] + scope.launch { + delay(10) + responseProducer(consumer) + } + null + } + .whenever(mockClient) + .sendMessage( + any(), + any>>(), + any>(), + isNull(), + ) + } + + private fun createTaskEvent(state: TaskState, text: String): TaskEvent { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status( + TaskStatus( + state, + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(TextPart(text))).build(), + null, + ) + ) + .build() + return TaskEvent(task) + } + + private fun createStatusUpdateEvent(state: TaskState, text: String): ClientEvent { + val task = Task.builder().id("task-1").contextId("context-1").status(TaskStatus(state)).build() + val update = + TaskStatusUpdateEvent.builder() + .taskId("task-1") + .contextId("context-1") + .status( + TaskStatus( + state, + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(TextPart(text))).build(), + null, + ) + ) + .build() + return TaskUpdateEvent(task, update) + } + + private fun createPartialEvent(text: String, append: Boolean, lastChunk: Boolean): ClientEvent { + return createTestEvent(TextPart(text), TaskState.TASK_STATE_WORKING, append, lastChunk) + } + + private fun createPartialFunctionCallEvent(name: String, id: String): ClientEvent { + val data = mapOf("name" to name, "id" to id, "args" to mapOf()) + val metadata = mapOf("adk_type" to "function_call") + return createTestEvent(DataPart(data, metadata), TaskState.TASK_STATE_WORKING, true, false) + } + + private fun createPartialFileEvent(uri: String, mimeType: String): ClientEvent { + return createTestEvent( + FilePart(FileWithUri(mimeType, "file", uri)), + TaskState.TASK_STATE_WORKING, + true, + false, + ) + } + + private fun createFinalEvent(text: String): ClientEvent { + return createTestEvent(TextPart(text), TaskState.TASK_STATE_COMPLETED, false, false) + } + + private fun createTestEvent( + part: A2APart<*>, + state: TaskState, + append: Boolean, + lastChunk: Boolean, + ): ClientEvent = createTestEvent(listOf(part), state, append, lastChunk) + + private fun createTestEvent( + parts: List>, + state: TaskState, + append: Boolean, + lastChunk: Boolean, + ): ClientEvent { + val artifact = Artifact.builder().artifactId("artifact-1").parts(parts).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(state)) + .artifacts(listOf(artifact)) + .build() + + if (state == TaskState.TASK_STATE_COMPLETED && !append && !lastChunk) { + return TaskEvent(task) + } + + val updateEvent = + TaskArtifactUpdateEvent.builder() + .lastChunk(lastChunk) + .append(append) + .contextId("context-1") + .artifact(artifact) + .taskId("task-id-1") + .build() + return TaskUpdateEvent(task, updateEvent) + } +} diff --git a/a2a/src/commonJvmAndroidTest/kotlin/com/google/adk/kt/a2a/converters/A2aConvertersTest.kt b/a2a/src/commonJvmAndroidTest/kotlin/com/google/adk/kt/a2a/converters/A2aConvertersTest.kt new file mode 100644 index 00000000..dedc26b9 --- /dev/null +++ b/a2a/src/commonJvmAndroidTest/kotlin/com/google/adk/kt/a2a/converters/A2aConvertersTest.kt @@ -0,0 +1,1085 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package com.google.adk.kt.a2a.converters + +import com.google.adk.kt.agents.InvocationContext +import com.google.adk.kt.events.Event +import com.google.adk.kt.sessions.Session +import com.google.adk.kt.sessions.SessionKey +import com.google.adk.kt.testing.DummyAgent +import com.google.adk.kt.testing.modelMessage +import com.google.adk.kt.testing.testSession +import com.google.adk.kt.testing.userMessage +import com.google.adk.kt.types.Blob +import com.google.adk.kt.types.Content +import com.google.adk.kt.types.FileData +import com.google.adk.kt.types.FunctionCall +import com.google.adk.kt.types.FunctionResponse +import com.google.adk.kt.types.GroundingMetadata +import com.google.adk.kt.types.Part +import com.google.adk.kt.types.Role +import com.google.adk.kt.types.UsageMetadata +import com.google.common.truth.Truth.assertThat +import java.util.Base64 +import kotlin.test.assertFailsWith +import org.a2aproject.sdk.client.MessageEvent +import org.a2aproject.sdk.client.TaskEvent +import org.a2aproject.sdk.client.TaskUpdateEvent +import org.a2aproject.sdk.spec.Artifact +import org.a2aproject.sdk.spec.DataPart +import org.a2aproject.sdk.spec.FilePart +import org.a2aproject.sdk.spec.FileWithBytes +import org.a2aproject.sdk.spec.FileWithUri +import org.a2aproject.sdk.spec.Message +import org.a2aproject.sdk.spec.Part as A2APart +import org.a2aproject.sdk.spec.Task +import org.a2aproject.sdk.spec.TaskArtifactUpdateEvent +import org.a2aproject.sdk.spec.TaskState +import org.a2aproject.sdk.spec.TaskStatus +import org.a2aproject.sdk.spec.TaskStatusUpdateEvent +import org.a2aproject.sdk.spec.TextPart +import org.junit.Test +import org.junit.runner.RunWith +import org.junit.runners.JUnit4 + +@RunWith(JUnit4::class) +class A2aConvertersTest { + + private val testAgent = DummyAgent(name = "test_agent") + + private val invocationContext = + InvocationContext( + invocationId = "invocation-1", + agent = testAgent, + branch = "main", + session = testSession(), + runConfig = null, + ) + + @Test + fun toAdk_withTextPart_returnsAdkTextPart() { + val textPart = TextPart("Hello") + val result = textPart.toAdk() + assertThat(result.text).isEqualTo("Hello") + } + + @Test + fun toA2A_withTextPart_returnsTextPart() { + val part = Part(text = "Hello") + val result = part.toA2A() + assertThat((result as TextPart).text).isEqualTo("Hello") + } + + @Test + fun toAdk_withFilePartUri_returnsAdkFilePart() { + val filePart = FilePart(FileWithUri("text/plain", "file.txt", "http://file.txt")) + val result = filePart.toAdk() + val fileData = result.fileData + assertThat(fileData).isNotNull() + assertThat(fileData!!.mimeType).isEqualTo("text/plain") + assertThat(fileData.fileUri).isEqualTo("http://file.txt") + } + + @Test + fun toA2A_withFileDataPart_returnsFilePartWithUri() { + val part = Part(fileData = FileData(mimeType = "text/plain", fileUri = "http://file.txt")) + val result = part.toA2A() + assertThat((result as FilePart).file.mimeType()).isEqualTo("text/plain") + assertThat((result.file as FileWithUri).uri()).isEqualTo("http://file.txt") + } + + @Test + fun toAdk_withFilePartBytes_returnsAdkBlobPart() { + val bytes = "file content".toByteArray() + val encoded = Base64.getEncoder().encodeToString(bytes) + val filePart = FilePart(FileWithBytes("text/plain", "file.txt", encoded)) + val result = filePart.toAdk() + val blob = result.inlineData + assertThat(blob).isNotNull() + assertThat(blob!!.mimeType).isEqualTo("text/plain") + assertThat(blob.displayName).isEqualTo("file.txt") + assertThat(blob.data).isNotNull() + assertThat(String(blob.data!!)).isEqualTo("file content") + } + + @Test + fun toA2A_withInlineDataPart_returnsFilePartWithBytes() { + val bytes = "content".toByteArray() + val part = + Part(inlineData = Blob(mimeType = "text/plain", displayName = "file.txt", data = bytes)) + val result = part.toA2A() + assertThat((result as FilePart).file.mimeType()).isEqualTo("text/plain") + assertThat(result.file.name()).isEqualTo("file.txt") + assertThat((result.file as FileWithBytes).bytes()) + .isEqualTo(Base64.getEncoder().encodeToString(bytes)) + } + + @Test + fun toAdk_withDataPartFunctionCall_returnsAdkFunctionCallPart() { + val data = mapOf("name" to "func", "id" to "1", "args" to mapOf()) + val metadata = mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL) + val dataPart = DataPart(data, metadata) + val result = dataPart.toAdk() + val functionCall = result.functionCall + assertThat(functionCall).isNotNull() + assertThat(functionCall!!.name).isEqualTo("func") + assertThat(functionCall.id).isEqualTo("1") + assertThat(functionCall.args).isEqualTo(mapOf()) + } + + @Test + fun toA2A_withFunctionCallPart_returnsDataPart() { + val part = Part(functionCall = FunctionCall(name = "func", id = "1", args = mapOf())) + val result = part.toA2A() + val dataPart = result as DataPart + val data = dataPart.data as Map<*, *> + assertThat(data["name"]).isEqualTo("func") + assertThat(data["id"]).isEqualTo("1") + assertThat(dataPart.metadata?.get(MetadataKeys.TYPE)).isEqualTo(TYPE_FUNCTION_CALL) + } + + @Test + fun toAdk_withDataPartFunctionResponse_returnsAdkFunctionResponsePart() { + val data = mapOf("name" to "func", "id" to "1", "response" to mapOf()) + val metadata = mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_RESPONSE) + val dataPart = DataPart(data, metadata) + val result = dataPart.toAdk() + val functionResponse = result.functionResponse + assertThat(functionResponse).isNotNull() + assertThat(functionResponse!!.name).isEqualTo("func") + assertThat(functionResponse.id).isEqualTo("1") + assertThat(functionResponse.response).isEqualTo(mapOf()) + } + + @Test + fun toA2A_withFunctionResponsePart_returnsDataPart() { + val part = + Part(functionResponse = FunctionResponse(name = "func", id = "1", response = mapOf())) + val result = part.toA2A() + val dataPart = result as DataPart + val data = dataPart.data as Map<*, *> + assertThat(data["name"]).isEqualTo("func") + assertThat(data["id"]).isEqualTo("1") + assertThat(dataPart.metadata?.get(MetadataKeys.TYPE)).isEqualTo(TYPE_FUNCTION_RESPONSE) + } + + @Test + fun toAdk_convertsAllSupportedParts() { + val a2aParts = + listOf(TextPart("text"), FilePart(FileWithUri("text/plain", "file.txt", "http://file.txt"))) + val result = a2aParts.toAdk() + assertThat(result.size).isEqualTo(2) + assertThat(result[0].text).isEqualTo("text") + assertThat(result[1].fileData).isNotNull() + } + + @Test + fun toAdk_withDataPartWithEmptyStringCoercedToEmptyMap() { + val data = mapOf("name" to "func", "id" to "1", "args" to "") + val metadata = mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL) + val dataPart = DataPart(data, metadata) + val result = dataPart.toAdk() + val functionCall = result.functionCall + assertThat(functionCall).isNotNull() + assertThat(functionCall!!.args).isEqualTo(mapOf()) + } + + @Test + fun toAdk_withDataPartWithNonMapCoercedToMap() { + val data = mapOf("name" to "func", "id" to "1", "args" to 123) + val metadata = mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL) + val dataPart = DataPart(data, metadata) + val result = dataPart.toAdk() + val functionCall = result.functionCall + assertThat(functionCall).isNotNull() + // gson's LONG_OR_DOUBLE number strategy decodes untyped integers as Long. + assertThat(functionCall!!.args).isEqualTo(mapOf("value" to 123L)) + } + + @Test + fun toAdk_withDataPartWithJsonStringCoercedToMap() { + val data = mapOf("name" to "func", "id" to "1", "args" to "{\"key\": \"value\"}") + val metadata = mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL) + val dataPart = DataPart(data, metadata) + val result = dataPart.toAdk() + val functionCall = result.functionCall + assertThat(functionCall).isNotNull() + assertThat(functionCall!!.args).isEqualTo(mapOf("key" to "value")) + } + + @Test + fun toAdk_withDataPartWithInvalidJsonStringCoercedToMap() { + val data = mapOf("name" to "func", "id" to "1", "args" to "{invalid}") + val metadata = mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL) + val dataPart = DataPart(data, metadata) + val result = dataPart.toAdk() + val functionCall = result.functionCall + assertThat(functionCall).isNotNull() + assertThat(functionCall!!.args).isEqualTo(mapOf("value" to "{invalid}")) + } + + @Test + fun toAdk_withFilePartBytes_handlesInvalidBase64() { + val filePart = FilePart(FileWithBytes("text/plain", "file.txt", "invalid-base64!")) + assertFailsWith { filePart.toAdk() } + } + + @Test + fun clientEventToEvent_withMessageEvent_returnsEvent() { + val a2aMessage = + Message.builder() + .messageId("msg-1") + .role(Message.Role.ROLE_USER) + .parts(listOf(TextPart("Hello"))) + .build() + val messageEvent = MessageEvent(a2aMessage) + + val result = messageEvent.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result!!.author).isEqualTo("test_agent") + assertThat(result.content?.parts?.get(0)?.text).isEqualTo("Hello") + } + + @Test + fun messageToEvent_convertsMessage() { + val a2aMessage = + Message.builder() + .messageId("msg-1") + .role(Message.Role.ROLE_USER) + .parts(listOf(TextPart("test-message"))) + .build() + + val result = a2aMessage.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.author).isEqualTo("test_agent") + assertThat(result.content?.role).isEqualTo("model") + assertThat(result.content?.parts?.get(0)?.text).isEqualTo("test-message") + } + + @Test + fun taskToEvent_withArtifacts_returnsEventFromLastArtifact() { + val a2aPart = TextPart("Artifact content") + val artifact = Artifact.builder().artifactId("artifact-1").parts(listOf(a2aPart)).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .artifacts(listOf(artifact)) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.content?.parts?.get(0)?.text).isEqualTo("Artifact content") + } + + @Test + fun taskToEvent_withStatusMessage_returnsEvent() { + val statusMessage = + Message.builder() + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart("Status message"))) + .build() + val status = TaskStatus(TaskState.TASK_STATE_WORKING, statusMessage, null) + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .artifacts(emptyList()) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.content?.parts?.get(0)?.text).isEqualTo("Status message") + } + + @Test + fun taskToEvent_withFailedState_setsErrorCode() { + val statusMessage = + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(TextPart("Task failed"))).build() + val status = TaskStatus(TaskState.TASK_STATE_FAILED, statusMessage, null) + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .artifacts(emptyList()) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.errorMessage).isEqualTo("Task failed") + } + + @Test + fun taskToEvent_withFailedStateAndMultipleStatusParts_keepsPartsAsContent() { + // Only a single-part status message becomes the error text; multi-part stays as content. + val statusMessage = + Message.builder() + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart("First"), TextPart("Second"))) + .build() + val status = TaskStatus(TaskState.TASK_STATE_FAILED, statusMessage, null) + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .artifacts(emptyList()) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result.content?.parts?.map { it.text }).containsExactly("First", "Second").inOrder() + assertThat(result.errorMessage).isEqualTo("A2A task failed") + } + + @Test + fun taskToEvent_withInputRequired_parsesLongRunningToolIds() { + val data = mapOf("name" to "myTool", "id" to "call_123", "args" to mapOf()) + val metadata = + mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL, MetadataKeys.IS_LONG_RUNNING to true) + val dataPart = DataPart(data, metadata) + + val statusData = + mapOf("name" to "messageTools", "id" to "msg_123", "args" to mapOf()) + val statusMetadata = + mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL, MetadataKeys.IS_LONG_RUNNING to true) + val statusDataPart = DataPart(statusData, statusMetadata) + + val statusMessage = + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(statusDataPart)).build() + val status = TaskStatus(TaskState.TASK_STATE_INPUT_REQUIRED, statusMessage, null) + + val artifact = Artifact.builder().artifactId("artifact-1").parts(listOf(dataPart)).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .artifacts(listOf(artifact)) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.longRunningToolIds).isEqualTo(setOf("call_123", "msg_123")) + } + + @Test + fun taskToEvent_withGroundingMetadata_returnsEvent() { + val statusMessage = + Message.builder() + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart("Status message"))) + .build() + val status = TaskStatus(TaskState.TASK_STATE_WORKING, statusMessage, null) + + val groundingMetadataJson = "{\"imageSearchQueries\":[\"test-query\"]}" + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .metadata(mapOf(MetadataKeys.GROUNDING to groundingMetadataJson)) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.groundingMetadata).isNotNull() + } + + @Test + fun taskToEvent_withCustomMetadata_returnsEvent() { + val statusMessage = + Message.builder() + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart("Status message"))) + .build() + val status = TaskStatus(TaskState.TASK_STATE_WORKING, statusMessage, null) + + val customMetadataMap = mapOf("test-key" to "test-value") + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .metadata(mapOf(MetadataKeys.CUSTOM to customMetadataMap)) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.customMetadata?.get(ADK_METADATA_TASK_ID)).isEqualTo("task-1") + assertThat(result.customMetadata?.get(ADK_METADATA_CONTEXT_ID)).isEqualTo("context-1") + assertThat(result.customMetadata?.get("test-key")).isEqualTo("test-value") + } + + @Test + fun taskToEvent_withErrorCode_returnsEvent() { + val statusMessage = + Message.builder() + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart("Status message"))) + .build() + val status = TaskStatus(TaskState.TASK_STATE_WORKING, statusMessage, null) + + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .metadata(mapOf(MetadataKeys.ERROR_CODE to "STOP")) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.errorCode).isEqualTo("STOP") + } + + @Test + fun clientEventToEvent_withTaskUpdateEventAndThought_returnsThoughtEvent() { + val statusMessage = + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(TextPart("thought-1"))).build() + val status = TaskStatus(TaskState.TASK_STATE_WORKING, statusMessage, null) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + + val updateEvent = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result!!.content?.parts?.get(0)?.text).isEqualTo("thought-1") + assertThat(result.content?.parts?.get(0)?.thought).isEqualTo(true) + } + + @Test + fun clientEventToEvent_withTaskArtifactUpdateEvent_withLastChunkTrue_returnsTaskEvent() { + val a2aPart = TextPart("Artifact content") + val artifact = Artifact.builder().artifactId("artifact-1").parts(listOf(a2aPart)).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .artifacts(listOf(artifact)) + .build() + + val updateEvent = + TaskArtifactUpdateEvent.builder() + .lastChunk(true) + .contextId("context-1") + .artifact(artifact) + .taskId("task-id-1") + .build() + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result!!.content?.parts?.get(0)?.text).isEqualTo("Artifact content") + } + + @Test + fun clientEventToEvent_withTaskArtifactUpdateEvent_withLastChunkFalse_returnsHandlingPartialEvent() { + val a2aPart = TextPart("Artifact content") + val artifact = Artifact.builder().artifactId("artifact-1").parts(listOf(a2aPart)).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .artifacts(listOf(artifact)) + .build() + + val updateEvent = + TaskArtifactUpdateEvent.builder() + .lastChunk(false) + .append(false) + .contextId("context-1") + .artifact(artifact) + .taskId("task-id-1") + .build() + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result!!.partial).isEqualTo(true) + } + + @Test + fun clientEventToEvent_withTaskArtifactUpdateEvent_lastChunkTrueNotAppend_returnsNonPartialEvent() { + val a2aPart = TextPart("Artifact content") + val artifact = Artifact.builder().artifactId("artifact-1").parts(listOf(a2aPart)).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .artifacts(listOf(artifact)) + .build() + + val updateEvent = + TaskArtifactUpdateEvent.builder() + .lastChunk(true) + .append(false) + .contextId("context-1") + .artifact(artifact) + .taskId("task-id-1") + .build() + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result!!.partial).isEqualTo(false) + } + + @Test + fun clientEventToEvent_withFinalTaskStatusUpdateEvent_withMessage_returnsEvent() { + val statusMessage = + Message.builder() + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart("Final status message"))) + .build() + val status = TaskStatus(TaskState.TASK_STATE_COMPLETED, statusMessage, null) + val updateEvent = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result!!.content?.parts?.get(0)?.text).isEqualTo("Final status message") + assertThat(result.partial).isEqualTo(false) + assertThat(result.turnComplete).isEqualTo(true) + } + + @Test + fun clientEventToEvent_withFailedTaskStatusUpdateEvent_returnsErrorEvent() { + val statusMessage = + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(TextPart("Task failed"))).build() + val status = TaskStatus(TaskState.TASK_STATE_FAILED, statusMessage, null) + val updateEvent = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result!!.errorMessage).isEqualTo("Task failed") + assertThat(result.turnComplete).isEqualTo(true) + } + + @Test + fun toA2aParts_validContent_returnsParts() { + val textPart = Part(text = "hello") + val content = Content(parts = listOf(textPart)) + val list = content.toA2aParts(false) + assertThat(list.size).isEqualTo(1) + assertThat((list[0] as TextPart).text).isEqualTo("hello") + } + + @Test + fun toA2aMessage_withUserAuthor_returnsUserRole() { + val event = Event(author = "user", content = userMessage("hello")) + val result = event.toA2aMessage() + assertThat(result.role).isEqualTo(Message.Role.ROLE_USER) + } + + @Test + fun toA2aMessage_withAgentAuthor_returnsAgentRole() { + val event = Event(author = "agent", content = modelMessage("hello")) + val result = event.toA2aMessage() + assertThat(result.role).isEqualTo(Message.Role.ROLE_AGENT) + } + + @Test + fun toA2aMessage_addsAuthorToMetadata() { + val event = Event(author = "test_author", content = userMessage("hello")) + val result = event.toA2aMessage() + assertThat(result.metadata?.get(MetadataKeys.AUTHOR)).isEqualTo("test_author") + } + + @Test + fun extractA2aParts_sessionHasEvents_returnsFormattedParts() { + val userEvent = Event(author = Role.USER, content = userMessage("hello")) + val agentEvent = Event(author = "test_agent", content = modelMessage("hi")) + val otherAgentEvent = Event(author = "other_agent", content = modelMessage("hey")) + + val session = + Session( + key = SessionKey(appName = "demo", userId = "user", id = "session-1"), + events = mutableListOf(userEvent, agentEvent, otherAgentEvent), + ) + + val mockAgent = DummyAgent(name = "test_agent") + val ctx = InvocationContext(agent = mockAgent, session = session, runConfig = null) + + val parts = ctx.extractA2aParts() + assertThat(parts.size).isEqualTo(2) + assertThat((parts[0] as TextPart).text).isEqualTo("For context:") + assertThat((parts[1] as TextPart).text).isEqualTo("[other_agent] said: hey") + } + + @Test + fun shouldBuffer_withTaskUpdateEventNonStatus_returnsTrue() { + val artifact = + Artifact.builder().artifactId("artifact-1").parts(listOf(TextPart("content"))).build() + val updateEvent = + TaskArtifactUpdateEvent.builder() + .artifact(artifact) + .taskId("task-1") + .contextId("context-1") + .build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_WORKING)) + .build() + val event = TaskUpdateEvent(task, updateEvent) + + assertThat(event.shouldBuffer()).isTrue() + } + + @Test + fun shouldBuffer_withTaskUpdateEventStatus_returnsFalse() { + val status = TaskStatus(TaskState.TASK_STATE_WORKING) + val updateEvent = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val event = TaskUpdateEvent(task, updateEvent) + + assertThat(event.shouldBuffer()).isFalse() + } + + @Test + fun shouldBuffer_withTaskEventWithArtifacts_returnsTrue() { + val artifact = + Artifact.builder().artifactId("artifact-1").parts(listOf(TextPart("content"))).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .artifacts(listOf(artifact)) + .status(TaskStatus(TaskState.TASK_STATE_WORKING)) + .build() + val event = TaskEvent(task) + + assertThat(event.shouldBuffer()).isTrue() + } + + @Test + fun shouldBuffer_withTaskEventWithoutArtifacts_returnsFalse() { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .artifacts(emptyList()) + .status(TaskStatus(TaskState.TASK_STATE_WORKING)) + .build() + val event = TaskEvent(task) + + assertThat(event.shouldBuffer()).isFalse() + } + + @Test + fun shouldBuffer_withOtherEvent_returnsTrue() { + val message = + Message.builder() + .messageId("msg-1") + .role(Message.Role.ROLE_USER) + .parts(listOf(TextPart("hello"))) + .build() + val event = MessageEvent(message) + + assertThat(event.shouldBuffer()).isTrue() + } + + @Test + fun shouldResetBuffer_withTaskUpdateEventArtifactNotAppendNotLast_returnsTrue() { + val artifact = + Artifact.builder().artifactId("artifact-1").parts(listOf(TextPart("content"))).build() + val updateEvent = + TaskArtifactUpdateEvent.builder() + .artifact(artifact) + .append(false) + .lastChunk(false) + .taskId("task-1") + .contextId("context-1") + .build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_WORKING)) + .build() + val event = TaskUpdateEvent(task, updateEvent) + + assertThat(event.shouldResetBuffer()).isTrue() + } + + @Test + fun shouldResetBuffer_withTaskUpdateEventArtifactAppend_returnsFalse() { + val artifact = + Artifact.builder().artifactId("artifact-1").parts(listOf(TextPart("content"))).build() + val updateEvent = + TaskArtifactUpdateEvent.builder() + .artifact(artifact) + .append(true) + .lastChunk(false) + .taskId("task-1") + .contextId("context-1") + .build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_WORKING)) + .build() + val event = TaskUpdateEvent(task, updateEvent) + + assertThat(event.shouldResetBuffer()).isFalse() + } + + @Test + fun shouldResetBuffer_withTaskUpdateEventArtifactLast_returnsFalse() { + val artifact = + Artifact.builder().artifactId("artifact-1").parts(listOf(TextPart("content"))).build() + val updateEvent = + TaskArtifactUpdateEvent.builder() + .artifact(artifact) + .append(false) + .lastChunk(true) + .taskId("task-1") + .contextId("context-1") + .build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_WORKING)) + .build() + val event = TaskUpdateEvent(task, updateEvent) + + assertThat(event.shouldResetBuffer()).isFalse() + } + + @Test + fun shouldResetBuffer_withTaskUpdateEventNonArtifact_returnsFalse() { + val status = TaskStatus(TaskState.TASK_STATE_WORKING) + val updateEvent = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val event = TaskUpdateEvent(task, updateEvent) + + assertThat(event.shouldResetBuffer()).isFalse() + } + + @Test + fun shouldResetBuffer_withNonTaskUpdateEvent_returnsFalse() { + val message = + Message.builder() + .messageId("msg-1") + .role(Message.Role.ROLE_USER) + .parts(listOf(TextPart("hello"))) + .build() + val event = MessageEvent(message) + + assertThat(event.shouldResetBuffer()).isFalse() + } + + @Test + fun shouldResetBuffer_withTaskEvent_returnsTrue() { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_WORKING)) + .build() + val event = TaskEvent(task) + + assertThat(event.shouldResetBuffer()).isTrue() + } + + @Test + fun isCompleted_withTaskEventCompleted_returnsTrue() { + val status = TaskStatus(TaskState.TASK_STATE_COMPLETED) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val event = TaskEvent(task) + + assertThat(event.isCompleted()).isTrue() + } + + @Test + fun isCompleted_withTaskEventNotCompleted_returnsFalse() { + val status = TaskStatus(TaskState.TASK_STATE_WORKING) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val event = TaskEvent(task) + + assertThat(event.isCompleted()).isFalse() + } + + @Test + fun isCompleted_withTaskUpdateEventCompleted_returnsTrue() { + val status = TaskStatus(TaskState.TASK_STATE_COMPLETED) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val updateEvent = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val event = TaskUpdateEvent(task, updateEvent) + + assertThat(event.isCompleted()).isTrue() + } + + @Test + fun isCompleted_withTaskUpdateEventNotCompleted_returnsFalse() { + val status = TaskStatus(TaskState.TASK_STATE_WORKING) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val updateEvent = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val event = TaskUpdateEvent(task, updateEvent) + + assertThat(event.isCompleted()).isFalse() + } + + @Test + fun isCompleted_withOtherEvent_returnsFalse() { + val message = + Message.builder() + .messageId("msg-1") + .role(Message.Role.ROLE_USER) + .parts(listOf(TextPart("hello"))) + .build() + val event = MessageEvent(message) + + assertThat(event.isCompleted()).isFalse() + } + + @Test + fun clientEventToEvent_withTaskArtifactUpdateEvent_withLastChunkAndPartial_returnsNull() { + val a2aPart = TextPart("Artifact content") + val artifact = Artifact.builder().artifactId("artifact-1").parts(listOf(a2aPart)).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .artifacts(listOf(artifact)) + .build() + + val updateEvent = + TaskArtifactUpdateEvent.builder() + .lastChunk(true) + .metadata(mapOf(MetadataKeys.PARTIAL to true)) + .contextId("context-1") + .artifact(artifact) + .taskId("task-id-1") + .build() + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNull() + } + + @Test + fun clientEventToEvent_withTaskArtifactUpdateEvent_withEmptyParts_returnsNull() { + val partsList = mutableListOf>(TextPart("dummy")) + val artifact = Artifact.builder().artifactId("artifact-1").parts(partsList).build() + partsList.clear() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .artifacts(listOf(artifact)) + .build() + + val updateEvent = + TaskArtifactUpdateEvent.builder() + .lastChunk(true) + .contextId("context-1") + .artifact(artifact) + .taskId("task-id-1") + .build() + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNull() + } + + @Test + fun taskToEvent_withNonInputRequiredState_assertLongRunningToolIdsIsEmpty() { + val data = mapOf("name" to "myTool", "id" to "call_123", "args" to mapOf()) + val metadata = + mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL, MetadataKeys.IS_LONG_RUNNING to true) + val dataPart = DataPart(data, metadata) + + val statusData = + mapOf("name" to "messageTools", "id" to "msg_123", "args" to mapOf()) + val statusMetadata = + mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_CALL, MetadataKeys.IS_LONG_RUNNING to true) + val statusDataPart = DataPart(statusData, statusMetadata) + + val statusMessage = + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(statusDataPart)).build() + val status = TaskStatus(TaskState.TASK_STATE_WORKING, statusMessage, null) + + val artifact = Artifact.builder().artifactId("artifact-1").parts(listOf(dataPart)).build() + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .artifacts(listOf(artifact)) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.longRunningToolIds).isEmpty() + } + + @Test + fun clientEventToEvent_withFailedTaskStatusUpdateEvent_noTextPart_returnsFallbackErrorMessage() { + val nonTextPart = FilePart(FileWithUri("text/plain", "file.txt", "http://file.txt")) + val statusMessage = + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(nonTextPart)).build() + val status = TaskStatus(TaskState.TASK_STATE_FAILED, statusMessage, null) + val updateEvent = TaskStatusUpdateEvent("task-1", status, "context-1", null) + val task = Task.builder().id("task-1").contextId("context-1").status(status).build() + val event = TaskUpdateEvent(task, updateEvent) + + val result = event.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result!!.errorMessage).isEqualTo(DEFAULT_ERROR_MESSAGE) + assertThat(result.turnComplete).isEqualTo(true) + } + + @Test + fun taskToEvent_withFailedState_noTextPart_returnsFallbackErrorMessage() { + val nonTextPart = FilePart(FileWithUri("text/plain", "file.txt", "http://file.txt")) + val statusMessage = + Message.builder().role(Message.Role.ROLE_AGENT).parts(listOf(nonTextPart)).build() + val status = TaskStatus(TaskState.TASK_STATE_FAILED, statusMessage, null) + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .artifacts(emptyList()) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result).isNotNull() + assertThat(result.errorMessage).isEqualTo(DEFAULT_ERROR_MESSAGE) + } + + @Test + fun toAdk_withUnsupportedPartType_throwsException() { + val unsupportedPart = object : A2APart {} + val exception = assertFailsWith { unsupportedPart.toAdk() } + assertThat(exception.message).contains("Unsupported A2A Part type") + } + + @Test + fun toA2A_withUnsupportedAdkPart_throwsException() { + val emptyPart = Part() + val exception = assertFailsWith { emptyPart.toA2A() } + assertThat(exception.message).contains("Unsupported ADK Part content") + } + + @Test + fun toAdk_withDataPartFunctionResponseWithNonMapCoercedToMap() { + val data = mapOf("name" to "func", "id" to "1", "response" to 456) + val metadata = mapOf(MetadataKeys.TYPE to TYPE_FUNCTION_RESPONSE) + val dataPart = DataPart(data, metadata) + val result = dataPart.toAdk() + val functionResponse = result.functionResponse + assertThat(functionResponse).isNotNull() + // gson's LONG_OR_DOUBLE number strategy decodes untyped integers as Long. + assertThat(functionResponse!!.response).isEqualTo(mapOf("value" to 456L)) + } + + @Test + fun taskToEvent_withUsageMetadata_returnsEvent() { + val statusMessage = + Message.builder() + .role(Message.Role.ROLE_AGENT) + .parts(listOf(TextPart("Status message"))) + .build() + val status = TaskStatus(TaskState.TASK_STATE_WORKING, statusMessage, null) + + val usageJson = "{\"promptTokenCount\":5,\"totalTokenCount\":12}" + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(status) + .metadata(mapOf(MetadataKeys.USAGE to usageJson)) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result.usageMetadata).isNotNull() + assertThat(result.usageMetadata?.promptTokenCount).isEqualTo(5) + assertThat(result.usageMetadata?.totalTokenCount).isEqualTo(12) + } + + @Test + fun taskToEvent_withInputRequiredState_emptyContent_returnsFinalEvent() { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_INPUT_REQUIRED)) + .artifacts(emptyList()) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result.turnComplete).isTrue() + assertThat(result.content).isNull() + } + + @Test + fun taskToEvent_withCompletedState_emptyContent_returnsFinalEvent() { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_COMPLETED)) + .artifacts(emptyList()) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result.turnComplete).isTrue() + assertThat(result.content).isNull() + } + + @Test + fun taskToEvent_withWorkingState_emptyContent_returnsEmptyEvent() { + val task = + Task.builder() + .id("task-1") + .contextId("context-1") + .status(TaskStatus(TaskState.TASK_STATE_WORKING)) + .artifacts(emptyList()) + .build() + + val result = task.toAdkEvent(invocationContext) + assertThat(result.content?.parts).isEmpty() + assertThat(result.turnComplete).isNotEqualTo(true) + } + + @Test + fun serializerFor_knownMetadataTypes_returnsSerializer() { + assertThat(serializerFor(GroundingMetadata::class)).isNotNull() + assertThat(serializerFor(UsageMetadata::class)).isNotNull() + } + + @Test + fun serializerFor_unknownType_returnsNull() { + assertThat(serializerFor(String::class)).isNull() + } +} diff --git a/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/agent/LegacyA2AAgent.kt b/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/agent/LegacyA2AAgent.kt index 7c635a30..8fd1e194 100644 --- a/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/agent/LegacyA2AAgent.kt +++ b/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/agent/LegacyA2AAgent.kt @@ -232,6 +232,10 @@ internal class LegacyA2AAgent( } /** Factory function to create a [BaseRemoteA2AAgent] for JVM/Android. */ +@Deprecated( + "Use A2AAgent with the A2A v1.0 SDK: build from an AgentCard or its URL " + + "(A2AAgent(name, agentCard, ...) / A2AAgent(name, agentCardUrl, ...)) instead of passing a Client." +) fun JvmA2AAgent( name: String, client: Client, diff --git a/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/converters/LegacyA2aConverters.kt b/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/converters/LegacyA2aConverters.kt index c7e3a5a9..71dcbaa1 100644 --- a/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/converters/LegacyA2aConverters.kt +++ b/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/converters/LegacyA2aConverters.kt @@ -71,11 +71,6 @@ private val metadataParser = private val PENDING_STATES = setOf(TaskState.WORKING, TaskState.SUBMITTED) -// DataPart types -internal const val TYPE_FUNCTION_CALL = "function_call" -internal const val TYPE_FUNCTION_RESPONSE = "function_response" -internal const val DEFAULT_ERROR_MESSAGE = "A2A task failed" - /** Converts a A2A [ClientEvent] to an ADK [Event]. */ internal fun ClientEvent.toAdkEvent(invocationContext: InvocationContext): Event? { return when (this) { diff --git a/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/jvm/A2AAgent.kt b/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/jvm/A2AAgent.kt new file mode 100644 index 00000000..786c1627 --- /dev/null +++ b/a2a/src/jvmMain/kotlin/com/google/adk/kt/a2a/jvm/A2AAgent.kt @@ -0,0 +1,81 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.jvm + +import com.google.adk.kt.a2a.agent.A2AAgentImpl +import com.google.adk.kt.a2a.agent.BaseRemoteA2AAgent +import com.google.adk.kt.a2a.agent.resolveAgentCard +import com.google.adk.kt.agents.BaseAgent +import com.google.adk.kt.callbacks.AfterAgentCallback +import com.google.adk.kt.callbacks.BeforeAgentCallback +import org.a2aproject.sdk.client.Client +import org.a2aproject.sdk.client.config.ClientConfig +import org.a2aproject.sdk.client.http.A2AHttpClient +import org.a2aproject.sdk.client.http.JdkA2AHttpClient +import org.a2aproject.sdk.client.transport.jsonrpc.JSONRPCTransport +import org.a2aproject.sdk.client.transport.jsonrpc.JSONRPCTransportConfig +import org.a2aproject.sdk.spec.AgentCard + +/** + * Builds a JVM A2A agent from an already-resolved [agentCard], wiring up the client so the caller + * never supplies a client and card separately. + */ +fun A2AAgent( + name: String, + agentCard: AgentCard, + httpClient: A2AHttpClient = JdkA2AHttpClient(), + streaming: Boolean = true, + subAgents: List = emptyList(), + beforeAgentCallbacks: List = emptyList(), + afterAgentCallbacks: List = emptyList(), +): BaseRemoteA2AAgent = + A2AAgentImpl( + name = name, + a2aClient = + Client.builder(agentCard) + .clientConfig(ClientConfig.Builder().setStreaming(streaming).build()) + .withTransport(JSONRPCTransport::class.java, JSONRPCTransportConfig(httpClient)) + .build(), + agentCard = agentCard, + streaming = streaming, + subAgents = subAgents, + beforeAgentCallbacks = beforeAgentCallbacks, + afterAgentCallbacks = afterAgentCallbacks, + ) + +/** + * Builds a JVM A2A agent from [agentCardUrl], auto-fetching the [AgentCard] from the remote agent's + * `/.well-known/agent-card.json` (like ADK Python/Go). Suspends on the network fetch. + */ +suspend fun A2AAgent( + name: String, + agentCardUrl: String, + httpClient: A2AHttpClient = JdkA2AHttpClient(), + streaming: Boolean = true, + subAgents: List = emptyList(), + beforeAgentCallbacks: List = emptyList(), + afterAgentCallbacks: List = emptyList(), +): BaseRemoteA2AAgent = + A2AAgent( + name = name, + agentCard = resolveAgentCard(httpClient, agentCardUrl), + httpClient = httpClient, + streaming = streaming, + subAgents = subAgents, + beforeAgentCallbacks = beforeAgentCallbacks, + afterAgentCallbacks = afterAgentCallbacks, + ) diff --git a/a2a/src/jvmTest/kotlin/com/google/adk/kt/a2a/agent/LegacyA2AAgentTest.kt b/a2a/src/jvmTest/kotlin/com/google/adk/kt/a2a/agent/LegacyA2AAgentTest.kt index e9c86571..a1d04191 100644 --- a/a2a/src/jvmTest/kotlin/com/google/adk/kt/a2a/agent/LegacyA2AAgentTest.kt +++ b/a2a/src/jvmTest/kotlin/com/google/adk/kt/a2a/agent/LegacyA2AAgentTest.kt @@ -600,6 +600,7 @@ class LegacyA2AAgentTest { verifyNoInteractions(mockClient) } + @Suppress("DEPRECATION") // Exercises the deprecated JvmA2AAgent on purpose. private fun createTestAgent( streaming: Boolean = true, beforeCallbacks: List = emptyList(), diff --git a/a2a/src/jvmTest/kotlin/com/google/adk/kt/a2a/jvm/A2AAgentTest.kt b/a2a/src/jvmTest/kotlin/com/google/adk/kt/a2a/jvm/A2AAgentTest.kt new file mode 100644 index 00000000..246b842f --- /dev/null +++ b/a2a/src/jvmTest/kotlin/com/google/adk/kt/a2a/jvm/A2AAgentTest.kt @@ -0,0 +1,79 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.a2a.jvm + +import com.google.common.truth.Truth.assertThat +import kotlinx.coroutines.test.runTest +import okhttp3.mockwebserver.MockResponse +import okhttp3.mockwebserver.MockWebServer +import org.a2aproject.sdk.jsonrpc.common.json.JsonUtil +import org.a2aproject.sdk.spec.AgentCapabilities +import org.a2aproject.sdk.spec.AgentCard +import org.a2aproject.sdk.spec.AgentInterface +import org.a2aproject.sdk.spec.TransportProtocol +import org.junit.After +import org.junit.Before +import org.junit.Test +import org.junit.runner.RunWith +import org.junit.runners.JUnit4 + +/** + * Verifies the JVM agent-card auto-fetch: [A2AAgent] fetches and parses a card from the standard + * `/.well-known/agent-card.json` endpoint (served here by [MockWebServer]). + */ +@RunWith(JUnit4::class) +class A2AAgentTest { + + private lateinit var server: MockWebServer + + @Before + fun setUp() { + server = MockWebServer() + server.start() + } + + @After + fun tearDown() { + server.shutdown() + } + + private fun agentCard(url: String): AgentCard = + AgentCard.builder() + .name("remote-agent") + .description("Remote Agent") + .url(url) + .version("1.0.0") + .defaultInputModes(listOf("text")) + .defaultOutputModes(listOf("text")) + .skills(listOf()) + .supportedInterfaces(listOf(AgentInterface(TransportProtocol.JSONRPC.asString(), url))) + .capabilities(AgentCapabilities.builder().streaming(false).build()) + .build() + + @Test + fun a2aAgent_autoFetchesCard_andPopulatesDescription() = runTest { + val baseUrl = server.url("/").toString() + server.enqueue(MockResponse().setBody(JsonUtil.toJson(agentCard(baseUrl)))) + + val agent = A2AAgent(name = "remote-agent", agentCardUrl = baseUrl, streaming = false) + + assertThat(agent.description).isEqualTo("Remote Agent") + assertThat(recordedPath()).isEqualTo("/.well-known/agent-card.json") + } + + private fun recordedPath(): String? = server.takeRequest().path +} diff --git a/examples/src/main/kotlin/com/google/adk/kt/examples/a2a/A2AAgentDemo.kt b/examples/src/main/kotlin/com/google/adk/kt/examples/a2a/A2AAgentDemo.kt new file mode 100644 index 00000000..febbe34e --- /dev/null +++ b/examples/src/main/kotlin/com/google/adk/kt/examples/a2a/A2AAgentDemo.kt @@ -0,0 +1,44 @@ +/* + * Copyright 2026 Google LLC + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package com.google.adk.kt.examples.a2a + +import com.google.adk.kt.a2a.jvm.A2AAgent +import kotlinx.coroutines.runBlocking + +/** + * Example agent demonstrating how to use [A2AAgent] to communicate with a remote A2A-compliant + * agent on the JVM. + * + * This demo showcases: + * 1. Auto-fetching the remote agent's `AgentCard` from its `/.well-known/agent-card.json` endpoint + * (the common case). To supply a pre-built card instead, use the `A2AAgent(name, agentCard)` + * overload. + * 2. Talking to the agent over the JSON-RPC transport backed by `JdkA2AHttpClient`, the JVM HTTP + * client. (On Android, use `androidA2AAgent`, which uses `AndroidA2AHttpClient`.) The factory + * injects the transport explicitly rather than relying on SDK ServiceLoader auto-resolution. + * 3. Wrapping the remote agent as a standard ADK [com.google.adk.kt.agents.Agent]. + */ +object A2AAgentDemo { + + @JvmField + val rootAgent = run { + val agentUrl = System.getenv("A2A_AGENT_URL") ?: "http://localhost:8080/a2a" + val agentName = System.getenv("A2A_AGENT_NAME") ?: "remote-agent" + // Non-streaming so the Client uses `message/send`, not SSE `message/stream`. + runBlocking { A2AAgent(name = agentName, agentCardUrl = agentUrl, streaming = false) } + } +} diff --git a/examples/src/main/kotlin/com/google/adk/kt/examples/a2a/JvmA2AAgentDemo.kt b/examples/src/main/kotlin/com/google/adk/kt/examples/a2a/JvmA2AAgentDemo.kt deleted file mode 100644 index 0c632cf1..00000000 --- a/examples/src/main/kotlin/com/google/adk/kt/examples/a2a/JvmA2AAgentDemo.kt +++ /dev/null @@ -1,77 +0,0 @@ -/* - * Copyright 2026 Google LLC - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ - -package com.google.adk.kt.examples.a2a - -import com.google.adk.kt.a2a.agent.JvmA2AAgent -import io.a2a.client.Client -import io.a2a.client.transport.jsonrpc.JSONRPCTransport -import io.a2a.client.transport.jsonrpc.JSONRPCTransportConfig -import io.a2a.spec.AgentCapabilities -import io.a2a.spec.AgentCard - -/** - * Example agent demonstrating how to use [RemoteA2AAgent] to communicate with a remote - * A2A-compliant agent. - * - * This demo showcases: - * 1. Initializing an [io.a2a.client.Client] with a specific transport (REST). - * 2. Using [ServiceLoader] (implicit via A2ACardResolver) for platform-specific HTTP client - * resolution (JDK vs Android). - * 3. Wrapping the remote agent as a standard ADK [com.google.adk.kt.agents.Agent]. - */ -object JvmA2AAgentDemo { - - @JvmField - val rootAgent = run { - println("Starting JvmA2AAgentDemo...") - - val agentUrl = System.getenv("A2A_AGENT_URL") ?: "http://localhost:8080/a2a" - val agentName = System.getenv("A2A_AGENT_NAME") ?: "remote-agent" - - val agentCard = - AgentCard.Builder() - .name(agentName) - .url(agentUrl) - .description("A remote A2A agent") - .version("1.0.0") - .protocolVersion("0.3.0") - .preferredTransport("JSONRPC") - .defaultInputModes(listOf("text")) - .defaultOutputModes(listOf("text")) - .capabilities( - // Advertise non-streaming so the a2a Client uses `message/send` instead - // of SSE `message/stream`. This is what selects the transport in - // io.a2a.client.Client (it checks agentCard.capabilities().streaming()). - AgentCapabilities.Builder() - .streaming(false) - .pushNotifications(false) - .stateTransitionHistory(false) - .build() - ) - .skills(emptyList()) - .build() - - val a2aClient = - Client.builder(agentCard) - .withTransport(JSONRPCTransport::class.java, JSONRPCTransportConfig()) - .build() - - // Use non-streaming (message/send) so the demo works against any A2A server - // regardless of its streaming support. - JvmA2AAgent(name = agentName, client = a2aClient, agentCard = agentCard, streaming = false) - } -} diff --git a/gradle/libs.versions.toml b/gradle/libs.versions.toml index dc79402a..3713dcee 100644 --- a/gradle/libs.versions.toml +++ b/gradle/libs.versions.toml @@ -1,7 +1,7 @@ [versions] # go/keep-sorted start -a2a = "0.3.2.Final" +a2a = "1.0.0.Final" a2a-legacy = "0.3.3.Final" androidx-compose-ui = "1.11.2" androidx-core = "1.16.0" @@ -50,11 +50,12 @@ snakeyaml = "2.2" a2a-legacy-sdk-client = { module = "io.github.a2asdk:a2a-java-sdk-client", version.ref = "a2a-legacy" } a2a-legacy-sdk-common = { module = "io.github.a2asdk:a2a-java-sdk-common", version.ref = "a2a-legacy" } a2a-legacy-sdk-spec = { module = "io.github.a2asdk:a2a-java-sdk-spec", version.ref = "a2a-legacy" } -a2a-sdk-client = { module = "io.github.a2asdk:a2a-java-sdk-client", version.ref = "a2a" } -a2a-sdk-common = { module = "io.github.a2asdk:a2a-java-sdk-common", version.ref = "a2a" } -a2a-sdk-spec = { module = "io.github.a2asdk:a2a-java-sdk-spec", version.ref = "a2a" } -a2a-sdk-transport-jsonrpc = { module = "io.github.a2asdk:a2a-java-sdk-transport-jsonrpc", version.ref = "a2a" } -a2a-sdk-transport-rest = { module = "io.github.a2asdk:a2a-java-sdk-transport-rest", version.ref = "a2a" } +a2a-sdk-client = { module = "org.a2aproject.sdk:a2a-java-sdk-client", version.ref = "a2a" } +a2a-sdk-common = { module = "org.a2aproject.sdk:a2a-java-sdk-common", version.ref = "a2a" } +a2a-sdk-http-client-android = { module = "org.a2aproject.sdk:a2a-java-sdk-http-client-android", version.ref = "a2a" } +a2a-sdk-spec = { module = "org.a2aproject.sdk:a2a-java-sdk-spec", version.ref = "a2a" } +a2a-sdk-transport-jsonrpc = { module = "org.a2aproject.sdk:a2a-java-sdk-client-transport-jsonrpc", version.ref = "a2a" } +a2a-sdk-transport-rest = { module = "org.a2aproject.sdk:a2a-java-sdk-client-transport-rest", version.ref = "a2a" } androidx-compose-ui-test-junit4 = { module = "androidx.compose.ui:ui-test-junit4", version.ref = "androidx-compose-ui" } androidx-compose-ui-test-manifest = { module = "androidx.compose.ui:ui-test-manifest", version.ref = "androidx-compose-ui" } androidx-core = { module = "androidx.core:core", version.ref = "androidx-core" }