diff --git a/src/api/providers/__tests__/openai-native.test.ts b/src/api/providers/__tests__/openai-native.test.ts index f39d49307f..d6a855849c 100644 --- a/src/api/providers/__tests__/openai-native.test.ts +++ b/src/api/providers/__tests__/openai-native.test.ts @@ -130,11 +130,20 @@ describe("OpenAiNativeHandler", () => { }) mockCreate.mockResolvedValueOnce({ - choices: [{ message: { content: null } }], - usage: { - prompt_tokens: 0, - completion_tokens: 0, - total_tokens: 0, + [Symbol.asyncIterator]: async function* () { + yield { + choices: [ + { + delta: { content: null }, + index: 0, + }, + ], + usage: { + prompt_tokens: 0, + completion_tokens: 0, + total_tokens: 0, + }, + } }, }) @@ -144,10 +153,7 @@ describe("OpenAiNativeHandler", () => { results.push(result) } - expect(results).toEqual([ - { type: "text", text: "" }, - { type: "usage", inputTokens: 0, outputTokens: 0 }, - ]) + expect(results).toEqual([{ type: "usage", inputTokens: 0, outputTokens: 0 }]) // Verify developer role is used for system prompt with o1 model expect(mockCreate).toHaveBeenCalledWith({ @@ -156,6 +162,8 @@ describe("OpenAiNativeHandler", () => { { role: "developer", content: "Formatting re-enabled\n" + systemPrompt }, { role: "user", content: "Hello!" }, ], + stream: true, + stream_options: { include_usage: true }, }) }) diff --git a/src/api/providers/openai-native.ts b/src/api/providers/openai-native.ts index 1a4f9e613a..8feeafdb96 100644 --- a/src/api/providers/openai-native.ts +++ b/src/api/providers/openai-native.ts @@ -56,9 +56,11 @@ export class OpenAiNativeHandler implements ApiHandler, SingleCompletionHandler }, ...convertToOpenAiMessages(messages), ], + stream: true, + stream_options: { include_usage: true }, }) - yield* this.yieldResponseData(response) + yield* this.handleStreamResponse(response) } private async *handleO3FamilyMessage(