Skip to content
Open
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
20 changes: 13 additions & 7 deletions go/adk/pkg/models/anthropic_adk.go
Original file line number Diff line number Diff line change
Expand Up @@ -264,14 +264,15 @@ func runAnthropicStreaming(ctx context.Context, m *AnthropicModel, params anthro
inputJSON string
})
var stopReason anthropic.StopReason
var inputTokens, outputTokens int64
var inputTokens, outputTokens, cacheReadInputTokens int64

for stream.Next() {
event := stream.Current()

switch e := event.AsAny().(type) {
case anthropic.MessageStartEvent:
inputTokens = e.Message.Usage.InputTokens
cacheReadInputTokens = e.Message.Usage.CacheReadInputTokens
case anthropic.ContentBlockStartEvent:
idx := int(e.Index)
if e.ContentBlock.Type == "tool_use" {
Expand Down Expand Up @@ -309,6 +310,9 @@ func runAnthropicStreaming(ctx context.Context, m *AnthropicModel, params anthro
case anthropic.MessageDeltaEvent:
stopReason = e.Delta.StopReason
outputTokens = e.Usage.OutputTokens
if e.Usage.JSON.CacheReadInputTokens.Valid() {
cacheReadInputTokens = e.Usage.CacheReadInputTokens
}
}
}

Expand Down Expand Up @@ -339,10 +343,11 @@ func runAnthropicStreaming(ctx context.Context, m *AnthropicModel, params anthro
}

var usage *genai.GenerateContentResponseUsageMetadata
if inputTokens > 0 || outputTokens > 0 {
if inputTokens > 0 || outputTokens > 0 || cacheReadInputTokens > 0 {
usage = &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: int32(inputTokens),
CandidatesTokenCount: int32(outputTokens),
PromptTokenCount: int32(inputTokens),
CandidatesTokenCount: int32(outputTokens),
CachedContentTokenCount: cachedTokenCount(cacheReadInputTokens),
}
}
resp := &model.LLMResponse{
Expand Down Expand Up @@ -386,10 +391,11 @@ func runAnthropicNonStreaming(ctx context.Context, m *AnthropicModel, params ant

// Build usage metadata
var usage *genai.GenerateContentResponseUsageMetadata
if message.Usage.InputTokens > 0 || message.Usage.OutputTokens > 0 {
if message.Usage.InputTokens > 0 || message.Usage.OutputTokens > 0 || message.Usage.CacheReadInputTokens > 0 {
usage = &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: int32(message.Usage.InputTokens),
CandidatesTokenCount: int32(message.Usage.OutputTokens),
PromptTokenCount: int32(message.Usage.InputTokens),
CandidatesTokenCount: int32(message.Usage.OutputTokens),
CachedContentTokenCount: cachedTokenCount(message.Usage.CacheReadInputTokens),
}
}

Expand Down
62 changes: 62 additions & 0 deletions go/adk/pkg/models/anthropic_adk_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
package models

import (
"context"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"testing"

"github.com/anthropics/anthropic-sdk-go"
"github.com/anthropics/anthropic-sdk-go/option"
"google.golang.org/adk/v2/model"
"google.golang.org/genai"
)

// messageResponse is the anthropic Messages API payload served by the mock server.
const anthropicMessageResponse = `{
"id":"msg_01","type":"message","role":"assistant","model":"claude-sonnet-4-20250514",
"content":[{"type":"text","text":"pong"}],
"stop_reason":"end_turn","stop_sequence":null,
"usage":{"input_tokens":10,"output_tokens":5,"cache_read_input_tokens":8,"cache_creation_input_tokens":0}
}`

// TestAnthropicNonStreamingCachedTokens verifies CacheReadInputTokens flows into
// GenerateContentResponseUsageMetadata.CachedContentTokenCount for non-streaming calls.
func TestAnthropicNonStreamingCachedTokens(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, anthropicMessageResponse)
}))
defer srv.Close()

client := anthropic.NewClient(
option.WithAPIKey("test"),
option.WithBaseURL(srv.URL),
)
m := &AnthropicModel{
Config: &AnthropicConfig{Model: "claude-sonnet-4-20250514"},
Client: client,
Logger: slog.New(slog.DiscardHandler),
}

var got *model.LLMResponse
for resp, err := range m.GenerateContent(context.Background(), &model.LLMRequest{
Contents: []*genai.Content{{Role: "user", Parts: []*genai.Part{{Text: "ping"}}}},
}, false) {
if err != nil {
t.Fatalf("GenerateContent error: %v", err)
}
got = resp
}
if got == nil || got.UsageMetadata == nil {
t.Fatalf("usage metadata = %#v", got)
}
if got.UsageMetadata.PromptTokenCount != 10 {
t.Fatalf("PromptTokenCount = %d, want 10", got.UsageMetadata.PromptTokenCount)
}
if got.UsageMetadata.CachedContentTokenCount != 8 {
t.Fatalf("CachedContentTokenCount = %d, want 8", got.UsageMetadata.CachedContentTokenCount)
}
}
14 changes: 8 additions & 6 deletions go/adk/pkg/models/bedrock.go
Original file line number Diff line number Diff line change
Expand Up @@ -409,9 +409,10 @@ func (m *BedrockModel) generateStreaming(ctx context.Context, modelId string, me
if meta, ok := event.(*types.ConverseStreamOutputMemberMetadata); ok {
if meta.Value.Usage != nil {
usageMetadata = &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: aws.ToInt32(meta.Value.Usage.InputTokens),
CandidatesTokenCount: aws.ToInt32(meta.Value.Usage.OutputTokens),
TotalTokenCount: aws.ToInt32(meta.Value.Usage.TotalTokens),
PromptTokenCount: aws.ToInt32(meta.Value.Usage.InputTokens),
CandidatesTokenCount: aws.ToInt32(meta.Value.Usage.OutputTokens),
TotalTokenCount: aws.ToInt32(meta.Value.Usage.TotalTokens),
CachedContentTokenCount: cachedTokenCount(int64(aws.ToInt32(meta.Value.Usage.CacheReadInputTokens))),
}
}
}
Expand Down Expand Up @@ -537,9 +538,10 @@ func (m *BedrockModel) generateNonStreaming(ctx context.Context, modelId string,
var usageMetadata *genai.GenerateContentResponseUsageMetadata
if output.Usage != nil {
usageMetadata = &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: aws.ToInt32(output.Usage.InputTokens),
CandidatesTokenCount: aws.ToInt32(output.Usage.OutputTokens),
TotalTokenCount: aws.ToInt32(output.Usage.TotalTokens),
PromptTokenCount: aws.ToInt32(output.Usage.InputTokens),
CandidatesTokenCount: aws.ToInt32(output.Usage.OutputTokens),
TotalTokenCount: aws.ToInt32(output.Usage.TotalTokens),
CachedContentTokenCount: cachedTokenCount(int64(aws.ToInt32(output.Usage.CacheReadInputTokens))),
}
}

