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 447b6ed10..20b78921a 100644 --- a/core/src/main/java/com/google/adk/models/GeminiLlmConnection.java +++ b/core/src/main/java/com/google/adk/models/GeminiLlmConnection.java @@ -127,6 +127,7 @@ private void handleServerMessage(LiveServerMessage message) { /** Converts a server message into the standardized LlmResponse format. */ static Optional convertToServerResponse(LiveServerMessage message) { LlmResponse.Builder builder = LlmResponse.builder(); + boolean hasRelevantData = false; if (message.serverContent().isPresent()) { LiveServerContent serverContent = message.serverContent().get(); @@ -140,6 +141,7 @@ static Optional convertToServerResponse(LiveServerMessage message) // overwrite the audio modelTurn content. serverContent.outputTranscription().ifPresent(builder::outputTranscription); serverContent.inputTranscription().ifPresent(builder::inputTranscription); + hasRelevantData = true; } else if (message.toolCall().isPresent()) { LiveServerToolCall toolCall = message.toolCall().get(); toolCall @@ -154,24 +156,37 @@ static Optional convertToServerResponse(LiveServerMessage message) } }); builder.partial(false).turnComplete(false); - } else if (message.usageMetadata().isPresent()) { - logger.debug("Received usage metadata: {}", message.usageMetadata().get()); - return Optional.empty(); + hasRelevantData = true; } else if (message.toolCallCancellation().isPresent()) { logger.debug("Received tool call cancellation: {}", message.toolCallCancellation().get()); builder.interrupted(true).turnComplete(true); - return Optional.of(builder.build()); + hasRelevantData = true; } else if (message.setupComplete().isPresent()) { logger.debug("Received setup complete."); return Optional.empty(); - } else { + } else if (message.usageMetadata().isEmpty()) { logger.warn("Received unknown or empty server message: {}", message.toJson()); builder .errorCode(new FinishReason("Unknown server message.")) .errorMessage("Received unknown server message."); + hasRelevantData = true; + } + + if (message.usageMetadata().isPresent()) { + logger.debug("Received usage metadata: {}", message.usageMetadata().get()); + builder.usageMetadata( + GeminiUtil.toGenerateContentResponseUsageMetadata(message.usageMetadata().get())); + if (!hasRelevantData) { + builder.partial(false).turnComplete(false); + } + hasRelevantData = true; + } + + if (hasRelevantData) { + return Optional.of(builder.build()); } - return Optional.of(builder.build()); + return Optional.empty(); } /** Handles errors that occur *during* the initial connection attempt. */ diff --git a/core/src/main/java/com/google/adk/models/azure/AzureRealtimeLlmConnection.java b/core/src/main/java/com/google/adk/models/azure/AzureRealtimeLlmConnection.java index 4142dd997..48839bb65 100644 --- a/core/src/main/java/com/google/adk/models/azure/AzureRealtimeLlmConnection.java +++ b/core/src/main/java/com/google/adk/models/azure/AzureRealtimeLlmConnection.java @@ -7,6 +7,8 @@ import com.google.genai.types.Blob; import com.google.genai.types.Content; import com.google.genai.types.FunctionCall; +import com.google.genai.types.GenerateContentResponseUsageMetadata; +import com.google.genai.types.ModalityTokenCount; import com.google.genai.types.Part; import com.google.genai.types.Transcription; import io.reactivex.rxjava3.core.Completable; @@ -548,6 +550,33 @@ private void handleResponseDone(JSONObject event) { "Realtime token usage — input: {}, output: {}", usage.optInt("input_tokens", 0), usage.optInt("output_tokens", 0)); + + GenerateContentResponseUsageMetadata.Builder usageBuilder = + GenerateContentResponseUsageMetadata.builder() + .promptTokenCount(usage.optInt("input_tokens", 0)) + .candidatesTokenCount(usage.optInt("output_tokens", 0)) + .totalTokenCount(usage.optInt("total_tokens", 0)); + + JSONObject inputDetails = usage.optJSONObject("input_token_details"); + if (inputDetails != null && inputDetails.has("audio_tokens")) { + usageBuilder.promptTokensDetails( + ImmutableList.of( + ModalityTokenCount.builder() + .modality(com.google.genai.types.MediaModality.Known.AUDIO) + .tokenCount(inputDetails.optInt("audio_tokens", 0)) + .build())); + } + + JSONObject outputDetails = usage.optJSONObject("output_token_details"); + if (outputDetails != null && outputDetails.has("audio_tokens")) { + usageBuilder.candidatesTokensDetails( + ImmutableList.of( + ModalityTokenCount.builder() + .modality(com.google.genai.types.MediaModality.Known.AUDIO) + .tokenCount(outputDetails.optInt("audio_tokens", 0)) + .build())); + } + responseProcessor.onNext(LlmResponse.builder().usageMetadata(usageBuilder.build()).build()); } } } 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 5577a8a47..6a95e1532 100644 --- a/core/src/test/java/com/google/adk/models/GeminiLlmConnectionTest.java +++ b/core/src/test/java/com/google/adk/models/GeminiLlmConnectionTest.java @@ -152,13 +152,15 @@ public void convertToServerResponse_withToolCall_mapsContentWithFunctionCall() { } @Test - public void convertToServerResponse_withUsageMetadata_returnsEmpty() { - LiveServerMessage message = - LiveServerMessage.builder().usageMetadata(UsageMetadata.builder().build()).build(); + public void convertToServerResponse_withUsageMetadata_returnsResponseWithUsage() { + UsageMetadata usageMetadata = UsageMetadata.builder().promptTokenCount(10).build(); + LiveServerMessage message = LiveServerMessage.builder().usageMetadata(usageMetadata).build(); Optional result = GeminiLlmConnection.convertToServerResponse(message); - assertThat(result.isPresent()).isFalse(); + assertThat(result.isPresent()).isTrue(); + assertThat(result.get().usageMetadata()).isPresent(); + assertThat(result.get().usageMetadata().get().promptTokenCount()).hasValue(10); } @Test