mirror of
https://github.com/vxcontrol/langchaingo.git
synced 2026-07-19 21:54:17 -04:00
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:
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user