From 5ffb145c133e2edd8d19f1fc4a0080514f4f9c3a Mon Sep 17 00:00:00 2001 From: Matt Rubens Date: Fri, 14 Feb 2025 19:26:04 -0500 Subject: [PATCH] Enable streaming for o1 --- .../providers/__tests__/openai-native.test.ts | 26 ++++++++++++------- src/api/providers/openai-native.ts | 4 ++- 2 files changed, 20 insertions(+), 10 deletions(-) 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(