Skip to content
Open
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
359 changes: 358 additions & 1 deletion src/index.test.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import { test, expect } from 'vitest'
import { z } from 'zod'
import { createFallback } from './index.js'
import { createFallback, TimeToFirstTokenTimeoutError } from './index.js'
import { createOpenAI } from '@ai-sdk/openai'
import { createGroq } from '@ai-sdk/groq'
import { createAnthropic } from '@ai-sdk/anthropic'
Expand Down Expand Up @@ -390,6 +390,363 @@ function sleep(ms: number) {
return new Promise((resolve) => setTimeout(resolve, ms))
}

test('doStream falls back when first token is too slow', async () => {
const slowModel = new MockLanguageModelV4({
provider: 'mock-slow',
modelId: 'slow-model',
doStream: async () => ({
stream: new ReadableStream<LanguageModelV4StreamPart>({
async start(controller) {
controller.enqueue({ type: 'stream-start', warnings: [] })
// Delay 500ms before emitting the first real token
await sleep(500)
controller.enqueue({ type: 'text-start', id: 't1' })
controller.enqueue({ type: 'text-delta', id: 't1', delta: 'slow response' })
controller.enqueue({
type: 'finish',
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
})
controller.close()
},
}),
}),
})

const fastModel = new MockLanguageModelV4({
provider: 'mock-fast',
modelId: 'fast-model',
doStream: async () => ({
stream: new ReadableStream<LanguageModelV4StreamPart>({
start(controller) {
controller.enqueue({ type: 'stream-start', warnings: [] })
controller.enqueue({ type: 'text-start', id: 't1' })
controller.enqueue({ type: 'text-delta', id: 't1', delta: 'fast response' })
controller.enqueue({
type: 'finish',
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
})
controller.close()
},
}),
}),
})

let errorCalled = false
const model = createFallback({
models: [slowModel, fastModel],
timeToFirstTokenTimeout: 100,
onError(error, modelId) {
errorCalled = true
expect(error).toBeInstanceOf(TimeToFirstTokenTimeoutError)
expect(modelId).toBe('slow-model')
},
})

const result = await model.doStream({ prompt: [] })
const reader = result.stream.getReader()
const chunks: LanguageModelV4StreamPart[] = []
while (true) {
const { done, value } = await reader.read()
if (done) break
chunks.push(value)
}

expect(errorCalled).toBe(true)
expect(model.currentModelIndex).toBe(1)
const textChunks = chunks.filter((c) => c.type === 'text-delta')
expect(textChunks[0]).toMatchObject({ delta: 'fast response' })
})

test('doStream does NOT fall back when first token is fast enough', async () => {
const fastModel = new MockLanguageModelV4({
provider: 'mock-fast',
modelId: 'primary-fast',
doStream: async () => ({
stream: new ReadableStream<LanguageModelV4StreamPart>({
async start(controller) {
controller.enqueue({ type: 'stream-start', warnings: [] })
// Small delay, well within timeout
await sleep(20)
controller.enqueue({ type: 'text-start', id: 't1' })
controller.enqueue({ type: 'text-delta', id: 't1', delta: 'primary response' })
controller.enqueue({
type: 'finish',
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
})
controller.close()
},
}),
}),
})

const fallbackModel = new MockLanguageModelV4({
provider: 'mock-fallback',
modelId: 'fallback-model',
doStream: async () => ({
stream: new ReadableStream<LanguageModelV4StreamPart>({
start(controller) {
controller.enqueue({ type: 'stream-start', warnings: [] })
controller.enqueue({ type: 'text-start', id: 't1' })
controller.enqueue({ type: 'text-delta', id: 't1', delta: 'fallback response' })
controller.close()
},
}),
}),
})

const model = createFallback({
models: [fastModel, fallbackModel],
timeToFirstTokenTimeout: 2000,
})

const result = await model.doStream({ prompt: [] })
const reader = result.stream.getReader()
const chunks: LanguageModelV4StreamPart[] = []
while (true) {
const { done, value } = await reader.read()
if (done) break
chunks.push(value)
}

expect(model.currentModelIndex).toBe(0)
const textChunks = chunks.filter((c) => c.type === 'text-delta')
expect(textChunks[0]).toMatchObject({ delta: 'primary response' })
})

test('doGenerate falls back when response is too slow', async () => {
const generateResult: LanguageModelV4GenerateResult = {
content: [{ type: 'text', text: 'fast result' }],
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
warnings: [],
}

const slowModel = new MockLanguageModelV4({
provider: 'mock-slow',
modelId: 'slow-generate',
doGenerate: async () => {
await sleep(500)
return generateResult
},
})

const fastModel = new MockLanguageModelV4({
provider: 'mock-fast',
modelId: 'fast-generate',
doGenerate: generateResult,
})

let errorCalled = false
const model = createFallback({
models: [slowModel, fastModel],
timeToFirstTokenTimeout: 100,
onError(error, modelId) {
errorCalled = true
expect(error).toBeInstanceOf(TimeToFirstTokenTimeoutError)
expect(modelId).toBe('slow-generate')
},
})

const result = await model.doGenerate({ prompt: [] })

expect(errorCalled).toBe(true)
expect(model.currentModelIndex).toBe(1)
expect(result.content).toEqual([{ type: 'text', text: 'fast result' }])
})

