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
106 changes: 28 additions & 78 deletions src/api/providers/__tests__/fireworks.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import OpenAI from "openai"
import { type FireworksModelId, fireworksDefaultModelId, fireworksModels } from "@roo-code/types"

import { FireworksHandler } from "../fireworks"
import { asyncStreamFrom, collectStream } from "../../../test-utils/stream"

// Create mock functions
const mockCreate = vi.fn()
Expand All @@ -29,18 +30,18 @@ describe("FireworksHandler", () => {
beforeEach(() => {
vi.clearAllMocks()
// Set up default mock implementation
mockCreate.mockImplementation(async () => ({
[Symbol.asyncIterator]: async function* () {
yield {
mockCreate.mockImplementation(async () =>
asyncStreamFrom([
{
choices: [
{
delta: { content: "Test response" },
index: 0,
},
],
usage: null,
}
yield {
},
{
choices: [
{
delta: {},
Expand All @@ -52,9 +53,9 @@ describe("FireworksHandler", () => {
completion_tokens: 5,
total_tokens: 15,
},
}
},
}))
},
]),
)
handler = new FireworksHandler({ fireworksApiKey: "test-key" })
})

Expand Down Expand Up @@ -436,19 +437,7 @@ describe("FireworksHandler", () => {
it("createMessage should yield text content from stream", async () => {
const testContent = "This is test content from Fireworks stream"

mockCreate.mockImplementationOnce(() => {
return {
[Symbol.asyncIterator]: () => ({
next: vi
.fn()
.mockResolvedValueOnce({
done: false,
value: { choices: [{ delta: { content: testContent } }] },
})
.mockResolvedValueOnce({ done: true }),
}),
}
})
mockCreate.mockImplementationOnce(() => asyncStreamFrom([{ choices: [{ delta: { content: testContent } }] }]))

const stream = handler.createMessage("system prompt", [])
const firstChunk = await stream.next()
Expand All @@ -458,19 +447,9 @@ describe("FireworksHandler", () => {
})

it("createMessage should yield usage data from stream", async () => {
mockCreate.mockImplementationOnce(() => {
return {
[Symbol.asyncIterator]: () => ({
next: vi
.fn()
.mockResolvedValueOnce({
done: false,
value: { choices: [{ delta: {} }], usage: { prompt_tokens: 10, completion_tokens: 20 } },
})
.mockResolvedValueOnce({ done: true }),
}),
}
})
mockCreate.mockImplementationOnce(() =>
asyncStreamFrom([{ choices: [{ delta: {} }], usage: { prompt_tokens: 10, completion_tokens: 20 } }]),
)

const stream = handler.createMessage("system prompt", [])
const firstChunk = await stream.next()
Expand All @@ -487,15 +466,7 @@ describe("FireworksHandler", () => {
fireworksApiKey: "test-fireworks-api-key",
})

mockCreate.mockImplementationOnce(() => {
return {
[Symbol.asyncIterator]: () => ({
async next() {
return { done: true }
},
}),
}
})
mockCreate.mockImplementationOnce(() => asyncStreamFrom([]))

const systemPrompt = "Test system prompt for Fireworks"
const messages: Anthropic.Messages.MessageParam[] = [{ role: "user", content: "Test message for Fireworks" }]
Expand Down Expand Up @@ -523,13 +494,7 @@ describe("FireworksHandler", () => {
fireworksApiKey: "test-fireworks-api-key",
})

mockCreate.mockImplementationOnce(() => ({
[Symbol.asyncIterator]: () => ({
async next() {
return { done: true }
},
}),
}))
mockCreate.mockImplementationOnce(() => asyncStreamFrom([]))

const messageGenerator = handlerWithModel.createMessage("system", [])
await messageGenerator.next()
Expand All @@ -549,13 +514,7 @@ describe("FireworksHandler", () => {
fireworksApiKey: "test-fireworks-api-key",
})

mockCreate.mockImplementationOnce(() => ({
[Symbol.asyncIterator]: () => ({
async next() {
return { done: true }
},
}),
}))
mockCreate.mockImplementationOnce(() => asyncStreamFrom([]))

const messageGenerator = handlerWithModel.createMessage("system", [])
await messageGenerator.next()
Expand All @@ -577,13 +536,7 @@ describe("FireworksHandler", () => {
modelTemperature: 0.7,
})

mockCreate.mockImplementationOnce(() => ({
[Symbol.asyncIterator]: () => ({
async next() {
return { done: true }
},
}),
}))
mockCreate.mockImplementationOnce(() => asyncStreamFrom([]))

const messageGenerator = handlerWithModel.createMessage("system", [])
await messageGenerator.next()
Expand All @@ -610,27 +563,27 @@ describe("FireworksHandler", () => {
})

it("createMessage should handle stream with multiple chunks", async () => {
mockCreate.mockImplementationOnce(async () => ({
[Symbol.asyncIterator]: async function* () {
yield {
mockCreate.mockImplementationOnce(async () =>
asyncStreamFrom([
{
choices: [
{
delta: { content: "Hello" },
index: 0,
},
],
usage: null,
}
yield {
},
{
choices: [
{
delta: { content: " world" },
index: 0,
},
],
usage: null,
}
yield {
},
{
choices: [
{
delta: {},
Expand All @@ -642,18 +595,15 @@ describe("FireworksHandler", () => {
completion_tokens: 10,
total_tokens: 15,
},
}
},
}))
},
]),
)

const systemPrompt = "You are a helpful assistant."
const messages: Anthropic.Messages.MessageParam[] = [{ role: "user", content: "Hi" }]

const stream = handler.createMessage(systemPrompt, messages)
const chunks = []
for await (const chunk of stream) {
chunks.push(chunk)
}
const chunks = await collectStream(stream)

expect(chunks[0]).toEqual({ type: "text", text: "Hello" })
expect(chunks[1]).toEqual({ type: "text", text: " world" })
Expand Down
Loading
Loading