Skip to content
Merged
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
36 changes: 33 additions & 3 deletions internal/transformer/stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -724,6 +724,22 @@ func usageInfoToAnthropic(usage *types.UsageInfo) *types.Usage {
}
}

// responsesUsageToAnthropic maps Responses terminal usage to the Anthropic
// terminal message_delta usage. Responses usage has no cache split, so the
// fields map 1:1. A missing terminal event keeps the old zero behavior.
func responsesUsageToAnthropic(usage *types.ResponsesUsage) *types.Usage {
if usage == nil {
return &types.Usage{
InputTokens: 0,
OutputTokens: 0,
}
}
return &types.Usage{
InputTokens: usage.InputTokens,
OutputTokens: usage.OutputTokens,
}
}

// writeContentBlockStop writes a content_block_stop SSE event at the given index.
func writeContentBlockStop(w http.ResponseWriter, index int) error {
return writeSSEEvent(w, types.MessageEvent{
Expand Down Expand Up @@ -818,6 +834,7 @@ func (h *StreamHandler) ProxyResponsesStream(
reasoningStarted := false
hasToolUse := false
startedToolCalls := make(map[string]int)
var terminalUsage *types.ResponsesUsage
readBuf := readBufPool.Get().(*[]byte)
defer readBufPool.Put(readBuf)

Expand All @@ -836,7 +853,7 @@ func (h *StreamHandler) ProxyResponsesStream(
for i := 0; i < n; i++ {
b := (*readBuf)[i]
if b == '\n' {
if err := h.processResponsesSSELine(w, flusher, lineBuf, &contentIndex, &contentStarted, &reasoningStarted, &hasToolUse, startedToolCalls, originalModel); err != nil {
if err := h.processResponsesSSELine(w, flusher, lineBuf, &contentIndex, &contentStarted, &reasoningStarted, &hasToolUse, startedToolCalls, originalModel, &terminalUsage); err != nil {
return err
}
lineBuf = lineBuf[:0]
Expand All @@ -848,7 +865,7 @@ func (h *StreamHandler) ProxyResponsesStream(

if err == io.EOF {
if len(lineBuf) > 0 {
if err := h.processResponsesSSELine(w, flusher, lineBuf, &contentIndex, &contentStarted, &reasoningStarted, &hasToolUse, startedToolCalls, originalModel); err != nil {
if err := h.processResponsesSSELine(w, flusher, lineBuf, &contentIndex, &contentStarted, &reasoningStarted, &hasToolUse, startedToolCalls, originalModel, &terminalUsage); err != nil {
return err
}
}
Expand Down Expand Up @@ -903,7 +920,7 @@ func (h *StreamHandler) ProxyResponsesStream(
Delta: &types.Delta{
StopReason: stopReason,
},
Usage: &types.Usage{InputTokens: 0, OutputTokens: 0},
Usage: responsesUsageToAnthropic(terminalUsage),
}
if err := writeSSEEvent(w, msgDelta); err != nil {
return ErrClientDisconnected
Expand All @@ -930,6 +947,7 @@ func (h *StreamHandler) processResponsesSSELine(
hasToolUse *bool,
startedToolCalls map[string]int,
originalModel string,
terminalUsage **types.ResponsesUsage,
) error {
line = bytes.TrimSpace(line)
if len(line) == 0 || !bytes.HasPrefix(line, []byte("data: ")) {
Expand All @@ -946,6 +964,18 @@ func (h *StreamHandler) processResponsesSSELine(
return nil
}

// Terminal event carries the only usage in a Responses stream.
// Flat shape: {"type":"response.completed","usage":{...}}.
// Nested shape: {"type":"response.completed","response":{"usage":{...}}}.
if chunk.Type == "response.completed" {
if chunk.Usage != nil {
*terminalUsage = chunk.Usage
} else if chunk.Response != nil && chunk.Response.Usage != nil {
*terminalUsage = chunk.Response.Usage
}
return nil
}

if chunk.Type == "response.output_text.delta" && chunk.Delta != "" {
if !*contentStarted {
*contentStarted = true
Expand Down
99 changes: 99 additions & 0 deletions internal/transformer/stream_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1694,3 +1694,102 @@ func TestProxyResponsesStream_ToolCall_ItemField(t *testing.T) {
t.Errorf("event[4] = %+v, want message_delta stop_reason tool_use", events[4])
}
}

// TestProxyResponsesStream_TerminalUsageFlat verifies a flat usage object on
// response.completed lands in the terminal message_delta usage.
func TestProxyResponsesStream_TerminalUsageFlat(t *testing.T) {
handler := NewStreamHandler()
w := newMockResponseWriter()
body := sseLines(
`{"type":"response.output_text.delta","delta":"Hi"}`,
`{"type":"response.completed","usage":{"input_tokens":100,"output_tokens":25}}`,
)

ctx, cancel := context.WithCancel(context.Background())
defer cancel()

if err := handler.ProxyResponsesStream(w, body, "muse-spark-1.3-contributor", ctx, 0, cancel); err != nil {
t.Fatalf("ProxyResponsesStream error: %v", err)
}

events := parseSSEEvents(t, w.buf.String())
if len(events) != 6 {
t.Fatalf("expected 6 events, got %d: %+v", len(events), events)
}
delta := events[4]
if delta.Type != "message_delta" {
t.Fatalf("event[4].Type = %q, want message_delta", delta.Type)
}
if delta.Usage == nil {
t.Fatalf("event[4].Usage = nil, want 100/25")
}
if delta.Usage.InputTokens != 100 || delta.Usage.OutputTokens != 25 {
t.Errorf("event[4].Usage = %+v, want input 100 output 25", delta.Usage)
}
}

// TestProxyResponsesStream_TerminalUsageNested verifies a nested
// response.usage object on response.completed lands in message_delta usage.
func TestProxyResponsesStream_TerminalUsageNested(t *testing.T) {
handler := NewStreamHandler()
w := newMockResponseWriter()
body := sseLines(
`{"type":"response.output_text.delta","delta":"Hi"}`,
`{"type":"response.completed","response":{"usage":{"input_tokens":200,"output_tokens":30}}}`,
)

ctx, cancel := context.WithCancel(context.Background())
defer cancel()

if err := handler.ProxyResponsesStream(w, body, "muse-spark-1.3-contributor", ctx, 0, cancel); err != nil {
t.Fatalf("ProxyResponsesStream error: %v", err)
}

events := parseSSEEvents(t, w.buf.String())
if len(events) != 6 {
t.Fatalf("expected 6 events, got %d: %+v", len(events), events)
}
delta := events[4]
if delta.Type != "message_delta" {
t.Fatalf("event[4].Type = %q, want message_delta", delta.Type)
}
if delta.Usage == nil {
t.Fatalf("event[4].Usage = nil, want 200/30")
}
if delta.Usage.InputTokens != 200 || delta.Usage.OutputTokens != 30 {
t.Errorf("event[4].Usage = %+v, want input 200 output 30", delta.Usage)
}
}

// TestProxyResponsesStream_TerminalUsageMissing verifies a bare
// response.completed keeps zero usage without crashing.
func TestProxyResponsesStream_TerminalUsageMissing(t *testing.T) {
handler := NewStreamHandler()
w := newMockResponseWriter()
body := sseLines(
`{"type":"response.output_text.delta","delta":"Hi"}`,
`{"type":"response.completed"}`,
)

ctx, cancel := context.WithCancel(context.Background())
defer cancel()

if err := handler.ProxyResponsesStream(w, body, "muse-spark-1.3-contributor", ctx, 0, cancel); err != nil {
t.Fatalf("ProxyResponsesStream error: %v", err)
}

events := parseSSEEvents(t, w.buf.String())
if len(events) != 6 {
t.Fatalf("expected 6 events, got %d: %+v", len(events), events)
}
delta := events[4]
if delta.Type != "message_delta" {
t.Fatalf("event[4].Type = %q, want message_delta", delta.Type)
}
if delta.Usage == nil {
t.Fatalf("event[4].Usage = nil, want zero usage")
}
if delta.Usage.InputTokens != 0 || delta.Usage.OutputTokens != 0 {
t.Errorf("event[4].Usage = %+v, want 0/0", delta.Usage)
}
}
21 changes: 14 additions & 7 deletions pkg/types/zen.go
Original file line number Diff line number Diff line change
Expand Up @@ -76,13 +76,20 @@ type ResponsesUsage struct {

// ResponsesChunk represents a streaming chunk from the Responses API.
type ResponsesChunk struct {
Type string `json:"type"`
ID string `json:"id,omitempty"`
ItemID string `json:"item_id,omitempty"`
Delta string `json:"delta,omitempty"`
Item *ResponsesOutput `json:"item,omitempty"`
Output []ResponsesOutput `json:"output,omitempty"`
Usage *ResponsesUsage `json:"usage,omitempty"`
Type string `json:"type"`
ID string `json:"id,omitempty"`
ItemID string `json:"item_id,omitempty"`
Delta string `json:"delta,omitempty"`
Item *ResponsesOutput `json:"item,omitempty"`
Output []ResponsesOutput `json:"output,omitempty"`
Usage *ResponsesUsage `json:"usage,omitempty"`
Response *ResponsesResult `json:"response,omitempty"`
}

// ResponsesResult carries the completed response object inside a
// response.completed event. Only usage is decoded; other fields are ignored.
type ResponsesResult struct {
Usage *ResponsesUsage `json:"usage,omitempty"`
}

// Google Gemini API types.
Expand Down
Loading