fix(translator/claude): defer message_delta and cache streaming usage

- Cache token usage metrics from streaming chunks instead of prematurely finalizing content blocks.
- Defer `message_delta` and `message_stop` emissions until encountering a finish reason, a trailing usage chunk, or stream completion.

Closes: #5419
This commit is contained in:
Luis Pater
2026-09-03 22:14:30 +08:00
parent 728ea8b855
commit f804fb5f30
2 changed files with 152 additions and 11 deletions

View File

@@ -58,6 +58,11 @@ type ConvertOpenAIResponseToAnthropicParams struct {
ThinkingContentBlockIndex int
// Next available content block index
NextContentBlockIndex int
// Usage metrics cached from streaming chunks
UsageInputTokens int64
UsageOutputTokens int64
UsageCachedTokens int64
UsageCacheWriteTokens int64
}
// ToolCallAccumulator holds the state for accumulating tool call data
@@ -101,6 +106,10 @@ func ConvertOpenAIResponseToClaude(_ context.Context, _ string, originalRequestR
TextContentBlockIndex: -1,
ThinkingContentBlockIndex: -1,
NextContentBlockIndex: 0,
UsageInputTokens: 0,
UsageOutputTokens: 0,
UsageCachedTokens: 0,
UsageCacheWriteTokens: 0,
}
}
@@ -304,16 +313,23 @@ func convertOpenAIStreamingChunkToAnthropic(rawJSON []byte, param *ConvertOpenAI
// Don't send message_delta here - wait for usage info or [DONE]
}
// Handle usage information separately (this comes in a later chunk)
// Only process if usage has actual values (not null)
if !param.MessageDeltaSent && (param.FinishReason != "" || param.SawToolCall) {
usage := root.Get("usage")
if usage.Exists() && usage.Type != gjson.Null {
finalizeOpenAIAnthropicContentBlocks(param, &results)
inputTokens, outputTokens, cachedTokens, cacheWriteTokens := extractOpenAIUsage(usage)
emitAnthropicMessageDelta(param, &results, inputTokens, outputTokens, cachedTokens, cacheWriteTokens)
emitMessageStopIfNeeded(param, &results)
}
// Cache usage information whenever present
usage := root.Get("usage")
hasUsage := usage.Exists() && usage.Type != gjson.Null
if hasUsage {
param.UsageInputTokens, param.UsageOutputTokens, param.UsageCachedTokens, param.UsageCacheWriteTokens = extractOpenAIUsage(usage)
}
// Emit message_delta and message_stop only when generation is finished:
// 1. Upstream provided a finish_reason, or
// 2. Upstream sent a trailing usage-only chunk (choices array is empty or absent) after content/tools started.
isTrailingUsageChunk := hasUsage && !root.Get("choices.0").Exists() &&
(param.FinishReason != "" || param.SawToolCall || param.TextContentBlockStarted || param.ThinkingContentBlockStarted || param.ContentAccumulator.Len() > 0)
if !param.MessageDeltaSent && (param.FinishReason != "" || isTrailingUsageChunk) && hasUsage {
finalizeOpenAIAnthropicContentBlocks(param, &results)
emitAnthropicMessageDelta(param, &results, param.UsageInputTokens, param.UsageOutputTokens, param.UsageCachedTokens, param.UsageCacheWriteTokens)
emitMessageStopIfNeeded(param, &results)
}
return results
@@ -326,7 +342,7 @@ func convertOpenAIDoneToAnthropic(param *ConvertOpenAIResponseToAnthropicParams)
finalizeOpenAIAnthropicContentBlocks(param, &results)
if !param.MessageDeltaSent {
emitAnthropicMessageDelta(param, &results, 0, 0, 0, 0)
emitAnthropicMessageDelta(param, &results, param.UsageInputTokens, param.UsageOutputTokens, param.UsageCachedTokens, param.UsageCacheWriteTokens)
}
emitMessageStopIfNeeded(param, &results)

View File

@@ -518,6 +518,131 @@ func TestStreamingTool_UsageWithoutFinishReasonEmitsMessageDelta(t *testing.T) {
}
}
func TestStreamingTool_PerChunkUsagePreservesToolArguments(t *testing.T) {
events := runStream(t, streamReq,
`{"id":"c1","model":"m","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"Skill","arguments":""}}]},"finish_reason":null}],"usage":{"prompt_tokens":191,"completion_tokens":5}}`,
`{"id":"c1","model":"m","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"function":{"arguments":"{\"skill\": \"stop-s"}}]},"finish_reason":null}],"usage":{"prompt_tokens":191,"completion_tokens":10}}`,
`{"id":"c1","model":"m","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"function":{"arguments":"lop\"}"}}]},"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":191,"completion_tokens":15}}`,
)
starts := toolUseStarts(events)
if len(starts) != 1 {
t.Fatalf("expected one tool_use start, got %d (starts=%+v)", len(starts), starts)
}
if name := gjson.Get(starts[0].Payload, "content_block.name").String(); name != "Skill" {
t.Fatalf("tool name = %q, want %q", name, "Skill")
}
if id := gjson.Get(starts[0].Payload, "content_block.id").String(); id != "call_1" {
t.Fatalf("tool id = %q, want %q", id, "call_1")
}
var deltas []sseEvent
for _, e := range events {
if e.Type == "content_block_delta" && gjson.Get(e.Payload, "delta.type").String() == "input_json_delta" {
deltas = append(deltas, e)
}
}
if len(deltas) == 0 {
t.Fatalf("expected at least one input_json_delta, got none (events=%+v)", events)
}
var mergedArgs strings.Builder
for _, d := range deltas {
mergedArgs.WriteString(gjson.Get(d.Payload, "delta.partial_json").String())
}
if merged := mergedArgs.String(); merged != `{"skill": "stop-slop"}` {
t.Fatalf("merged arguments = %q, want %q", merged, `{"skill": "stop-slop"}`)
}
if got := countByType(events, "message_delta"); got != 1 {
t.Fatalf("expected exactly one message_delta, got %d (events=%+v)", got, events)
}
if got := lastStopReason(events); got != "tool_use" {
t.Fatalf("stop_reason = %q, want %q", got, "tool_use")
}
if got := countByType(events, "message_stop"); got != 1 {
t.Fatalf("expected exactly one message_stop, got %d (events=%+v)", got, events)
}
var deltaEvent *sseEvent
for _, e := range events {
if e.Type == "message_delta" {
deltaEvent = &e
break
}
}
if deltaEvent == nil {
t.Fatalf("missing message_delta event")
}
if input := gjson.Get(deltaEvent.Payload, "usage.input_tokens").Int(); input != 191 {
t.Fatalf("input_tokens = %d, want 191", input)
}
if output := gjson.Get(deltaEvent.Payload, "usage.output_tokens").Int(); output != 15 {
t.Fatalf("output_tokens = %d, want 15", output)
}
}
func TestStreamingTool_PerChunkUsageOmittedFinishReasonPreservesToolArguments(t *testing.T) {
events := runStream(t, streamReq,
`{"id":"c1","model":"m","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"Skill","arguments":""}}]},"finish_reason":null}],"usage":{"prompt_tokens":191,"completion_tokens":5}}`,
`{"id":"c1","model":"m","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"function":{"arguments":"{\"skill\": \"stop-s"}}]},"finish_reason":null}],"usage":{"prompt_tokens":191,"completion_tokens":10}}`,
`{"id":"c1","model":"m","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[{"index":0,"function":{"arguments":"lop\"}"}}]},"finish_reason":null}],"usage":{"prompt_tokens":191,"completion_tokens":15}}`,
)
starts := toolUseStarts(events)
if len(starts) != 1 {
t.Fatalf("expected one tool_use start, got %d (starts=%+v)", len(starts), starts)
}
if name := gjson.Get(starts[0].Payload, "content_block.name").String(); name != "Skill" {
t.Fatalf("tool name = %q, want %q", name, "Skill")
}
var deltas []sseEvent
for _, e := range events {
if e.Type == "content_block_delta" && gjson.Get(e.Payload, "delta.type").String() == "input_json_delta" {
deltas = append(deltas, e)
}
}
if len(deltas) == 0 {
t.Fatalf("expected at least one input_json_delta, got none (events=%+v)", events)
}
var mergedArgs strings.Builder
for _, d := range deltas {
mergedArgs.WriteString(gjson.Get(d.Payload, "delta.partial_json").String())
}
if merged := mergedArgs.String(); merged != `{"skill": "stop-slop"}` {
t.Fatalf("merged arguments = %q, want %q", merged, `{"skill": "stop-slop"}`)
}
if got := countByType(events, "message_delta"); got != 1 {
t.Fatalf("expected exactly one message_delta, got %d (events=%+v)", got, events)
}
if got := lastStopReason(events); got != "tool_use" {
t.Fatalf("stop_reason = %q, want %q", got, "tool_use")
}
if got := countByType(events, "message_stop"); got != 1 {
t.Fatalf("expected exactly one message_stop, got %d (events=%+v)", got, events)
}
var deltaEvent *sseEvent
for _, e := range events {
if e.Type == "message_delta" {
deltaEvent = &e
break
}
}
if deltaEvent == nil {
t.Fatalf("missing message_delta event")
}
if input := gjson.Get(deltaEvent.Payload, "usage.input_tokens").Int(); input != 191 {
t.Fatalf("input_tokens = %d, want 191", input)
}
if output := gjson.Get(deltaEvent.Payload, "usage.output_tokens").Int(); output != 15 {
t.Fatalf("output_tokens = %d, want 15", output)
}
}
func TestStreamingTool_OmittedToolCallIndexPreservesParallelCalls(t *testing.T) {
events := runStream(t, streamReq,
`{"id":"c1","model":"m","choices":[{"index":0,"delta":{"role":"assistant","tool_calls":[