fix(llm-proxy): reserve quota before the provider call and stop cutting off long streams
This commit is contained in:
parent
2f94d24c99
commit
b306034bfc
7 changed files with 1489 additions and 307 deletions
378
server/supabase/functions/llm-proxy/handler.test.ts
Normal file
378
server/supabase/functions/llm-proxy/handler.test.ts
Normal file
|
|
@ -0,0 +1,378 @@
|
|||
// Regression tests for llm-proxy:
|
||||
// * quota is reserved before the provider call, so parallel requests past the
|
||||
// allowance never reach Anthropic;
|
||||
// * the provider's first-byte deadline no longer cuts off a long stream, while
|
||||
// idle/total stream deadlines and client disconnects still stop it;
|
||||
// * provider-side failures refund the reservation.
|
||||
import type { QuotaFeature, QuotaPeriod, QuotaReservation } from '../_shared/quota.ts'
|
||||
import {
|
||||
createLlmProxyHandler,
|
||||
type LlmProxyDeps,
|
||||
type LlmQuotaStore,
|
||||
type ProviderFetch,
|
||||
} from './handler.ts'
|
||||
import { createSseCompletionTracker } from './provider-deadline.ts'
|
||||
|
||||
function assert(condition: boolean, message: string): asserts condition {
|
||||
if (!condition) throw new Error(message)
|
||||
}
|
||||
|
||||
function assertEquals<T>(actual: T, expected: T, message: string): void {
|
||||
const a = JSON.stringify(actual)
|
||||
const e = JSON.stringify(expected)
|
||||
if (a !== e) throw new Error(`${message}: expected ${e}, got ${a}`)
|
||||
}
|
||||
|
||||
const delay = (ms: number) => new Promise<void>((resolve) => setTimeout(resolve, ms))
|
||||
const encoder = new TextEncoder()
|
||||
|
||||
/** In-memory reserve/finalize with the same in-flight semantics as reserve_llm_quota. */
|
||||
class FakeLlmQuotaStore implements LlmQuotaStore {
|
||||
used = 0
|
||||
reserveCalls = 0
|
||||
finalized: Array<{ id: string; succeeded: boolean }> = []
|
||||
private reservations = new Map<string, 'reserved' | 'completed' | 'released'>()
|
||||
constructor(private readonly limit: number, private readonly tier: 'free' | 'pro' = 'pro') {}
|
||||
|
||||
readTier(): Promise<'free' | 'pro'> {
|
||||
return Promise.resolve(this.tier)
|
||||
}
|
||||
|
||||
reserve(
|
||||
_userId: string,
|
||||
reservationId: string,
|
||||
_feature: QuotaFeature,
|
||||
_baseLimit: number,
|
||||
period: QuotaPeriod,
|
||||
): Promise<QuotaReservation> {
|
||||
this.reserveCalls++
|
||||
if (this.used >= this.limit) {
|
||||
return Promise.resolve({
|
||||
allowed: false, reservationId: null, status: 'denied', current: this.used, limit: this.limit,
|
||||
period, tier: this.tier, overageCredits: 0, consumedFrom: 'none',
|
||||
})
|
||||
}
|
||||
this.used++
|
||||
this.reservations.set(reservationId, 'reserved')
|
||||
return Promise.resolve({
|
||||
allowed: true, reservationId, status: 'reserved', current: this.used, limit: this.limit,
|
||||
period, tier: this.tier, overageCredits: 0, consumedFrom: 'base',
|
||||
})
|
||||
}
|
||||
|
||||
finalize(reservationId: string, succeeded: boolean): Promise<'completed' | 'released'> {
|
||||
this.finalized.push({ id: reservationId, succeeded })
|
||||
if (this.reservations.get(reservationId) === 'reserved' && !succeeded) this.used--
|
||||
const status = succeeded ? 'completed' : 'released'
|
||||
this.reservations.set(reservationId, status)
|
||||
return Promise.resolve(status)
|
||||
}
|
||||
}
|
||||
|
||||
interface UpstreamPlan {
|
||||
chunks: string[]
|
||||
intervalMs: number
|
||||
/** Stop producing after this many chunks and hang (no more data, no EOF). */
|
||||
stallAfter?: number
|
||||
}
|
||||
|
||||
/**
|
||||
* Mock of Deno's fetch body semantics: aborting the request signal after the
|
||||
* headers arrived errors the response body stream.
|
||||
*/
|
||||
function slowSseResponse(plan: UpstreamPlan, signal: AbortSignal): { response: Response; state: { aborted: boolean } } {
|
||||
const state = { aborted: false }
|
||||
let index = 0
|
||||
let timer: ReturnType<typeof setTimeout> | null = null
|
||||
let wake: (() => void) | null = null
|
||||
let streamController: ReadableStreamDefaultController<Uint8Array> | null = null
|
||||
const onAbort = () => {
|
||||
state.aborted = true
|
||||
if (timer !== null) clearTimeout(timer)
|
||||
timer = null
|
||||
wake?.()
|
||||
try {
|
||||
streamController?.error(signal.reason)
|
||||
} catch {
|
||||
// closed
|
||||
}
|
||||
}
|
||||
signal.addEventListener('abort', onAbort, { once: true })
|
||||
const body = new ReadableStream<Uint8Array>({
|
||||
start(controller) {
|
||||
streamController = controller
|
||||
},
|
||||
async pull(controller) {
|
||||
if (state.aborted) return
|
||||
if (plan.stallAfter !== undefined && index >= plan.stallAfter) {
|
||||
await new Promise<void>((resolve) => {
|
||||
wake = resolve
|
||||
})
|
||||
return
|
||||
}
|
||||
await new Promise<void>((resolve) => {
|
||||
wake = resolve
|
||||
timer = setTimeout(resolve, plan.intervalMs)
|
||||
})
|
||||
timer = null
|
||||
if (state.aborted) return
|
||||
if (index >= plan.chunks.length) {
|
||||
signal.removeEventListener('abort', onAbort)
|
||||
controller.close()
|
||||
return
|
||||
}
|
||||
controller.enqueue(encoder.encode(plan.chunks[index++]))
|
||||
},
|
||||
cancel() {
|
||||
if (timer !== null) clearTimeout(timer)
|
||||
signal.removeEventListener('abort', onAbort)
|
||||
},
|
||||
})
|
||||
return {
|
||||
response: new Response(body, { status: 200, headers: { 'Content-Type': 'text/event-stream' } }),
|
||||
state,
|
||||
}
|
||||
}
|
||||
|
||||
const SSE_CHUNKS = [
|
||||
'event: message_start\ndata: {"type":"message_start"}\n\n',
|
||||
'event: content_block_delta\ndata: {"type":"content_block_delta","delta":{"type":"text_delta","text":"1. "}}\n\n',
|
||||
'event: content_block_delta\ndata: {"type":"content_block_delta","delta":{"type":"text_delta","text":"item"}}\n\n',
|
||||
'event: content_block_delta\ndata: {"type":"content_block_delta","delta":{"type":"text_delta","text":" two"}}\n\n',
|
||||
'event: message_delta\ndata: {"type":"message_delta"}\n\n',
|
||||
'event: message_st',
|
||||
'op\ndata: {"type":"message_stop"}\n\n',
|
||||
]
|
||||
|
||||
function buildDeps(
|
||||
store: FakeLlmQuotaStore,
|
||||
fetchProvider: ProviderFetch,
|
||||
timeouts: LlmProxyDeps['timeouts'] = {},
|
||||
): LlmProxyDeps {
|
||||
let seq = 0
|
||||
return {
|
||||
authenticate: () => Promise.resolve({ id: 'user-1' }),
|
||||
providerKey: () => 'test-key',
|
||||
quotaStore: () => store,
|
||||
issueGenerationReceipt: () => Promise.resolve('00000000-0000-4000-8000-000000000001'),
|
||||
fetchProvider,
|
||||
newReservationId: () => `res-${++seq}`,
|
||||
timeouts,
|
||||
}
|
||||
}
|
||||
|
||||
function llmRequest(stream: boolean, signal?: AbortSignal): Request {
|
||||
return new Request('http://localhost/llm-proxy', {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({
|
||||
messages: [{ role: 'user', content: 'Write a detailed itemised summary.' }],
|
||||
max_tokens: 2048,
|
||||
stream,
|
||||
}),
|
||||
signal,
|
||||
})
|
||||
}
|
||||
|
||||
function okJson(): Response {
|
||||
return new Response(JSON.stringify({ content: [{ type: 'text', text: 'hello' }] }), {
|
||||
status: 200,
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
})
|
||||
}
|
||||
|
||||
Deno.test('stream longer than the first-byte deadline is relayed in full and charged once', async () => {
|
||||
const store = new FakeLlmQuotaStore(10)
|
||||
const upstreams: Array<{ aborted: boolean }> = []
|
||||
const fetchProvider: ProviderFetch = (_input, init) => {
|
||||
const { response, state } = slowSseResponse({ chunks: SSE_CHUNKS, intervalMs: 25 }, init.signal)
|
||||
upstreams.push(state)
|
||||
return Promise.resolve(response)
|
||||
}
|
||||
// 7 chunks x 25 ms ≈ 200 ms of streaming against a 40 ms first-byte deadline.
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider, {
|
||||
ttfbMs: 40,
|
||||
streamIdleMs: 500,
|
||||
streamTotalMs: 5_000,
|
||||
}))
|
||||
|
||||
const resp = await handler(llmRequest(true))
|
||||
assertEquals(resp.status, 200, 'status')
|
||||
const text = await resp.text()
|
||||
assertEquals(text, SSE_CHUNKS.join(''), 'relayed body')
|
||||
assert(upstreams.length === 1 && !upstreams[0].aborted, 'upstream must not be aborted')
|
||||
assertEquals(store.finalized, [{ id: 'res-1', succeeded: true }], 'finalized')
|
||||
assertEquals(store.used, 1, 'one unit spent')
|
||||
})
|
||||
|
||||
Deno.test('stream that stalls past the idle deadline is aborted and refunded', async () => {
|
||||
const store = new FakeLlmQuotaStore(10)
|
||||
const upstreams: Array<{ aborted: boolean }> = []
|
||||
const fetchProvider: ProviderFetch = (_input, init) => {
|
||||
const { response, state } = slowSseResponse(
|
||||
{ chunks: SSE_CHUNKS, intervalMs: 5, stallAfter: 2 },
|
||||
init.signal,
|
||||
)
|
||||
upstreams.push(state)
|
||||
return Promise.resolve(response)
|
||||
}
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider, {
|
||||
ttfbMs: 1_000,
|
||||
streamIdleMs: 60,
|
||||
streamTotalMs: 5_000,
|
||||
}))
|
||||
|
||||
const resp = await handler(llmRequest(true))
|
||||
assertEquals(resp.status, 200, 'status')
|
||||
let errored = false
|
||||
try {
|
||||
await resp.text()
|
||||
} catch (err) {
|
||||
errored = err instanceof DOMException && err.name === 'TimeoutError'
|
||||
}
|
||||
assert(errored, 'client stream must error with a timeout')
|
||||
assert(upstreams.length === 1 && upstreams[0].aborted, 'upstream must be aborted')
|
||||
assertEquals(store.finalized, [{ id: 'res-1', succeeded: false }], 'refunded')
|
||||
assertEquals(store.used, 0, 'unit returned')
|
||||
})
|
||||
|
||||
Deno.test('stream that exceeds the total deadline is aborted and refunded', async () => {
|
||||
const store = new FakeLlmQuotaStore(10)
|
||||
const many = Array.from({ length: 200 }, () => 'event: ping\ndata: {"type":"ping"}\n\n')
|
||||
const fetchProvider: ProviderFetch = (_input, init) =>
|
||||
Promise.resolve(slowSseResponse({ chunks: many, intervalMs: 10 }, init.signal).response)
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider, {
|
||||
ttfbMs: 1_000,
|
||||
streamIdleMs: 500,
|
||||
streamTotalMs: 120,
|
||||
}))
|
||||
|
||||
const resp = await handler(llmRequest(true))
|
||||
let errored = false
|
||||
try {
|
||||
await resp.text()
|
||||
} catch {
|
||||
errored = true
|
||||
}
|
||||
assert(errored, 'client stream must error')
|
||||
assertEquals(store.finalized, [{ id: 'res-1', succeeded: false }], 'refunded')
|
||||
})
|
||||
|
||||
Deno.test('stream that ends without message_stop is refunded', async () => {
|
||||
const store = new FakeLlmQuotaStore(10)
|
||||
const fetchProvider: ProviderFetch = (_input, init) =>
|
||||
Promise.resolve(slowSseResponse({ chunks: SSE_CHUNKS.slice(0, 3), intervalMs: 5 }, init.signal).response)
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider, { ttfbMs: 1_000 }))
|
||||
|
||||
const resp = await handler(llmRequest(true))
|
||||
await resp.text()
|
||||
assertEquals(store.finalized, [{ id: 'res-1', succeeded: false }], 'refunded')
|
||||
})
|
||||
|
||||
Deno.test('client disconnect cancels the upstream stream and keeps the charge', async () => {
|
||||
const store = new FakeLlmQuotaStore(10)
|
||||
const upstreams: Array<{ aborted: boolean }> = []
|
||||
const fetchProvider: ProviderFetch = (_input, init) => {
|
||||
const { response, state } = slowSseResponse({ chunks: SSE_CHUNKS, intervalMs: 20 }, init.signal)
|
||||
upstreams.push(state)
|
||||
return Promise.resolve(response)
|
||||
}
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider, { ttfbMs: 1_000 }))
|
||||
|
||||
const resp = await handler(llmRequest(true))
|
||||
const reader = resp.body!.getReader()
|
||||
await reader.read()
|
||||
await reader.cancel('client left')
|
||||
await delay(10)
|
||||
assert(upstreams.length === 1 && upstreams[0].aborted, 'upstream must be aborted')
|
||||
assertEquals(store.finalized, [{ id: 'res-1', succeeded: true }], 'charged')
|
||||
})
|
||||
|
||||
Deno.test('parallel requests past the allowance never reach the provider', async () => {
|
||||
const store = new FakeLlmQuotaStore(2)
|
||||
let providerCalls = 0
|
||||
let release: () => void = () => {}
|
||||
const gate = new Promise<void>((resolve) => {
|
||||
release = resolve
|
||||
})
|
||||
const fetchProvider: ProviderFetch = async () => {
|
||||
providerCalls++
|
||||
await gate
|
||||
return okJson()
|
||||
}
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider))
|
||||
|
||||
const pending = Array.from({ length: 6 }, () => handler(llmRequest(false)))
|
||||
await delay(10)
|
||||
// All reservations are decided while the first two provider calls are still in flight.
|
||||
assertEquals(providerCalls, 2, 'provider calls while in flight')
|
||||
release()
|
||||
const statuses = (await Promise.all(pending)).map((r) => r.status).sort()
|
||||
assertEquals(statuses, [200, 200, 429, 429, 429, 429], 'statuses')
|
||||
assertEquals(providerCalls, 2, 'total provider calls')
|
||||
assertEquals(store.used, 2, 'units spent')
|
||||
})
|
||||
|
||||
Deno.test('provider error refunds the reservation', async () => {
|
||||
const store = new FakeLlmQuotaStore(1)
|
||||
const fetchProvider: ProviderFetch = () =>
|
||||
Promise.resolve(new Response('overloaded', { status: 529 }))
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider))
|
||||
|
||||
const resp = await handler(llmRequest(false))
|
||||
assertEquals(resp.status, 502, 'status')
|
||||
await resp.body?.cancel()
|
||||
assertEquals(store.finalized, [{ id: 'res-1', succeeded: false }], 'refunded')
|
||||
assertEquals(store.used, 0, 'unit returned')
|
||||
})
|
||||
|
||||
Deno.test('first-byte timeout returns 504 and refunds the reservation', async () => {
|
||||
const store = new FakeLlmQuotaStore(1)
|
||||
const fetchProvider: ProviderFetch = (_input, init) =>
|
||||
new Promise<Response>((_resolve, reject) => {
|
||||
const signal = init.signal
|
||||
signal.addEventListener('abort', () => reject(signal.reason), { once: true })
|
||||
})
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider, { ttfbMs: 30 }))
|
||||
|
||||
const resp = await handler(llmRequest(false))
|
||||
assertEquals(resp.status, 504, 'status')
|
||||
assertEquals(await resp.json(), { error: 'provider_timeout' }, 'body')
|
||||
assertEquals(store.finalized, [{ id: 'res-1', succeeded: false }], 'refunded')
|
||||
})
|
||||
|
||||
Deno.test('denied reservation returns 429 without calling the provider', async () => {
|
||||
const store = new FakeLlmQuotaStore(0, 'free')
|
||||
let providerCalls = 0
|
||||
const fetchProvider: ProviderFetch = () => {
|
||||
providerCalls++
|
||||
return Promise.resolve(okJson())
|
||||
}
|
||||
const handler = createLlmProxyHandler(buildDeps(store, fetchProvider))
|
||||
|
||||
const resp = await handler(llmRequest(false))
|
||||
assertEquals(resp.status, 429, 'status')
|
||||
const body = await resp.json()
|
||||
assertEquals(body.error, 'quota_exceeded', 'error')
|
||||
assertEquals(body.model, 'claude-haiku-4-5-20251001', 'model')
|
||||
assertEquals(body.tier, 'free', 'tier')
|
||||
assertEquals(providerCalls, 0, 'provider calls')
|
||||
})
|
||||
|
||||
Deno.test('non-stream success settles the reservation as completed', async () => {
|
||||
const store = new FakeLlmQuotaStore(5)
|
||||
const handler = createLlmProxyHandler(buildDeps(store, () => Promise.resolve(okJson())))
|
||||
|
||||
const resp = await handler(llmRequest(false))
|
||||
assertEquals(resp.status, 200, 'status')
|
||||
assertEquals(await resp.json(), { content: [{ type: 'text', text: 'hello' }] }, 'body')
|
||||
assertEquals(store.finalized, [{ id: 'res-1', succeeded: true }], 'completed')
|
||||
})
|
||||
|
||||
Deno.test('SSE completion tracker finds message_stop split across chunks', () => {
|
||||
const tracker = createSseCompletionTracker()
|
||||
tracker.push(encoder.encode('event: message_st'))
|
||||
assert(!tracker.completed, 'not yet')
|
||||
tracker.push(encoder.encode('op\ndata: {}\n\n'))
|
||||
assert(tracker.completed, 'completed')
|
||||
})
|
||||
Loading…
Add table
Add a link
Reference in a new issue