From f79b68a0dbd6f02dc48604c2957aa3061aee433e Mon Sep 17 00:00:00 2001 From: Dmitry Ng <19asdek91@gmail.com> Date: Wed, 28 Jan 2026 18:05:51 +0300 Subject: [PATCH] llms/bedrock: standardize token field names for cross-provider compatibility - Updated token generation information across multiple provider implementations (AI21, Amazon, Anthropic, Meta, Nova) to include standardized field names: `PromptTokens`, `CompletionTokens`, and `TotalTokens`. - This change enhances consistency in token reporting across different LLM providers, facilitating easier integration and comparison of results. - Adjusted parsing logic to ensure new fields are populated correctly during streaming responses. --- .../internal/bedrockclient/provider_ai21.go | 9 +++++++++ .../internal/bedrockclient/provider_amazon.go | 9 +++++++++ .../bedrockclient/provider_anthropic.go | 17 ++++++++++++++++- .../internal/bedrockclient/provider_meta.go | 9 +++++++++ .../internal/bedrockclient/provider_nova.go | 9 +++++++++ 5 files changed, 52 insertions(+), 1 deletion(-) diff --git a/llms/bedrock/internal/bedrockclient/provider_ai21.go b/llms/bedrock/internal/bedrockclient/provider_ai21.go index 77909457..d130995a 100644 --- a/llms/bedrock/internal/bedrockclient/provider_ai21.go +++ b/llms/bedrock/internal/bedrockclient/provider_ai21.go @@ -220,6 +220,10 @@ func createAi21Completion(ctx context.Context, client *bedrockruntime.Client, mo "id": output.ID, "input_tokens": int32(len(output.Prompt.Tokens)), "output_tokens": int32(len(completion.Data.Tokens)), + // Standardized field names for cross-provider compatibility + "PromptTokens": int32(len(output.Prompt.Tokens)), + "CompletionTokens": int32(len(completion.Data.Tokens)), + "TotalTokens": int32(len(output.Prompt.Tokens)) + int32(len(completion.Data.Tokens)), }, } } @@ -341,9 +345,14 @@ func parseAi21StreamingResponse(ctx context.Context, client *bedrockruntime.Clie // Set token counts if available if resp.Usage.PromptTokens > 0 { contentchoices[0].GenerationInfo["input_tokens"] = resp.Usage.PromptTokens + contentchoices[0].GenerationInfo["PromptTokens"] = resp.Usage.PromptTokens } if resp.Usage.CompletionTokens > 0 { contentchoices[0].GenerationInfo["output_tokens"] = resp.Usage.CompletionTokens + contentchoices[0].GenerationInfo["CompletionTokens"] = resp.Usage.CompletionTokens + } + if resp.Usage.PromptTokens > 0 || resp.Usage.CompletionTokens > 0 { + contentchoices[0].GenerationInfo["TotalTokens"] = resp.Usage.PromptTokens + resp.Usage.CompletionTokens } } } diff --git a/llms/bedrock/internal/bedrockclient/provider_amazon.go b/llms/bedrock/internal/bedrockclient/provider_amazon.go index 83d84107..e74107a9 100644 --- a/llms/bedrock/internal/bedrockclient/provider_amazon.go +++ b/llms/bedrock/internal/bedrockclient/provider_amazon.go @@ -133,6 +133,10 @@ func createAmazonCompletion(ctx context.Context, GenerationInfo: map[string]any{ "input_tokens": output.InputTextTokenCount, "output_tokens": result.TokenCount, + // Standardized field names for cross-provider compatibility + "PromptTokens": output.InputTextTokenCount, + "CompletionTokens": result.TokenCount, + "TotalTokens": output.InputTextTokenCount + result.TokenCount, }, } } @@ -183,9 +187,14 @@ func parseAmazonStreamingResponse(ctx context.Context, client *bedrockruntime.Cl // Set token counts if resp.InputTextTokenCount > 0 { contentchoices[0].GenerationInfo["input_tokens"] = resp.InputTextTokenCount + contentchoices[0].GenerationInfo["PromptTokens"] = resp.InputTextTokenCount } if resp.OutputTextTokenCount > 0 { contentchoices[0].GenerationInfo["output_tokens"] = resp.OutputTextTokenCount + contentchoices[0].GenerationInfo["CompletionTokens"] = resp.OutputTextTokenCount + } + if resp.InputTextTokenCount > 0 || resp.OutputTextTokenCount > 0 { + contentchoices[0].GenerationInfo["TotalTokens"] = resp.InputTextTokenCount + resp.OutputTextTokenCount } } } diff --git a/llms/bedrock/internal/bedrockclient/provider_anthropic.go b/llms/bedrock/internal/bedrockclient/provider_anthropic.go index 3a305e7c..ee96dcab 100644 --- a/llms/bedrock/internal/bedrockclient/provider_anthropic.go +++ b/llms/bedrock/internal/bedrockclient/provider_anthropic.go @@ -277,6 +277,10 @@ func createAnthropicCompletion(ctx context.Context, GenerationInfo: map[string]any{ "input_tokens": output.Usage.InputTokens, "output_tokens": output.Usage.OutputTokens, + // Standardized field names for cross-provider compatibility + "PromptTokens": output.Usage.InputTokens, + "CompletionTokens": output.Usage.OutputTokens, + "TotalTokens": output.Usage.InputTokens + output.Usage.OutputTokens, }, } Contentchoices = append(Contentchoices, choice) @@ -368,6 +372,8 @@ func parseStreamingCompletionResponse(ctx context.Context, client *bedrockruntim switch resp.Type { case "message_start": contentchoices[0].GenerationInfo["input_tokens"] = resp.Message.Usage.InputTokens + contentchoices[0].GenerationInfo["PromptTokens"] = resp.Message.Usage.InputTokens + contentchoices[0].GenerationInfo["TotalTokens"] = resp.Message.Usage.InputTokens case "content_block_start": if resp.ContentBlock.Type == "tool_use" { currentToolCall = &streaming.ToolCall{ @@ -425,7 +431,16 @@ func parseStreamingCompletionResponse(ctx context.Context, client *bedrockruntim } case "message_delta": contentchoices[0].StopReason = resp.Delta.StopReason - contentchoices[0].GenerationInfo["output_tokens"] = resp.Usage.OutputTokens + inputTokens := resp.Message.Usage.InputTokens + outputTokens := resp.Message.Usage.OutputTokens + if inputTokens == 0 { + if v, ok := contentchoices[0].GenerationInfo["input_tokens"].(int32); ok { + inputTokens = v + } + } + contentchoices[0].GenerationInfo["output_tokens"] = outputTokens + contentchoices[0].GenerationInfo["CompletionTokens"] = outputTokens + contentchoices[0].GenerationInfo["TotalTokens"] = inputTokens + outputTokens } } } diff --git a/llms/bedrock/internal/bedrockclient/provider_meta.go b/llms/bedrock/internal/bedrockclient/provider_meta.go index 3a497bcb..78a7bb2f 100644 --- a/llms/bedrock/internal/bedrockclient/provider_meta.go +++ b/llms/bedrock/internal/bedrockclient/provider_meta.go @@ -120,6 +120,10 @@ func createMetaCompletion(ctx context.Context, GenerationInfo: map[string]any{ "input_tokens": output.PromptTokenCount, "output_tokens": output.GenerationTokenCount, + // Standardized field names for cross-provider compatibility + "PromptTokens": output.PromptTokenCount, + "CompletionTokens": output.GenerationTokenCount, + "TotalTokens": output.PromptTokenCount + output.GenerationTokenCount, }, }, }, @@ -167,9 +171,14 @@ func parseMetaStreamingResponse(ctx context.Context, client *bedrockruntime.Clie // Set token counts if resp.PromptTokenCount > 0 { contentchoices[0].GenerationInfo["input_tokens"] = resp.PromptTokenCount + contentchoices[0].GenerationInfo["PromptTokens"] = resp.PromptTokenCount } if resp.GenerationTokenCount > 0 { contentchoices[0].GenerationInfo["output_tokens"] = resp.GenerationTokenCount + contentchoices[0].GenerationInfo["CompletionTokens"] = resp.GenerationTokenCount + } + if resp.PromptTokenCount > 0 || resp.GenerationTokenCount > 0 { + contentchoices[0].GenerationInfo["TotalTokens"] = resp.PromptTokenCount + resp.GenerationTokenCount } } } diff --git a/llms/bedrock/internal/bedrockclient/provider_nova.go b/llms/bedrock/internal/bedrockclient/provider_nova.go index 2d79b80c..e18fd9c9 100644 --- a/llms/bedrock/internal/bedrockclient/provider_nova.go +++ b/llms/bedrock/internal/bedrockclient/provider_nova.go @@ -226,6 +226,10 @@ func createNovaCompletion(ctx context.Context, GenerationInfo: map[string]any{ "input_tokens": output.Usage.InputTokens, "output_tokens": output.Usage.OutputTokens, + // Standardized field names for cross-provider compatibility + "PromptTokens": output.Usage.InputTokens, + "CompletionTokens": output.Usage.OutputTokens, + "TotalTokens": output.Usage.InputTokens + output.Usage.OutputTokens, }, } } @@ -372,6 +376,7 @@ func parseNovaStreamingResponse(ctx context.Context, client *bedrockruntime.Clie // Check for message start (contains input tokens) if resp.MessageStart.Usage.InputTokens > 0 { contentchoices[0].GenerationInfo["input_tokens"] = resp.MessageStart.Usage.InputTokens + contentchoices[0].GenerationInfo["PromptTokens"] = resp.MessageStart.Usage.InputTokens } // Check for message delta (contains stop reason and output tokens) @@ -380,6 +385,10 @@ func parseNovaStreamingResponse(ctx context.Context, client *bedrockruntime.Clie } if resp.MessageDelta.Usage.OutputTokens > 0 { contentchoices[0].GenerationInfo["output_tokens"] = resp.MessageDelta.Usage.OutputTokens + contentchoices[0].GenerationInfo["CompletionTokens"] = resp.MessageDelta.Usage.OutputTokens + } + if resp.MessageStart.Usage.InputTokens > 0 || resp.MessageDelta.Usage.OutputTokens > 0 { + contentchoices[0].GenerationInfo["TotalTokens"] = resp.MessageStart.Usage.InputTokens + resp.MessageDelta.Usage.OutputTokens } } }