Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions go/adk/pkg/a2a/executor.go
Original file line number Diff line number Diff line change
Expand Up @@ -270,7 +270,11 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.RequestCont
lastNonPartialParts a2atype.ContentParts
hitlParts a2atype.ContentParts
runErr error
usage turnUsage
)
// Resumed tasks (HITL cycles, follow-up messages) carry the previously
// persisted total; seed it so kagent_usage_total stays a task-lifetime sum.
usage.seedFromTask(reqCtx.StoredTask)

for adkEvent, adkErr := range r.Run(ctx, userID, sessionID, content, runConfig) {
if adkErr != nil {
Expand All @@ -287,6 +291,9 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.RequestCont
invocationSpan.SetAttributes(attribute.String("gcp.vertex.agent.invocation_id", invocationID))
}

// Aggregate token usage for the terminal status update.
usage.add(adkEvent)

// Build per-event metadata (inherits baseMeta + adds invocation_id, usage etc.).
eventMeta := buildEventMeta(baseMeta, adkEvent)

Expand All @@ -299,6 +306,7 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.RequestCont
a2atype.TextPart{Text: fmt.Sprintf("LLM error: %s %s", adkEvent.ErrorCode, adkEvent.ErrorMessage)})
failed := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateFailed, errMsg)
failed.Final = true
usage.stamp(eventMeta)
failed.Metadata = eventMeta
return queue.Write(ctx, failed)
}
Expand All @@ -311,6 +319,7 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.RequestCont
a2atype.TextPart{Text: fmt.Sprintf("LLM error: %s %s", adkEvent.ErrorCode, adkEvent.ErrorMessage)})
failed := a2atype.NewStatusUpdateEvent(reqCtx, a2atype.TaskStateFailed, errMsg)
failed.Final = true
usage.stamp(eventMeta)
failed.Metadata = eventMeta
return queue.Write(ctx, failed)
}
Expand Down Expand Up @@ -383,6 +392,7 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.RequestCont
if invocationID != "" {
finalMeta[adka2a.ToA2AMetaKey("invocation_id")] = invocationID
}
usage.stamp(finalMeta)

if runErr != nil {
errMsg := newAgentMessage(reqCtx, a2atype.TextPart{Text: runErr.Error()})
Expand Down
110 changes: 110 additions & 0 deletions go/adk/pkg/a2a/usage.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,110 @@
package a2a

import (
"math"

a2atype "github.com/a2aproject/a2a-go/a2a"
adksession "google.golang.org/adk/v2/session"
"google.golang.org/genai"
)

// turnUsage accumulates token usage across the ADK events of one execution so
// the aggregated total can be emitted on terminal status updates. Partial
// (streaming chunk) events are skipped: each LLM call reports its usage on the
// final non-partial event, so summing partials would double-count.
type turnUsage struct {
promptTokens int64
completionTokens int64
thoughtsTokens int64
cachedContentTokens int64
totalTokens int64
modelVersion string
}

func (u *turnUsage) add(event *adksession.Event) {
if event == nil || event.Partial || event.UsageMetadata == nil {
return
}
u.promptTokens += int64(event.UsageMetadata.PromptTokenCount)
u.completionTokens += int64(event.UsageMetadata.CandidatesTokenCount)
u.thoughtsTokens += int64(event.UsageMetadata.ThoughtsTokenCount)
u.cachedContentTokens += int64(event.UsageMetadata.CachedContentTokenCount)
u.totalTokens += int64(event.UsageMetadata.TotalTokenCount)
if event.ModelVersion != "" {
u.modelVersion = event.ModelVersion
}
}

// seedFromTask primes the accumulator with the total already persisted on a
// resumed task, so tasks spanning multiple executions (HITL input-required
// cycles, follow-up messages) report a task-lifetime total instead of the last
// segment only.
func (u *turnUsage) seedFromTask(task *a2atype.Task) {
if task == nil || task.Metadata == nil {
return
}
prior, ok := task.Metadata[GetKAgentMetadataKey("usage_total")].(map[string]any)
if !ok {
return
}
u.promptTokens += metadataTokenCount(prior["promptTokenCount"])
u.completionTokens += metadataTokenCount(prior["candidatesTokenCount"])
u.thoughtsTokens += metadataTokenCount(prior["thoughtsTokenCount"])
u.cachedContentTokens += metadataTokenCount(prior["cachedContentTokenCount"])
u.totalTokens += metadataTokenCount(prior["totalTokenCount"])
if modelVersion, ok := prior["modelVersion"].(string); ok && modelVersion != "" {
u.modelVersion = modelVersion
}
}

// metadataTokenCount reads a numeric token count from stored task metadata.
// Counts are float64 after a JSON round-trip but keep an integer type with
// in-memory task stores.
func metadataTokenCount(value any) int64 {
switch n := value.(type) {
case float64:
return int64(n)
case int64:
return n
case int32:
return int64(n)
case int:
return int64(n)
default:
return 0
}
}

func (u *turnUsage) empty() bool {
return u.promptTokens == 0 && u.completionTokens == 0 && u.totalTokens == 0
}

// stamp attaches the aggregate to meta under kagent_usage_total. The value is
// serialized exactly like the per-event adk_usage_metadata (same genai type,
// same JSON mapping) plus modelVersion, so consumers can share one parser.
func (u *turnUsage) stamp(meta map[string]any) {
if u.empty() {
return
}
total, err := toA2AMetadataMap(&genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: clampInt32(u.promptTokens),
CandidatesTokenCount: clampInt32(u.completionTokens),
ThoughtsTokenCount: clampInt32(u.thoughtsTokens),
CachedContentTokenCount: clampInt32(u.cachedContentTokens),
TotalTokenCount: clampInt32(u.totalTokens),
})
if err != nil || total == nil {
return
}
if u.modelVersion != "" {
total["modelVersion"] = u.modelVersion
}
meta[GetKAgentMetadataKey("usage_total")] = total
}

