diff --git a/providers/openai/language_model.go b/providers/openai/language_model.go index f9b03ff0c..b46651c8b 100644 --- a/providers/openai/language_model.go +++ b/providers/openai/language_model.go @@ -341,7 +341,12 @@ func (o languageModel) Stream(ctx context.Context, call fantasy.Call) (fantasy.S for stream.Next() { chunk := stream.Current() acc.AddChunk(chunk) - usage, providerMetadata = o.streamUsageFunc(chunk, extraContext, providerMetadata) + // Some OpenAI-compatible backends emit cumulative usage on + // delta chunks and end with a usage-less finish chunk; keep + // the last usage-bearing result instead of zeroing it. + if chunkUsage, chunkMetadata := o.streamUsageFunc(chunk, extraContext, providerMetadata); chunkUsage != (fantasy.Usage{}) { + usage, providerMetadata = chunkUsage, chunkMetadata + } if len(chunk.Choices) == 0 { continue } @@ -870,8 +875,11 @@ func (o languageModel) streamObjectWithJSONMode(ctx context.Context, call fantas for stream.Next() { chunk := stream.Current() - // Update usage - usage, providerMetadata = o.streamUsageFunc(chunk, make(map[string]any), providerMetadata) + // Update usage, ignoring usage-less chunks so a trailing + // finish chunk cannot zero previously reported usage. + if chunkUsage, chunkMetadata := o.streamUsageFunc(chunk, make(map[string]any), providerMetadata); chunkUsage != (fantasy.Usage{}) { + usage, providerMetadata = chunkUsage, chunkMetadata + } if len(chunk.Choices) == 0 { continue diff --git a/providers/openai/stream_usage_test.go b/providers/openai/stream_usage_test.go new file mode 100644 index 000000000..635d797f3 --- /dev/null +++ b/providers/openai/stream_usage_test.go @@ -0,0 +1,96 @@ +package openai + +import ( + "context" + "testing" + + "charm.land/fantasy" + "github.com/stretchr/testify/require" +) + +// Some OpenAI-compatible backends report cumulative usage on delta +// chunks and end with a usage-less finish chunk; the last +// usage-bearing chunk must win. +func TestStreamUsageSurvivesTrailingUsagelessChunk(t *testing.T) { + t.Parallel() + + server := newStreamingMockServer() + defer server.close() + + server.chunks = []string{ + `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", + `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", + `data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"laguna-xs","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":null}` + "\n\n", + "data: [DONE]\n\n", + } + + provider, err := New( + WithAPIKey("test-api-key"), + WithBaseURL(server.server.URL), + ) + require.NoError(t, err) + model, err := provider.LanguageModel(t.Context(), "laguna-xs") + require.NoError(t, err) + + stream, err := model.Stream(context.Background(), fantasy.Call{Prompt: testPrompt}) + require.NoError(t, err) + + parts, err := collectStreamParts(stream) + require.NoError(t, err) + + var finish *fantasy.StreamPart + for i, part := range parts { + if part.Type == fantasy.StreamPartTypeFinish { + finish = &parts[i] + } + } + require.NotNil(t, finish) + require.Equal(t, fantasy.FinishReasonStop, finish.FinishReason) + // prompt_tokens includes cached tokens; input is reported net of cache. + require.Equal(t, int64(20), finish.Usage.InputTokens) + require.Equal(t, int64(80), finish.Usage.CacheReadTokens) + require.Equal(t, int64(5), finish.Usage.OutputTokens) + require.Equal(t, int64(105), finish.Usage.TotalTokens) +} + +func TestStreamObjectUsageSurvivesTrailingUsagelessChunk(t *testing.T) { + t.Parallel() + + server := newStreamingMockServer() + defer server.close() + + server.chunks = []string{ + `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", + `data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"laguna-xs","choices":[{"index":0,"delta":{},"finish_reason":"stop"}],"usage":null}` + "\n\n", + "data: [DONE]\n\n", + } + + provider, err := New( + WithAPIKey("test-api-key"), + WithBaseURL(server.server.URL), + ) + require.NoError(t, err) + model, err := provider.LanguageModel(t.Context(), "laguna-xs") + require.NoError(t, err) + + stream, err := model.StreamObject(context.Background(), fantasy.ObjectCall{ + Prompt: testPrompt, + Schema: fantasy.Schema{ + Type: "object", + Properties: map[string]*fantasy.Schema{ + "answer": {Type: "string"}, + }, + Required: []string{"answer"}, + }, + }) + require.NoError(t, err) + + parts := collectObjectStreamParts(stream) + require.NotEmpty(t, parts) + finish := parts[len(parts)-1] + require.Equal(t, fantasy.ObjectStreamPartTypeFinish, finish.Type) + require.Equal(t, int64(20), finish.Usage.InputTokens) + require.Equal(t, int64(80), finish.Usage.CacheReadTokens) + require.Equal(t, int64(5), finish.Usage.OutputTokens) + require.Equal(t, int64(105), finish.Usage.TotalTokens) +}