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.
This commit is contained in:
Dmitry Ng
2026-01-28 18:05:51 +03:00
parent 63e01f2f31
commit f79b68a0db
5 changed files with 52 additions and 1 deletions
@@ -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
}
}
}
@@ -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
}
}
}
@@ -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
}
}
}
@@ -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
}
}
}
@@ -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
}
}
}