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
13 changes: 8 additions & 5 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 @@ -341,8 +342,9 @@ func runAnthropicStreaming(ctx context.Context, m *AnthropicModel, params anthro
var usage *genai.GenerateContentResponseUsageMetadata
if inputTokens > 0 || outputTokens > 0 {
usage = &genai.GenerateContentResponseUsageMetadata{
PromptTokenCount: int32(inputTokens),
CandidatesTokenCount: int32(outputTokens),
PromptTokenCount: int32(inputTokens),
CandidatesTokenCount: int32(outputTokens),
CachedContentTokenCount: int32(cacheReadInputTokens),
}
}
resp := &model.LLMResponse{
Expand Down Expand Up @@ -388,8 +390,9 @@ func runAnthropicNonStreaming(ctx context.Context, m *AnthropicModel, params ant
var usage *genai.GenerateContentResponseUsageMetadata
if message.Usage.InputTokens > 0 || message.Usage.OutputTokens > 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: int32(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"
"net/http"
"net/http/httptest"
"testing"

"github.com/anthropics/anthropic-sdk-go"
"github.com/anthropics/anthropic-sdk-go/option"
"github.com/go-logr/logr"
"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: logr.Discard(),
}

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: 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: 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: int32(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: int32(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":11}
}
}`)

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 != 11 {
t.Fatalf("cachedContentTokenCount = %d, want 11", resp.UsageMetadata.CachedContentTokenCount)
}
}

func TestExtractThoughtSignatureFromStreamingToolCallChunk(t *testing.T) {
Expand Down
9 changes: 7 additions & 2 deletions go/adk/pkg/models/openai_responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -342,9 +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),
PromptTokenCount: int32(u.InputTokens),
CandidatesTokenCount: int32(u.OutputTokens),
CachedContentTokenCount: int32(u.InputTokensDetails.CachedTokens),
}
}

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 @@ -127,7 +127,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":42},"output_tokens_details":{"reasoning_tokens":0}}
}`)
var resp responses.Response
if err := json.Unmarshal(raw, &resp); err != nil {
Expand All @@ -146,6 +146,9 @@ func TestResponseToLLMResponse(t *testing.T) {
if out.UsageMetadata == nil || out.UsageMetadata.PromptTokenCount != 10 {
t.Fatalf("usage = %#v", out.UsageMetadata)
}
if out.UsageMetadata.CachedContentTokenCount != 42 {
t.Fatalf("cachedContentTokenCount = %d, want 42", out.UsageMetadata.CachedContentTokenCount)
}
}

func TestOpenAIModel_GenerateContent_Responses(t *testing.T) {
Expand All @@ -160,7 +163,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":7},"output_tokens_details":{"reasoning_tokens":0}}
}`))
}))
defer srv.Close()
Expand Down Expand Up @@ -194,6 +197,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 == nil || got.UsageMetadata == nil || got.UsageMetadata.CachedContentTokenCount != 7 {
t.Fatalf("cachedContentTokenCount = %#v, want 7", got.UsageMetadata)
}
}

func TestOpenAIModel_GenerateContent_ResponsesStreaming(t *testing.T) {
Expand All @@ -208,7 +214,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":5},"output_tokens_details":{"reasoning_tokens":0}}}}`)
_, _ = io.WriteString(w, "data: [DONE]\n\n")
}))
defer srv.Close()
Expand Down Expand Up @@ -244,6 +250,9 @@ func TestOpenAIModel_GenerateContent_ResponsesStreaming(t *testing.T) {
if final == nil || final.Content.Parts[0].Text != "hello" {
t.Fatalf("final = %#v", final)
}
if final == nil || final.UsageMetadata == nil || final.UsageMetadata.CachedContentTokenCount != 5 {
t.Fatalf("cachedContentTokenCount = %#v, want 5", final.UsageMetadata)
}
}

func TestGenaiContentsToResponsesInput_Image(t *testing.T) {
Expand Down
13 changes: 13 additions & 0 deletions python/packages/kagent-adk/src/kagent/adk/models/_openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -280,6 +280,17 @@ def _convert_tools_to_openai(tools: list[types.Tool]) -> list[ChatCompletionTool
return openai_tools


def _cached_prompt_tokens(usage: Any) -> int:
"""Return OpenAI cached prompt-token count (0 when absent).

OpenAI reports prompt-cache hits via usage.prompt_tokens_details.cached_tokens.
"""
details = getattr(usage, "prompt_tokens_details", None)
if details is None:
return 0
return getattr(details, "cached_tokens", 0) or 0


def _convert_openai_response_to_llm_response(response: ChatCompletion) -> LlmResponse:
"""Convert OpenAI response to LlmResponse."""
choice = response.choices[0]
Expand Down Expand Up @@ -319,6 +330,7 @@ def _convert_openai_response_to_llm_response(response: ChatCompletion) -> LlmRes
prompt_token_count=response.usage.prompt_tokens,
candidates_token_count=response.usage.completion_tokens,
total_token_count=response.usage.total_tokens,
cached_content_token_count=_cached_prompt_tokens(response.usage),
)

# Handle finish reason
Expand Down Expand Up @@ -522,6 +534,7 @@ async def generate_content_async(
prompt_token_count=chunk.usage.prompt_tokens,
candidates_token_count=chunk.usage.completion_tokens,
total_token_count=chunk.usage.total_tokens,
cached_content_token_count=_cached_prompt_tokens(chunk.usage),
)

# Yield final aggregated response with partial=False
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1061,3 +1061,31 @@ def test_round_trip_preserves_thought_signature_for_follow_up_tool_result(self):
tool_messages = [m for m in messages if m["role"] == "tool"]
assert len(tool_messages) == 1
assert tool_messages[0]["extra_content"] == {"google": {"thought_signature": "YWJj"}}


def test_usage_metadata_populates_cached_content_token_count(self):
# Openai reports prompt-cache hits via usage.prompt_tokens_details.cached_tokens.
class _MockPromptTokensDetails:
cached_tokens = 12

response = self._MockResponse(self._MockMessage(content="hi"))
response.usage.prompt_tokens_details = _MockPromptTokensDetails()

llm_response = _convert_openai_response_to_llm_response(response)

assert llm_response.usage_metadata is not None, "usage metadata should be populated"
assert llm_response.usage_metadata.cached_content_token_count == 12, (
"cached_content_token_count should equal cached prompt tokens",
llm_response.usage_metadata
)

def test_opens_metadata_cached_content_token_zero_when_absent(self):
response = self._MockResponse(self._MockMessage("hi"))

llm_response = _convert_openai_response_to_llm_response(response)

assert llm_response.usage_metadata is not None
assert llm_response.usage_metadata.cached_content_token_count == 0, (
"cached_content_token_count should default to 0 when provider omits it",
llm_response.usage_metadata.cached_content_token_count,
)
Loading