test('TTFT timeout retries even with custom shouldRetryThisError', async () => {
const generateResult: LanguageModelV4GenerateResult = {
content: [{ type: 'text', text: 'ok' }],
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
warnings: [],
}

const slowModel = new MockLanguageModelV4({
provider: 'mock-slow',
modelId: 'slow',
doGenerate: async () => {
await sleep(500)
return generateResult
},
})
const fastModel = new MockLanguageModelV4({
provider: 'mock-fast',
modelId: 'fast',
doGenerate: generateResult,
})

const model = createFallback({
models: [slowModel, fastModel],
timeToFirstTokenTimeout: 100,
// Custom predicate that rejects everything — TTFT should still retry
shouldRetryThisError: () => false,
})

const result = await model.doGenerate({ prompt: [] })
expect(model.currentModelIndex).toBe(1)
expect(result.content).toEqual([{ type: 'text', text: 'ok' }])
})

test('doStream TTFT covers slow doStream() startup', async () => {
// The doStream call itself is slow (simulating slow HTTP handshake)
const slowStartupModel = new MockLanguageModelV4({
provider: 'mock-slow-startup',
modelId: 'slow-startup',
doStream: async () => {
await sleep(500)
return {
stream: new ReadableStream<LanguageModelV4StreamPart>({
start(controller) {
controller.enqueue({ type: 'stream-start', warnings: [] })
controller.enqueue({ type: 'text-start', id: 't1' })
controller.enqueue({ type: 'text-delta', id: 't1', delta: 'slow' })
controller.enqueue({
type: 'finish',
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
})
controller.close()
},
}),
}
},
})

const fastModel = new MockLanguageModelV4({
provider: 'mock-fast',
modelId: 'fast-startup',
doStream: async () => ({
stream: new ReadableStream<LanguageModelV4StreamPart>({
start(controller) {
controller.enqueue({ type: 'stream-start', warnings: [] })
controller.enqueue({ type: 'text-start', id: 't1' })
controller.enqueue({ type: 'text-delta', id: 't1', delta: 'fast' })
controller.enqueue({
type: 'finish',
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
})
controller.close()
},
}),
}),
})

const model = createFallback({
models: [slowStartupModel, fastModel],
timeToFirstTokenTimeout: 100,
})

const result = await model.doStream({ prompt: [] })
const reader = result.stream.getReader()
const chunks: LanguageModelV4StreamPart[] = []
while (true) {
const { done, value } = await reader.read()
if (done) break
chunks.push(value)
}

expect(model.currentModelIndex).toBe(1)
const textChunks = chunks.filter((c) => c.type === 'text-delta')
expect(textChunks[0]).toMatchObject({ delta: 'fast' })
})

test('doStream TTFT does not clear on text-start, only on text-delta', async () => {
// Model emits text-start immediately but delays before text-delta
const model_with_slow_content = new MockLanguageModelV4({
provider: 'mock',
modelId: 'slow-content',
doStream: async () => ({
stream: new ReadableStream<LanguageModelV4StreamPart>({
async start(controller) {
controller.enqueue({ type: 'stream-start', warnings: [] })
// text-start arrives fast
controller.enqueue({ type: 'text-start', id: 't1' })
// but actual content is slow
await sleep(500)
controller.enqueue({ type: 'text-delta', id: 't1', delta: 'slow content' })
controller.enqueue({
type: 'finish',
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
})
controller.close()
},
}),
}),
})

const fastModel = new MockLanguageModelV4({
provider: 'mock-fast',
modelId: 'fast-content',
doStream: async () => ({
stream: new ReadableStream<LanguageModelV4StreamPart>({
start(controller) {
controller.enqueue({ type: 'stream-start', warnings: [] })
controller.enqueue({ type: 'text-start', id: 't1' })
controller.enqueue({ type: 'text-delta', id: 't1', delta: 'fast content' })
controller.enqueue({
type: 'finish',
finishReason: { unified: 'stop', raw: 'stop' },
usage: {
inputTokens: { total: 1, noCache: 1, cacheRead: undefined, cacheWrite: undefined },
outputTokens: { total: 1, text: 1, reasoning: undefined },
},
})
controller.close()
},
}),
}),
})

const model = createFallback({
models: [model_with_slow_content, fastModel],
timeToFirstTokenTimeout: 100,
})

const result = await model.doStream({ prompt: [] })
const reader = result.stream.getReader()
const chunks: LanguageModelV4StreamPart[] = []
while (true) {
const { done, value } = await reader.read()
if (done) break
chunks.push(value)
}

// Should have timed out because text-start doesn't count as content
expect(model.currentModelIndex).toBe(1)
const textChunks = chunks.filter((c) => c.type === 'text-delta')
expect(textChunks[0]).toMatchObject({ delta: 'fast content' })
})

test(
'handles overloaded_error from reader.read() and retries with fallback model',
async () => {
Expand Down
Loading