From ad628b52f11dedbbf35bc1b9c332d9cc5af92056 Mon Sep 17 00:00:00 2001 From: Guillaume Laforge Date: Wed, 23 Sep 2026 08:59:23 +0200 Subject: [PATCH] feat: handle usageUpdate events and populate UsageMetadata in AgentResponse --- .flattened-pom.xml | 4 +- antigravity-sdk-harness/.flattened-pom.xml | 4 +- antigravity-sdk-protocol/.flattened-pom.xml | 4 +- antigravity-sdk-wrapper/.flattened-pom.xml | 4 +- .../io/github/glaforge/antigravity/Agent.java | 199 ++++++++++++++++-- .../glaforge/antigravity/AgentResponse.java | 9 + .../glaforge/antigravity/UsageMetadata.java | 7 +- .../antigravity/ObservabilityTest.java | 17 +- .../antigravity/UsageMetadataParsingTest.java | 115 ++++++++++ .../references/api-reference.md | 4 + 10 files changed, 332 insertions(+), 35 deletions(-) create mode 100644 antigravity-sdk-wrapper/src/test/java/io/github/glaforge/antigravity/UsageMetadataParsingTest.java diff --git a/.flattened-pom.xml b/.flattened-pom.xml index 514b3c5..c62d85f 100644 --- a/.flattened-pom.xml +++ b/.flattened-pom.xml @@ -4,7 +4,7 @@ 4.0.0 io.github.glaforge.antigravity antigravity-sdk-parent - 0.2.14-SNAPSHOT + 0.2.16-SNAPSHOT pom Antigravity Java SDK A Java SDK for building agents using the Antigravity localharness. @@ -40,7 +40,7 @@ 3.25.3 UTF-8 2.17.1 - 0.2.14-SNAPSHOT + 0.2.16-SNAPSHOT diff --git a/antigravity-sdk-harness/.flattened-pom.xml b/antigravity-sdk-harness/.flattened-pom.xml index 4904e25..07cfd89 100644 --- a/antigravity-sdk-harness/.flattened-pom.xml +++ b/antigravity-sdk-harness/.flattened-pom.xml @@ -5,10 +5,10 @@ io.github.glaforge.antigravity antigravity-sdk-parent - 0.2.14-SNAPSHOT + 0.2.16-SNAPSHOT antigravity-sdk-harness - 0.2.14-SNAPSHOT + 0.2.16-SNAPSHOT Antigravity Java SDK - Harness Binaries Platform-specific native localharness Go binaries for offline deployment. https://github.com/glaforge/antigravity-java-sdk diff --git a/antigravity-sdk-protocol/.flattened-pom.xml b/antigravity-sdk-protocol/.flattened-pom.xml index 70ad795..ab88fda 100644 --- a/antigravity-sdk-protocol/.flattened-pom.xml +++ b/antigravity-sdk-protocol/.flattened-pom.xml @@ -5,10 +5,10 @@ io.github.glaforge.antigravity antigravity-sdk-parent - 0.2.14-SNAPSHOT + 0.2.16-SNAPSHOT antigravity-sdk-protocol - 0.2.14-SNAPSHOT + 0.2.16-SNAPSHOT Antigravity Java SDK - Protocol The Protocol Buffers definitions and generated classes for Antigravity Java SDK. https://github.com/glaforge/antigravity-java-sdk diff --git a/antigravity-sdk-wrapper/.flattened-pom.xml b/antigravity-sdk-wrapper/.flattened-pom.xml index f5ba69f..eebfe7f 100644 --- a/antigravity-sdk-wrapper/.flattened-pom.xml +++ b/antigravity-sdk-wrapper/.flattened-pom.xml @@ -5,10 +5,10 @@ io.github.glaforge.antigravity antigravity-sdk-parent - 0.2.14-SNAPSHOT + 0.2.16-SNAPSHOT antigravity-sdk-wrapper - 0.2.14-SNAPSHOT + 0.2.16-SNAPSHOT Antigravity Java SDK - Wrapper The main Java wrapper and client for the Antigravity localharness. https://github.com/glaforge/antigravity-java-sdk diff --git a/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/Agent.java b/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/Agent.java index a1df30c..767e514 100644 --- a/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/Agent.java +++ b/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/Agent.java @@ -64,6 +64,8 @@ import java.util.function.Consumer; import java.util.List; import java.util.Set; +import java.util.Map; +import java.util.LinkedHashMap; import java.util.ArrayList; import java.util.Collections; import java.nio.file.Path; @@ -92,6 +94,9 @@ public class Agent implements AutoCloseable, TriggerContext { private StringBuilder currentText; private StringBuilder currentThoughts; private UsageMetadata currentUsage; + private UsageMetadata cumulativeUsage = new UsageMetadata(0, 0, 0, 0, 0); + private UsageMetadata turnStartUsage; + private final ConcurrentMap trajectoryUsages = new ConcurrentHashMap<>(); private final List policies; private boolean hasStructuredOutput; private StringBuilder wsBuffer = new StringBuilder(); @@ -116,7 +121,35 @@ public SandboxStatus getSandboxStatus() { * @return the usage metadata */ public UsageMetadata getUsageMetadata() { - return currentUsage; + if (currentUsage != null) { + return currentUsage; + } + if (cumulativeUsage != null && turnStartUsage != null) { + UsageMetadata diff = cumulativeUsage.subtract(turnStartUsage); + if (diff.totalTokenCount() > 0) { + return diff; + } + } + return cumulativeUsage; + } + + /** + * Returns the cumulative token usage across all turns in this session. + * + * @return the total cumulative usage metadata + */ + public UsageMetadata getTotalUsage() { + return cumulativeUsage; + } + + /** + * Returns a map of trajectory ID to cumulative token usage for subagents and + * main agent. + * + * @return an unmodifiable map of trajectory usage metadata + */ + public Map getTrajectoryUsages() { + return Collections.unmodifiableMap(new LinkedHashMap<>(trajectoryUsages)); } /** @@ -982,6 +1015,7 @@ public CompletableFuture chatStream(List inputs, Cons this.currentText = new StringBuilder(); this.currentThoughts = new StringBuilder(); this.hasStructuredOutput = false; + this.turnStartUsage = this.cumulativeUsage; this.currentUsage = null; StringBuilder combinedText = new StringBuilder(); @@ -1165,6 +1199,76 @@ void handleIncomingMessage(WebSocket webSocket, String message) { this.sandboxStatus = new SandboxStatus(available, reason); warnIfSandboxUnavailable(this.sandboxStatus); } + + if (initResp.has("cumulativeUsage") || initResp.has("cumulative_usage")) { + JsonNode cumNode = initResp.has("cumulativeUsage") + ? initResp.get("cumulativeUsage") + : initResp.get("cumulative_usage"); + UsageMetadata parsed = parseUsageMetadata(cumNode); + if (parsed != null) { + this.cumulativeUsage = parsed; + this.turnStartUsage = parsed; + } + } + + JsonNode trajUsageNode = initResp.has("trajectoryUsage") + ? initResp.get("trajectoryUsage") + : initResp.get("trajectory_usage"); + if (trajUsageNode != null && trajUsageNode.isArray()) { + for (JsonNode entry : trajUsageNode) { + String trajId = entry.has("trajectoryId") + ? entry.get("trajectoryId").asText() + : entry.path("trajectory_id").asText(); + JsonNode uNode = entry.has("usage") ? entry.get("usage") : null; + if (trajId != null && !trajId.isEmpty() && uNode != null) { + UsageMetadata parsed = parseUsageMetadata(uNode); + if (parsed != null) { + trajectoryUsages.put(trajId, parsed); + } + } + } + } + } + + if (payload.has("usageUpdate") || payload.has("usage_update")) { + JsonNode usageUpdate = payload.has("usageUpdate") + ? payload.get("usageUpdate") + : payload.get("usage_update"); + + if (usageUpdate.has("total")) { + UsageMetadata newTotal = parseUsageMetadata(usageUpdate.get("total")); + if (newTotal != null) { + this.cumulativeUsage = newTotal; + this.currentUsage = (turnStartUsage != null) ? newTotal.subtract(turnStartUsage) : newTotal; + } + } + + JsonNode agentsNode = usageUpdate.has("agents") ? usageUpdate.get("agents") : null; + if (agentsNode != null && agentsNode.isArray()) { + for (JsonNode entry : agentsNode) { + String trajId = entry.has("trajectoryId") + ? entry.get("trajectoryId").asText() + : entry.path("trajectory_id").asText(); + JsonNode uNode = entry.has("usage") ? entry.get("usage") : null; + if (trajId != null && !trajId.isEmpty() && uNode != null) { + UsageMetadata parsed = parseUsageMetadata(uNode); + if (parsed != null) { + trajectoryUsages.put(trajId, parsed); + } + } + } + } + } + + if (payload.has("usageMetadata") || payload.has("usage_metadata")) { + JsonNode topUsage = payload.has("usageMetadata") + ? payload.get("usageMetadata") + : payload.get("usage_metadata"); + UsageMetadata parsed = parseUsageMetadata(topUsage); + if (parsed != null) { + this.cumulativeUsage = parsed; + this.currentUsage = (turnStartUsage != null) ? parsed.subtract(turnStartUsage) : parsed; + } } if (payload.has("stepUpdate")) { @@ -1223,20 +1327,14 @@ void handleIncomingMessage(WebSocket webSocket, String message) { } } - if (stepUpdate.has("usageMetadata")) { - JsonNode usage = stepUpdate.get("usageMetadata"); - List promptDetails = parseModalityDetails(usage.path("promptTokensDetails")); - List cacheDetails = parseModalityDetails(usage.path("cacheTokensDetails")); - List candidateDetails = parseModalityDetails( - usage.path("candidatesTokensDetails")); - List toolUseDetails = parseModalityDetails( - usage.path("toolUsePromptTokensDetails")); - String serviceTier = usage.has("serviceTier") ? usage.get("serviceTier").asText() : null; - - currentUsage = new UsageMetadata(usage.path("promptTokenCount").asInt(), - usage.path("cachedContentTokenCount").asInt(), usage.path("candidatesTokenCount").asInt(), - usage.path("thoughtsTokenCount").asInt(), usage.path("totalTokenCount").asInt(), - serviceTier, promptDetails, cacheDetails, candidateDetails, toolUseDetails); + if (stepUpdate.has("usageMetadata") || stepUpdate.has("usage_metadata")) { + JsonNode usage = stepUpdate.has("usageMetadata") + ? stepUpdate.get("usageMetadata") + : stepUpdate.get("usage_metadata"); + UsageMetadata parsed = parseUsageMetadata(usage); + if (parsed != null) { + this.currentUsage = parsed; + } } if (stepUpdate.has("state") && "STATE_ERROR".equals(stepUpdate.path("state").asText())) { @@ -1532,9 +1630,18 @@ else if (req.has("customToolCall")) { currentToolCallsPublisher.closeExceptionally(new AgentCancelledException()); } } else { + UsageMetadata finalUsage = currentUsage; + if (finalUsage == null && cumulativeUsage != null) { + UsageMetadata diff = (turnStartUsage != null) + ? cumulativeUsage.subtract(turnStartUsage) + : cumulativeUsage; + if (diff.totalTokenCount() > 0) { + finalUsage = diff; + } + } currentChatFuture .complete(new AgentResponse(currentText != null ? currentText.toString() : "", - currentThoughts != null ? currentThoughts.toString() : "", currentUsage)); + currentThoughts != null ? currentThoughts.toString() : "", finalUsage)); if (currentThoughtsPublisher != null) { currentThoughtsPublisher.close(); } @@ -1741,6 +1848,62 @@ private static String resolveGeminiApiKey() { return propKey != null ? propKey : (envKey != null ? envKey : localEnvKey); } + /** + * Parses a {@link UsageMetadata} object from a Jackson {@link JsonNode}, + * supporting both camelCase and snake_case field variants. + * + * @param usage + * the JSON node representing UsageMetadata + * @return the parsed UsageMetadata record, or null if node is null/missing + */ + static UsageMetadata parseUsageMetadata(JsonNode usage) { + if (usage == null || usage.isMissingNode() || usage.isNull()) { + return null; + } + int promptTokenCount = getIntField(usage, "promptTokenCount", "prompt_token_count"); + int cachedContentTokenCount = getIntField(usage, "cachedContentTokenCount", "cached_content_token_count"); + int candidatesTokenCount = getIntField(usage, "candidatesTokenCount", "candidates_token_count"); + int thoughtsTokenCount = getIntField(usage, "thoughtsTokenCount", "thoughts_token_count"); + int totalTokenCount = getIntField(usage, "totalTokenCount", "total_token_count"); + + String serviceTier = null; + if (usage.has("serviceTier")) { + serviceTier = usage.get("serviceTier").asText(); + } else if (usage.has("service_tier")) { + serviceTier = usage.get("service_tier").asText(); + } + + JsonNode promptDetailsNode = usage.has("promptTokensDetails") + ? usage.get("promptTokensDetails") + : usage.get("prompt_tokens_details"); + JsonNode cacheDetailsNode = usage.has("cacheTokensDetails") + ? usage.get("cacheTokensDetails") + : usage.get("cache_tokens_details"); + JsonNode candidateDetailsNode = usage.has("candidatesTokensDetails") + ? usage.get("candidatesTokensDetails") + : usage.get("candidates_tokens_details"); + JsonNode toolUseDetailsNode = usage.has("toolUsePromptTokensDetails") + ? usage.get("toolUsePromptTokensDetails") + : usage.get("tool_use_prompt_tokens_details"); + + List promptDetails = parseModalityDetails(promptDetailsNode); + List cacheDetails = parseModalityDetails(cacheDetailsNode); + List candidateDetails = parseModalityDetails(candidateDetailsNode); + List toolUseDetails = parseModalityDetails(toolUseDetailsNode); + + return new UsageMetadata(promptTokenCount, cachedContentTokenCount, candidatesTokenCount, thoughtsTokenCount, + totalTokenCount, serviceTier, promptDetails, cacheDetails, candidateDetails, toolUseDetails); + } + + private static int getIntField(JsonNode node, String camelCaseName, String snakeCaseName) { + if (node.has(camelCaseName)) { + return node.get(camelCaseName).asInt(0); + } else if (node.has(snakeCaseName)) { + return node.get(snakeCaseName).asInt(0); + } + return 0; + } + /** * Parses a JSON array of modality token details into a list of * {@link ModalityTokenCount} records. @@ -1749,14 +1912,14 @@ private static String resolveGeminiApiKey() { * the JSON node containing the modality details array * @return an unmodifiable list of ModalityTokenCount records */ - private static List parseModalityDetails(JsonNode node) { + static List parseModalityDetails(JsonNode node) { if (node == null || !node.isArray()) { return List.of(); } List list = new ArrayList<>(); for (JsonNode item : node) { String modStr = item.path("modality").asText(""); - long count = item.path("tokenCount").asLong(); + long count = item.has("tokenCount") ? item.get("tokenCount").asLong(0) : item.path("token_count").asLong(0); Modality mod = Modality.fromString(modStr); list.add(new ModalityTokenCount(mod, count)); } diff --git a/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/AgentResponse.java b/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/AgentResponse.java index fe06d77..c52b28b 100644 --- a/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/AgentResponse.java +++ b/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/AgentResponse.java @@ -37,6 +37,15 @@ public record AgentResponse(String text, String thoughts, UsageMetadata usageMet thoughts = thoughts != null ? thoughts : ""; } + /** + * Convenience alias for {@link #usageMetadata()}. + * + * @return the usage metadata + */ + public UsageMetadata usage() { + return usageMetadata; + } + /** * Parses the text response as JSON and maps it to the specified class. * diff --git a/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/UsageMetadata.java b/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/UsageMetadata.java index ec309d3..d5821c0 100644 --- a/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/UsageMetadata.java +++ b/antigravity-sdk-wrapper/src/main/java/io/github/glaforge/antigravity/UsageMetadata.java @@ -74,9 +74,12 @@ public UsageMetadata(int promptTokenCount, int cachedContentTokenCount, int cand * @return a new UsageMetadata with summed token counts and merged service tier */ public UsageMetadata add(UsageMetadata other) { - if (other == null) { + if (other == null || other.totalTokenCount() == 0) { return this; } + if (this.totalTokenCount == 0) { + return other; + } String mergedTier; if (this.serviceTier == null) { mergedTier = other.serviceTier(); @@ -111,7 +114,7 @@ public UsageMetadata plus(UsageMetadata other) { * @return a new UsageMetadata with subtracted token counts */ public UsageMetadata subtract(UsageMetadata other) { - if (other == null) { + if (other == null || other.totalTokenCount() == 0) { return this; } String mergedTier = this.serviceTier != null ? this.serviceTier : other.serviceTier(); diff --git a/antigravity-sdk-wrapper/src/test/java/io/github/glaforge/antigravity/ObservabilityTest.java b/antigravity-sdk-wrapper/src/test/java/io/github/glaforge/antigravity/ObservabilityTest.java index 0b28cb7..d794eaa 100644 --- a/antigravity-sdk-wrapper/src/test/java/io/github/glaforge/antigravity/ObservabilityTest.java +++ b/antigravity-sdk-wrapper/src/test/java/io/github/glaforge/antigravity/ObservabilityTest.java @@ -38,14 +38,17 @@ public void testUsageObservability() throws Exception { await().atMost(90, TimeUnit.SECONDS).until(future::isDone); AgentResponse response = future.get(); - // Verify usage metadata is populated if available + // Verify usage metadata is populated UsageMetadata usage = response.usageMetadata(); - if (usage != null) { - System.out.println("Metadata: " + usage); - assertTrue(usage.promptTokenCount() >= 0, "Prompt tokens should be >= 0"); - } else { - System.out.println("Metadata was not returned by the harness in this test run."); - } + assertNotNull(usage, "UsageMetadata should not be null in response"); + assertEquals(usage, response.usage(), "response.usage() should match response.usageMetadata()"); + assertTrue(usage.promptTokenCount() > 0, "Prompt tokens should be > 0"); + assertTrue(usage.totalTokenCount() > 0, "Total tokens should be > 0"); + + // Verify Agent getters + assertNotNull(agent.getUsageMetadata(), "agent.getUsageMetadata() should not be null"); + assertNotNull(agent.getTotalUsage(), "agent.getTotalUsage() should not be null"); + assertTrue(agent.getTotalUsage().totalTokenCount() > 0, "Cumulative total tokens should be > 0"); } }); } diff --git a/antigravity-sdk-wrapper/src/test/java/io/github/glaforge/antigravity/UsageMetadataParsingTest.java b/antigravity-sdk-wrapper/src/test/java/io/github/glaforge/antigravity/UsageMetadataParsingTest.java new file mode 100644 index 0000000..26f7843 --- /dev/null +++ b/antigravity-sdk-wrapper/src/test/java/io/github/glaforge/antigravity/UsageMetadataParsingTest.java @@ -0,0 +1,115 @@ +/* + * 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 + * + * https://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 io.github.glaforge.antigravity; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.json.JsonMapper; +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; + +import java.util.List; + +import static org.junit.jupiter.api.Assertions.*; + +@Tag("unit") +public class UsageMetadataParsingTest { + + private final JsonMapper mapper = JsonMapper.builder().build(); + + @Test + public void testParseCamelCaseUsageMetadata() throws Exception { + String json = """ + { + "promptTokenCount": "1480", + "cachedContentTokenCount": "100", + "candidatesTokenCount": "25", + "thoughtsTokenCount": "30", + "totalTokenCount": "1635", + "serviceTier": "priority", + "promptTokensDetails": [ + {"modality": "TEXT", "tokenCount": "1480"} + ] + } + """; + JsonNode node = mapper.readTree(json); + UsageMetadata usage = Agent.parseUsageMetadata(node); + + assertNotNull(usage); + assertEquals(1480, usage.promptTokenCount()); + assertEquals(100, usage.cachedContentTokenCount()); + assertEquals(25, usage.candidatesTokenCount()); + assertEquals(30, usage.thoughtsTokenCount()); + assertEquals(1635, usage.totalTokenCount()); + assertEquals("priority", usage.serviceTier()); + assertEquals(1, usage.promptTokensDetails().size()); + assertEquals(Modality.TEXT, usage.promptTokensDetails().get(0).modality()); + assertEquals(1480, usage.promptTokensDetails().get(0).tokenCount()); + } + + @Test + public void testParseSnakeCaseUsageMetadata() throws Exception { + String json = """ + { + "prompt_token_count": 500, + "cached_content_token_count": 50, + "candidates_token_count": 100, + "thoughts_token_count": 40, + "total_token_count": 690, + "service_tier": "standard", + "prompt_tokens_details": [ + {"modality": "TEXT", "token_count": 500} + ] + } + """; + JsonNode node = mapper.readTree(json); + UsageMetadata usage = Agent.parseUsageMetadata(node); + + assertNotNull(usage); + assertEquals(500, usage.promptTokenCount()); + assertEquals(50, usage.cachedContentTokenCount()); + assertEquals(100, usage.candidatesTokenCount()); + assertEquals(40, usage.thoughtsTokenCount()); + assertEquals(690, usage.totalTokenCount()); + assertEquals("standard", usage.serviceTier()); + assertEquals(1, usage.promptTokensDetails().size()); + assertEquals(Modality.TEXT, usage.promptTokensDetails().get(0).modality()); + assertEquals(500, usage.promptTokensDetails().get(0).tokenCount()); + } + + @Test + public void testAgentResponseUsageAlias() { + UsageMetadata usage = new UsageMetadata(100, 0, 50, 10, 160); + AgentResponse response = new AgentResponse("Hello", "Thinking...", usage); + + assertEquals(usage, response.usageMetadata()); + assertEquals(usage, response.usage()); + } + + @Test + public void testUsageMetadataArithmeticWithZeroIdentity() { + UsageMetadata initial = new UsageMetadata(0, 0, 0, 0, 0); + UsageMetadata turn1 = new UsageMetadata(100, 10, 50, 20, 180, "priority", + List.of(new ModalityTokenCount(Modality.TEXT, 100)), List.of(), List.of(), List.of()); + + // Subtraction against initial zero usage must preserve details + UsageMetadata diff = turn1.subtract(initial); + assertSame(turn1, diff); + + // Addition with zero usage must preserve details + UsageMetadata added = initial.add(turn1); + assertSame(turn1, added); + } +} diff --git a/skills/antigravity-sdk-java/references/api-reference.md b/skills/antigravity-sdk-java/references/api-reference.md index 50ef299..6b61030 100644 --- a/skills/antigravity-sdk-java/references/api-reference.md +++ b/skills/antigravity-sdk-java/references/api-reference.md @@ -583,6 +583,10 @@ if (usage != null) { System.out.println("Prompt modality: " + detail.modality() + " -> " + detail.tokenCount() + " tokens"); } } + +// Session cumulative usage and per-trajectory usage (for subagents) +UsageMetadata totalUsage = agent.getTotalUsage(); +Map trajectoryUsages = agent.getTrajectoryUsages(); ``` ---