diff --git a/core/src/main/java/com/google/adk/models/Gemini.java b/core/src/main/java/com/google/adk/models/Gemini.java index 8b4d95298..5195e570b 100644 --- a/core/src/main/java/com/google/adk/models/Gemini.java +++ b/core/src/main/java/com/google/adk/models/Gemini.java @@ -313,7 +313,17 @@ public BaseLlmConnection connect(LlmRequest llmRequest) { logger.debug("Connecting to model {}", effectiveModelName); logger.trace("Connection Config: {}", liveConnectConfig); - return new GeminiLlmConnection(apiClient, effectiveModelName, liveConnectConfig); + return new GeminiLlmConnection(connectLiveTransport(effectiveModelName, liveConnectConfig)); + } + + /** + * Opens the live transport the connection drives. Overridable so a test can supply an in-process + * {@link GeminiLiveTransport} double and exercise the real {@link GeminiLlmConnection} without a + * network. + */ + protected CompletableFuture connectLiveTransport( + String modelName, LiveConnectConfig config) { + return apiClient.async.live.connect(modelName, config).thenApply(GenAiLiveTransport::new); } private static final class StreamingResponseAggregator { diff --git a/core/src/main/java/com/google/adk/models/GeminiLiveTransport.java b/core/src/main/java/com/google/adk/models/GeminiLiveTransport.java new file mode 100644 index 000000000..bded74086 --- /dev/null +++ b/core/src/main/java/com/google/adk/models/GeminiLiveTransport.java @@ -0,0 +1,53 @@ +/* + * Copyright 2025 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.models; + +import com.google.genai.types.LiveSendClientContentParameters; +import com.google.genai.types.LiveSendRealtimeInputParameters; +import com.google.genai.types.LiveSendToolResponseParameters; +import com.google.genai.types.LiveServerMessage; +import java.util.concurrent.CompletableFuture; +import java.util.function.Consumer; + +/** + * The bidirectional live transport that {@link GeminiLlmConnection} drives. + * + *

{@link GeminiLlmConnection} holds the translation logic; this is the transport beneath it, so + * the connection can drive any implementation. The default one delegates to a genai live session. + */ +public interface GeminiLiveTransport { + + /** Sends a client-content turn to the transport. */ + CompletableFuture sendClientContent(LiveSendClientContentParameters params); + + /** Sends realtime input (audio, video, or text) to the transport. */ + CompletableFuture sendRealtimeInput(LiveSendRealtimeInputParameters params); + + /** Sends a tool response to the transport. */ + CompletableFuture sendToolResponse(LiveSendToolResponseParameters params); + + /** + * Registers the callback for messages the transport yields and a callback for when its receive + * stream ends. An implementation whose stream ends only when the client closes never invokes + * {@code onStreamEnd}; one backed by a finite script invokes it when the script is exhausted, so + * the run can end. + */ + CompletableFuture receive(Consumer onMessage, Runnable onStreamEnd); + + /** Closes the transport. */ + CompletableFuture close(); +} diff --git a/core/src/main/java/com/google/adk/models/GeminiLlmConnection.java b/core/src/main/java/com/google/adk/models/GeminiLlmConnection.java index 35dadd9fd..5b2b26814 100644 --- a/core/src/main/java/com/google/adk/models/GeminiLlmConnection.java +++ b/core/src/main/java/com/google/adk/models/GeminiLlmConnection.java @@ -19,13 +19,10 @@ import static com.google.common.collect.ImmutableList.toImmutableList; import com.google.common.collect.ImmutableList; -import com.google.genai.AsyncSession; -import com.google.genai.Client; import com.google.genai.types.Blob; import com.google.genai.types.Content; import com.google.genai.types.FinishReason; import com.google.genai.types.FunctionResponse; -import com.google.genai.types.LiveConnectConfig; import com.google.genai.types.LiveSendClientContentParameters; import com.google.genai.types.LiveSendRealtimeInputParameters; import com.google.genai.types.LiveSendToolResponseParameters; @@ -61,38 +58,26 @@ public final class GeminiLlmConnection implements BaseLlmConnection { private static final Logger logger = LoggerFactory.getLogger(GeminiLlmConnection.class); - private final Client apiClient; - private final String modelName; - private final LiveConnectConfig connectConfig; - private final CompletableFuture sessionFuture; + private final CompletableFuture transportFuture; private final PublishProcessor responseProcessor = PublishProcessor.create(); private final Flowable responseFlowable = responseProcessor.serialize(); private final CompositeDisposable disposables = new CompositeDisposable(); private final AtomicBoolean closed = new AtomicBoolean(false); /** - * Establishes a new connection. + * Establishes a new connection over the given live transport. * - * @param apiClient The API client for communication. - * @param modelName The specific Gemini model endpoint (e.g., "gemini-2.0-flash). - * @param connectConfig Configuration parameters for the live session. + * @param transportFuture The live transport the connection drives, once established. */ - GeminiLlmConnection(Client apiClient, String modelName, LiveConnectConfig connectConfig) { - this.apiClient = Objects.requireNonNull(apiClient); - this.modelName = Objects.requireNonNull(modelName); - this.connectConfig = Objects.requireNonNull(connectConfig); - - this.sessionFuture = - this.apiClient - .async - .live - .connect(this.modelName, this.connectConfig) + GeminiLlmConnection(CompletableFuture transportFuture) { + this.transportFuture = + Objects.requireNonNull(transportFuture) .whenCompleteAsync( - (session, throwable) -> { + (transport, throwable) -> { if (throwable != null) { handleConnectionError(throwable); - } else if (session != null) { - setupReceiver(session); + } else if (transport != null) { + setupReceiver(transport); } else if (!closed.get()) { handleConnectionError( new SocketException("WebSocket connection failed without explicit error.")); @@ -100,14 +85,14 @@ public final class GeminiLlmConnection implements BaseLlmConnection { }); } - /** Configures the session to forward incoming messages to the response processor. */ - private void setupReceiver(AsyncSession session) { + /** Configures the transport to forward incoming messages to the response processor. */ + private void setupReceiver(GeminiLiveTransport transport) { if (closed.get()) { - closeSessionIgnoringErrors(session); + closeTransportIgnoringErrors(transport); return; } - session - .receive(this::handleServerMessage) + transport + .receive(this::handleServerMessage, this::completeReceive) .exceptionally( error -> { handleReceiveError(error); @@ -115,6 +100,15 @@ private void setupReceiver(AsyncSession session) { }); } + /** Completes the response stream when the transport's receive stream ends, ending the run. */ + private void completeReceive() { + // Only a transport that ends its own stream reaches this, so it owns its close; none is issued. + if (closed.compareAndSet(false, true)) { + responseProcessor.onComplete(); + disposables.dispose(); + } + } + /** Processes messages received from the WebSocket server. */ private void handleServerMessage(LiveServerMessage message) { if (closed.get()) { @@ -241,7 +235,9 @@ private void handleReceiveError(Throwable throwable) { if (closed.compareAndSet(false, true)) { logger.error("Error during WebSocket receive operation", throwable); responseProcessor.onError(throwable); - sessionFuture.thenAccept(this::closeSessionIgnoringErrors).exceptionally(unusedError -> null); + transportFuture + .thenAccept(this::closeTransportIgnoringErrors) + .exceptionally(unusedError -> null); } } @@ -286,22 +282,22 @@ private List extractFunctionResponses(Content content) { @Override public Completable sendRealtime(Blob blob) { return Completable.fromFuture( - sessionFuture.thenCompose( - session -> - session.sendRealtimeInput( + transportFuture.thenCompose( + transport -> + transport.sendRealtimeInput( LiveSendRealtimeInputParameters.builder().media(blob).build()))); } /** Helper to send client content parameters. */ private Completable sendClientContentInternal(LiveSendClientContentParameters parameters) { return Completable.fromFuture( - sessionFuture.thenCompose(session -> session.sendClientContent(parameters))); + transportFuture.thenCompose(transport -> transport.sendClientContent(parameters))); } /** Helper to send tool response parameters. */ private Completable sendToolResponseInternal(LiveSendToolResponseParameters parameters) { return Completable.fromFuture( - sessionFuture.thenCompose(session -> session.sendToolResponse(parameters))); + transportFuture.thenCompose(transport -> transport.sendToolResponse(parameters))); } @Override @@ -331,26 +327,26 @@ private void closeInternal(Throwable throwable) { responseProcessor.onError(throwable); } - if (sessionFuture.isDone()) { - sessionFuture - .thenAccept(this::closeSessionIgnoringErrors) + if (transportFuture.isDone()) { + transportFuture + .thenAccept(this::closeTransportIgnoringErrors) .exceptionally(unusedError -> null); } else { - sessionFuture.cancel(false); + transportFuture.cancel(false); } disposables.dispose(); } } - /** Closes the AsyncSession safely, logging any errors. */ - private void closeSessionIgnoringErrors(AsyncSession session) { - if (session != null) { - session + /** Closes the transport safely, logging any errors. */ + private void closeTransportIgnoringErrors(GeminiLiveTransport transport) { + if (transport != null) { + transport .close() .exceptionally( closeError -> { - logger.warn("Error occurred while closing AsyncSession", closeError); + logger.warn("Error occurred while closing live transport", closeError); return null; // Suppress error during close }); } diff --git a/core/src/main/java/com/google/adk/models/GenAiLiveTransport.java b/core/src/main/java/com/google/adk/models/GenAiLiveTransport.java new file mode 100644 index 000000000..489b36ab5 --- /dev/null +++ b/core/src/main/java/com/google/adk/models/GenAiLiveTransport.java @@ -0,0 +1,63 @@ +/* + * Copyright 2025 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.models; + +import com.google.genai.AsyncSession; +import com.google.genai.types.LiveSendClientContentParameters; +import com.google.genai.types.LiveSendRealtimeInputParameters; +import com.google.genai.types.LiveSendToolResponseParameters; +import com.google.genai.types.LiveServerMessage; +import java.util.Objects; +import java.util.concurrent.CompletableFuture; +import java.util.function.Consumer; + +/** The default {@link GeminiLiveTransport}, delegating to a genai {@link AsyncSession}. */ +final class GenAiLiveTransport implements GeminiLiveTransport { + + private final AsyncSession session; + + GenAiLiveTransport(AsyncSession session) { + this.session = Objects.requireNonNull(session); + } + + @Override + public CompletableFuture sendClientContent(LiveSendClientContentParameters params) { + return session.sendClientContent(params); + } + + @Override + public CompletableFuture sendRealtimeInput(LiveSendRealtimeInputParameters params) { + return session.sendRealtimeInput(params); + } + + @Override + public CompletableFuture sendToolResponse(LiveSendToolResponseParameters params) { + return session.sendToolResponse(params); + } + + @Override + public CompletableFuture receive( + Consumer onMessage, Runnable onStreamEnd) { + // The genai session has no end-of-stream signal, so onStreamEnd never fires here. + return session.receive(onMessage); + } + + @Override + public CompletableFuture close() { + return session.close(); + } +} diff --git a/core/src/test/java/com/google/adk/models/GeminiLlmConnectionTest.java b/core/src/test/java/com/google/adk/models/GeminiLlmConnectionTest.java index b15a65852..d212d536f 100644 --- a/core/src/test/java/com/google/adk/models/GeminiLlmConnectionTest.java +++ b/core/src/test/java/com/google/adk/models/GeminiLlmConnectionTest.java @@ -23,6 +23,9 @@ import com.google.genai.types.Content; import com.google.genai.types.FunctionCall; import com.google.genai.types.GenerateContentResponseUsageMetadata; +import com.google.genai.types.LiveSendClientContentParameters; +import com.google.genai.types.LiveSendRealtimeInputParameters; +import com.google.genai.types.LiveSendToolResponseParameters; import com.google.genai.types.LiveServerContent; import com.google.genai.types.LiveServerMessage; import com.google.genai.types.LiveServerSetupComplete; @@ -31,7 +34,12 @@ import com.google.genai.types.Part; import com.google.genai.types.UsageMetadata; import io.reactivex.rxjava3.observers.TestObserver; +import io.reactivex.rxjava3.subscribers.TestSubscriber; import java.util.List; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.function.Consumer; import org.junit.Test; import org.junit.runner.RunWith; import org.junit.runners.JUnit4; @@ -326,4 +334,78 @@ public void convertToServerResponse_withContentAndUsageMetadata_emitsMultiple() .build(); assertThat(usageResponse.usageMetadata()).hasValue(expectedUsageMetadata); } + + @Test + public void receive_completesWhenTransportSignalsStreamEnd() { + // A transport that ends its own receive stream completes the connection's response stream, so a + // live run with no client close still terminates. + FakeTransport transport = new FakeTransport(/* endStreamOnReceive= */ true); + GeminiLlmConnection connection = + new GeminiLlmConnection(CompletableFuture.completedFuture(transport)); + + TestSubscriber subscriber = connection.receive().test(); + subscriber.awaitDone(5, TimeUnit.SECONDS); + + subscriber.assertComplete(); + subscriber.assertNoErrors(); + } + + @Test + public void close_afterTransportStreamEnd_isNoOp() { + FakeTransport transport = new FakeTransport(/* endStreamOnReceive= */ true); + GeminiLlmConnection connection = + new GeminiLlmConnection(CompletableFuture.completedFuture(transport)); + TestSubscriber subscriber = connection.receive().test(); + subscriber.awaitDone(5, TimeUnit.SECONDS); + subscriber.assertComplete(); + + connection.close(); // Already completed via stream-end; must not re-terminate or error. + + subscriber.assertComplete(); + subscriber.assertNoErrors(); + // completeReceive leaves the transport to close itself, so no close() is issued on this path. + assertThat(transport.closeCount.get()).isEqualTo(0); + } + + /** + * A minimal in-process {@link GeminiLiveTransport} that records closes and can end its stream. + */ + private static final class FakeTransport implements GeminiLiveTransport { + private final boolean endStreamOnReceive; + final AtomicInteger closeCount = new AtomicInteger(); + + FakeTransport(boolean endStreamOnReceive) { + this.endStreamOnReceive = endStreamOnReceive; + } + + @Override + public CompletableFuture sendClientContent(LiveSendClientContentParameters params) { + return CompletableFuture.completedFuture(null); + } + + @Override + public CompletableFuture sendRealtimeInput(LiveSendRealtimeInputParameters params) { + return CompletableFuture.completedFuture(null); + } + + @Override + public CompletableFuture sendToolResponse(LiveSendToolResponseParameters params) { + return CompletableFuture.completedFuture(null); + } + + @Override + public CompletableFuture receive( + Consumer onMessage, Runnable onStreamEnd) { + if (endStreamOnReceive) { + onStreamEnd.run(); + } + return CompletableFuture.completedFuture(null); + } + + @Override + public CompletableFuture close() { + closeCount.incrementAndGet(); + return CompletableFuture.completedFuture(null); + } + } }