func clampInt32(v int64) int32 {
if v > math.MaxInt32 {
return math.MaxInt32
}
return int32(v)
}
122 changes: 122 additions & 0 deletions go/adk/pkg/a2a/usage_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,122 @@
package a2a

import (
"math"
"testing"

a2atype "github.com/a2aproject/a2a-go/a2a"
"github.com/stretchr/testify/require"
"google.golang.org/adk/v2/model"
"google.golang.org/adk/v2/server/adka2a" //nolint:staticcheck // kagent still uses a2a-go v1; this ADK package is the compatibility adapter.
adksession "google.golang.org/adk/v2/session"
"google.golang.org/genai"
)

func usageEvent(prompt, completion, total int32, modelVersion string, partial bool) *adksession.Event {
return &adksession.Event{
LLMResponse: model.LLMResponse{
UsageMetadata: &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: prompt,
CandidatesTokenCount: completion,
TotalTokenCount: total,
},
ModelVersion: modelVersion,
Partial: partial,
},
}
}

func stampedTotal(t *testing.T, usage *turnUsage) map[string]any {
t.Helper()
meta := map[string]any{}
usage.stamp(meta)
total, ok := meta[GetKAgentMetadataKey("usage_total")].(map[string]any)
require.True(t, ok, "kagent_usage_total must be stamped")
return total
}

func TestTurnUsageAggregatesNonPartialEvents(t *testing.T) {
var usage turnUsage
require.True(t, usage.empty())

usage.add(usageEvent(100, 20, 120, "model-a", false))
usage.add(usageEvent(200, 30, 230, "model-b", false))

require.False(t, usage.empty())
require.Equal(t, map[string]any{
"promptTokenCount": float64(300),
"candidatesTokenCount": float64(50),
"totalTokenCount": float64(350),
"modelVersion": "model-b",
}, stampedTotal(t, &usage))
}

func TestTurnUsageSkipsPartialAndEmptyEvents(t *testing.T) {
var usage turnUsage

usage.add(nil)
usage.add(&adksession.Event{})
usage.add(usageEvent(999, 999, 999, "chunk-model", true))

require.True(t, usage.empty())
meta := map[string]any{}
usage.stamp(meta)
require.NotContains(t, meta, GetKAgentMetadataKey("usage_total"),
"empty usage must not stamp the key")

usage.add(usageEvent(10, 5, 15, "", false))
require.Equal(t, map[string]any{
"promptTokenCount": float64(10),
"candidatesTokenCount": float64(5),
"totalTokenCount": float64(15),
}, stampedTotal(t, &usage), "modelVersion must be omitted when no event carried one")
}

func TestTurnUsageShapeMatchesPerEventUsageMetadata(t *testing.T) {
event := usageEvent(10, 5, 15, "", false)

var usage turnUsage
usage.add(event)

perEventMeta := buildEventMeta(map[string]any{}, event)
require.Equal(t,
perEventMeta[adka2a.ToA2AMetaKey("usage_metadata")],
stampedTotal(t, &usage),
"kagent_usage_total must serialize identically to adk_usage_metadata")
}

func TestTurnUsageSeedFromTaskAccumulatesAcrossExecutions(t *testing.T) {
var usage turnUsage
// float64 values mimic a JSON round-trip through the task store.
usage.seedFromTask(&a2atype.Task{Metadata: map[string]any{
GetKAgentMetadataKey("usage_total"): map[string]any{
"promptTokenCount": float64(100),
"candidatesTokenCount": float64(20),
"totalTokenCount": float64(120),
"modelVersion": "model-a",
},
}})
usage.add(usageEvent(200, 30, 230, "", false))

require.Equal(t, map[string]any{
"promptTokenCount": float64(300),
"candidatesTokenCount": float64(50),
"totalTokenCount": float64(350),
"modelVersion": "model-a",
}, stampedTotal(t, &usage))
}

func TestTurnUsageSeedFromTaskIgnoresMissingOrMalformed(t *testing.T) {
var usage turnUsage
usage.seedFromTask(nil)
usage.seedFromTask(&a2atype.Task{})
usage.seedFromTask(&a2atype.Task{Metadata: map[string]any{
GetKAgentMetadataKey("usage_total"): "not-a-map",
}})
require.True(t, usage.empty())
}

func TestClampInt32(t *testing.T) {
require.Equal(t, int32(42), clampInt32(42))
require.Equal(t, int32(math.MaxInt32), clampInt32(math.MaxInt32+1))
}
Loading
Loading