Expand Down
17 changes: 10 additions & 7 deletions go/adk/pkg/models/openai_adk.go
Original file line number Diff line number Diff line change
Expand Up @@ -378,14 +378,15 @@ func runStreaming(ctx context.Context, m *OpenAIModel, params openai.ChatComplet
var aggregatedText strings.Builder
toolCallsAcc := make(map[int64]map[string]any)
var finishReason string
var promptTokens, completionTokens, totalTokens int64
var promptTokens, completionTokens, totalTokens, cachedTokens int64

for stream.Next() {
chunk := stream.Current()
if chunk.Usage.PromptTokens > 0 || chunk.Usage.CompletionTokens > 0 {
promptTokens = chunk.Usage.PromptTokens
completionTokens = chunk.Usage.CompletionTokens
totalTokens = chunk.Usage.TotalTokens
cachedTokens = chunk.Usage.PromptTokensDetails.CachedTokens
}
if len(chunk.Choices) == 0 {
continue
Expand Down Expand Up @@ -465,9 +466,10 @@ func runStreaming(ctx context.Context, m *OpenAIModel, params openai.ChatComplet
var usage *genai.GenerateContentResponseUsageMetadata
if promptTokens > 0 || completionTokens > 0 {
usage = &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: int32(promptTokens),
CandidatesTokenCount: int32(completionTokens),
TotalTokenCount: int32(totalTokens),
PromptTokenCount: int32(promptTokens),
CandidatesTokenCount: int32(completionTokens),
TotalTokenCount: int32(totalTokens),
CachedContentTokenCount: cachedTokenCount(cachedTokens),
}
}
resp := &model.LLMResponse{
Expand Down Expand Up @@ -527,9 +529,10 @@ func chatCompletionToLLMResponse(completion *openai.ChatCompletion) *model.LLMRe
var usage *genai.GenerateContentResponseUsageMetadata
if completion.Usage.PromptTokens > 0 || completion.Usage.CompletionTokens > 0 {
usage = &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: int32(completion.Usage.PromptTokens),
CandidatesTokenCount: int32(completion.Usage.CompletionTokens),
TotalTokenCount: int32(completion.Usage.TotalTokens),
PromptTokenCount: int32(completion.Usage.PromptTokens),
CandidatesTokenCount: int32(completion.Usage.CompletionTokens),
TotalTokenCount: int32(completion.Usage.TotalTokens),
CachedContentTokenCount: cachedTokenCount(completion.Usage.PromptTokensDetails.CachedTokens),
}
}
return &model.LLMResponse{
Expand Down
6 changes: 5 additions & 1 deletion go/adk/pkg/models/openai_adk_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -387,7 +387,8 @@ func TestChatCompletionToLLMResponse_PreservesThoughtSignature(t *testing.T) {
"usage":{
"prompt_tokens":3,
"completion_tokens":4,
"total_tokens":7
"total_tokens":7,
"prompt_tokens_details":{"cached_tokens":2}
}
}`)

Expand All @@ -414,6 +415,9 @@ func TestChatCompletionToLLMResponse_PreservesThoughtSignature(t *testing.T) {
if resp.UsageMetadata == nil || resp.UsageMetadata.PromptTokenCount != 3 || resp.UsageMetadata.CandidatesTokenCount != 4 {
t.Fatalf("usage metadata = %#v, want prompt=3 completion=4", resp.UsageMetadata)
}
if resp.UsageMetadata.CachedContentTokenCount != 2 {
t.Fatalf("cachedContentTokenCount = %d, want 2", resp.UsageMetadata.CachedContentTokenCount)
}
}

func TestExtractThoughtSignatureFromStreamingToolCallChunk(t *testing.T) {
Expand Down
6 changes: 5 additions & 1 deletion go/adk/pkg/models/openai_responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -342,10 +342,14 @@ func responsesUsageToGenai(u responses.ResponseUsage) *genai.GenerateContentResp
if u.InputTokens == 0 && u.OutputTokens == 0 {
return nil
}
// CachedContentTokenCount flows into A2A task usage and the llm_response trace
// attribute via GenerateContentResponseUsageMetadata. Emitting a dedicated
// gen_ai.client.token.usage observation (gen_ai.token.type="cached") is a
// follow-up; see #2669.
return &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: int32(u.InputTokens),
CandidatesTokenCount: int32(u.OutputTokens),
CachedContentTokenCount: int32(u.InputTokensDetails.CachedTokens),
CachedContentTokenCount: cachedTokenCount(u.InputTokensDetails.CachedTokens),
ThoughtsTokenCount: int32(u.OutputTokensDetails.ReasoningTokens),
}
}
Expand Down
15 changes: 12 additions & 3 deletions go/adk/pkg/models/openai_responses_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -128,7 +128,7 @@ func TestResponseToLLMResponse(t *testing.T) {
"status":"completed"
}
],
"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15,"input_tokens_details":{"cached_tokens":0},"output_tokens_details":{"reasoning_tokens":0}}
"usage":{"input_tokens":10,"output_tokens":5,"total_tokens":15,"input_tokens_details":{"cached_tokens":8},"output_tokens_details":{"reasoning_tokens":0}}
}`)
var resp responses.Response
if err := json.Unmarshal(raw, &resp); err != nil {
Expand All @@ -147,6 +147,9 @@ func TestResponseToLLMResponse(t *testing.T) {
if out.UsageMetadata == nil || out.UsageMetadata.PromptTokenCount != 10 {
t.Fatalf("usage = %#v", out.UsageMetadata)
}
if out.UsageMetadata.CachedContentTokenCount != 8 {
t.Fatalf("cachedContentTokenCount = %d, want 8", out.UsageMetadata.CachedContentTokenCount)
}
}

func TestOpenAIModel_GenerateContent_Responses(t *testing.T) {
Expand All @@ -161,7 +164,7 @@ func TestOpenAIModel_GenerateContent_Responses(t *testing.T) {
"id":"resp_1","object":"response","created_at":1,"status":"completed","model":"gpt-4o",
"output":[{"type":"message","id":"msg_1","role":"assistant","status":"completed",
"content":[{"type":"output_text","text":"pong","annotations":[]}]}],
"usage":{"input_tokens":3,"output_tokens":1,"total_tokens":4,"input_tokens_details":{"cached_tokens":0},"output_tokens_details":{"reasoning_tokens":0}}
"usage":{"input_tokens":3,"output_tokens":1,"total_tokens":4,"input_tokens_details":{"cached_tokens":2},"output_tokens_details":{"reasoning_tokens":0}}
}`))
}))
defer srv.Close()
Expand Down Expand Up @@ -195,6 +198,9 @@ func TestOpenAIModel_GenerateContent_Responses(t *testing.T) {
if got == nil || got.Content == nil || len(got.Content.Parts) != 1 || got.Content.Parts[0].Text != "pong" {
t.Fatalf("response = %#v", got)
}
if got.UsageMetadata == nil || got.UsageMetadata.CachedContentTokenCount != 2 {
t.Fatalf("cachedContentTokenCount = %#v, want 2", got.UsageMetadata)
}
}

func TestOpenAIModel_GenerateContent_ResponsesStreaming(t *testing.T) {
Expand All @@ -209,7 +215,7 @@ func TestOpenAIModel_GenerateContent_ResponsesStreaming(t *testing.T) {
}
write(`{"type":"response.output_text.delta","content_index":0,"delta":"hel","item_id":"msg_1","output_index":0,"sequence_number":1,"logprobs":[]}`)
write(`{"type":"response.output_text.delta","content_index":0,"delta":"lo","item_id":"msg_1","output_index":0,"sequence_number":2,"logprobs":[]}`)
write(`{"type":"response.completed","sequence_number":3,"response":{"id":"resp_1","object":"response","created_at":1,"status":"completed","model":"gpt-4o","output":[],"usage":{"input_tokens":1,"output_tokens":2,"total_tokens":3,"input_tokens_details":{"cached_tokens":0},"output_tokens_details":{"reasoning_tokens":0}}}}`)
write(`{"type":"response.completed","sequence_number":3,"response":{"id":"resp_1","object":"response","created_at":1,"status":"completed","model":"gpt-4o","output":[],"usage":{"input_tokens":1,"output_tokens":2,"total_tokens":3,"input_tokens_details":{"cached_tokens":1},"output_tokens_details":{"reasoning_tokens":0}}}}`)
_, _ = io.WriteString(w, "data: [DONE]\n\n")
}))
defer srv.Close()
Expand Down Expand Up @@ -245,6 +251,9 @@ func TestOpenAIModel_GenerateContent_ResponsesStreaming(t *testing.T) {
if final == nil || final.Content.Parts[0].Text != "hello" {
t.Fatalf("final = %#v", final)
}
if final.UsageMetadata == nil || final.UsageMetadata.CachedContentTokenCount != 1 {
t.Fatalf("cachedContentTokenCount = %#v, want 1", final.UsageMetadata)
}
}

func TestGenaiContentsToResponsesInput_Image(t *testing.T) {
Expand Down
7 changes: 7 additions & 0 deletions go/adk/pkg/models/usage.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
package models

import "math"

func cachedTokenCount(tokens int64) int32 {
return int32(min(max(tokens, 0), math.MaxInt32))
}
Loading