Skip to content

Commit 833854a

Browse files
authored
fix(providers/openai): keep last usage when stream ends with usage-less chunk (#52)
Some OpenAI-compatible backends (e.g. Poolside laguna-xs when tools are declared) report cumulative usage on every delta chunk and end the stream with a finish_reason chunk whose usage is null, without a trailing usage-only chunk. The stream loops reassigned usage from every chunk, so the trailing usage-less chunk wiped the real usage (and provider metadata) to zero, and consumers saw Usage{0,0,0}. Only adopt the stream usage hook's result when it actually reports usage, in both the chat-completions stream loop and the JSON-mode object stream loop.
1 parent 6f8df37 commit 833854a

2 files changed

Lines changed: 107 additions & 3 deletions

File tree

‎providers/openai/language_model.go‎

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -341,7 +341,12 @@ func (o languageModel) Stream(ctx context.Context, call fantasy.Call) (fantasy.S
341341
for stream.Next() {
342342
chunk := stream.Current()
343343
acc.AddChunk(chunk)
344-
usage, providerMetadata = o.streamUsageFunc(chunk, extraContext, providerMetadata)
344+
// Some OpenAI-compatible backends emit cumulative usage on
345+
// delta chunks and end with a usage-less finish chunk; keep
346+
// the last usage-bearing result instead of zeroing it.
347+
if chunkUsage, chunkMetadata := o.streamUsageFunc(chunk, extraContext, providerMetadata); chunkUsage != (fantasy.Usage{}) {
348+
usage, providerMetadata = chunkUsage, chunkMetadata
349+
}
345350
if len(chunk.Choices) == 0 {
346351
continue
347352
}
@@ -870,8 +875,11 @@ func (o languageModel) streamObjectWithJSONMode(ctx context.Context, call fantas
870875
for stream.Next() {
871876
chunk := stream.Current()
872877

873-
// Update usage
874-
usage, providerMetadata = o.streamUsageFunc(chunk, make(map[string]any), providerMetadata)
878+
// Update usage, ignoring usage-less chunks so a trailing
879+
// finish chunk cannot zero previously reported usage.
880+
if chunkUsage, chunkMetadata := o.streamUsageFunc(chunk, make(map[string]any), providerMetadata); chunkUsage != (fantasy.Usage{}) {
881+
usage, providerMetadata = chunkUsage, chunkMetadata
882+
}
875883

876884
if len(chunk.Choices) == 0 {
877885
continue
Lines changed: 96 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,96 @@
1+
package openai
2+
3+
import (
4+
"context"
5+
"testing"
6+
7+
"charm.land/fantasy"
8+
"github.com/stretchr/testify/require"
9+
)
10+
11+
// Some OpenAI-compatible backends report cumulative usage on delta
12+
// chunks and end with a usage-less finish chunk; the last
13+
// usage-bearing chunk must win.
14+
func TestStreamUsageSurvivesTrailingUsagelessChunk(t *testing.T) {
15+
t.Parallel()
16+
17+
server := newStreamingMockServer()
18+
defer server.close()
19+
20+
server.chunks = []string{
21+
`data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"laguna-xs","choices":[{"index":0,"delta":{"role":"assistant","content":"Hel"},"finish_reason":null}],"usage":{"prompt_tokens":100,"completion_tokens":1,"total_tokens":101}}` + "\n\n",
22+
`data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"laguna-xs","choices":[{"index":0,"delta":{"content":"lo"},"finish_reason":null}],"usage":{"prompt_tokens":100,"completion_tokens":5,"total_tokens":105,"prompt_tokens_details":{"cached_tokens":80}}}` + "\n\n",
23+
`data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"laguna-xs","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":null}` + "\n\n",
24+
"data: [DONE]\n\n",
25+
}
26+
27+
provider, err := New(
28+
WithAPIKey("test-api-key"),
29+
WithBaseURL(server.server.URL),
30+
)
31+
require.NoError(t, err)
32+
model, err := provider.LanguageModel(t.Context(), "laguna-xs")
33+
require.NoError(t, err)
34+
35+
stream, err := model.Stream(context.Background(), fantasy.Call{Prompt: testPrompt})
36+
require.NoError(t, err)
37+
38+
parts, err := collectStreamParts(stream)
39+
require.NoError(t, err)
40+
41+
var finish *fantasy.StreamPart
42+
for i, part := range parts {
43+
if part.Type == fantasy.StreamPartTypeFinish {
44+
finish = &parts[i]
45+
}
46+
}
47+
require.NotNil(t, finish)
48+
require.Equal(t, fantasy.FinishReasonStop, finish.FinishReason)
49+
// prompt_tokens includes cached tokens; input is reported net of cache.
50+
require.Equal(t, int64(20), finish.Usage.InputTokens)
51+
require.Equal(t, int64(80), finish.Usage.CacheReadTokens)
52+
require.Equal(t, int64(5), finish.Usage.OutputTokens)
53+
require.Equal(t, int64(105), finish.Usage.TotalTokens)
54+
}
55+
56+
func TestStreamObjectUsageSurvivesTrailingUsagelessChunk(t *testing.T) {
57+
t.Parallel()
58+
59+
server := newStreamingMockServer()
60+
defer server.close()
61+
62+
server.chunks = []string{
63+
`data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"laguna-xs","choices":[{"index":0,"delta":{"role":"assistant","content":"{\"answer\":\"hello\"}"},"finish_reason":null}],"usage":{"prompt_tokens":100,"completion_tokens":5,"total_tokens":105,"prompt_tokens_details":{"cached_tokens":80}}}` + "\n\n",
64+
`data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"laguna-xs","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":null}` + "\n\n",
65+
"data: [DONE]\n\n",
66+
}
67+
68+
provider, err := New(
69+
WithAPIKey("test-api-key"),
70+
WithBaseURL(server.server.URL),
71+
)
72+
require.NoError(t, err)
73+
model, err := provider.LanguageModel(t.Context(), "laguna-xs")
74+
require.NoError(t, err)
75+
76+
stream, err := model.StreamObject(context.Background(), fantasy.ObjectCall{
77+
Prompt: testPrompt,
78+
Schema: fantasy.Schema{
79+
Type: "object",
80+
Properties: map[string]*fantasy.Schema{
81+
"answer": {Type: "string"},
82+
},
83+
Required: []string{"answer"},
84+
},
85+
})
86+
require.NoError(t, err)
87+
88+
parts := collectObjectStreamParts(stream)
89+
require.NotEmpty(t, parts)
90+
finish := parts[len(parts)-1]
91+
require.Equal(t, fantasy.ObjectStreamPartTypeFinish, finish.Type)
92+
require.Equal(t, int64(20), finish.Usage.InputTokens)
93+
require.Equal(t, int64(80), finish.Usage.CacheReadTokens)
94+
require.Equal(t, int64(5), finish.Usage.OutputTokens)
95+
require.Equal(t, int64(105), finish.Usage.TotalTokens)
96+
}

0 commit comments

Comments
 (0)