fix(stream): isolate transport state by credential
This commit is contained in:
+3
-2
@@ -71,6 +71,7 @@ export function createCommandCodeTransportRouter(deps: TransportDependencies) {
|
|||||||
apiKey = options?.apiKey
|
apiKey = options?.apiKey
|
||||||
transport = "unknown"
|
transport = "unknown"
|
||||||
}
|
}
|
||||||
|
const requestApiKey = options?.apiKey
|
||||||
if (transport === "generate") return deps.streamGenerate(model, context, options)
|
if (transport === "generate") return deps.streamGenerate(model, context, options)
|
||||||
|
|
||||||
const output = deps.createStream()
|
const output = deps.createStream()
|
||||||
@@ -94,13 +95,13 @@ export function createCommandCodeTransportRouter(deps: TransportDependencies) {
|
|||||||
|
|
||||||
for await (const event of providerStream) {
|
for await (const event of providerStream) {
|
||||||
if (!upgradeRequired) {
|
if (!upgradeRequired) {
|
||||||
transport = "provider"
|
if (apiKey === requestApiKey) transport = "provider"
|
||||||
output.push(event)
|
output.push(event)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if (upgradeRequired) {
|
if (upgradeRequired) {
|
||||||
transport = "generate"
|
if (apiKey === requestApiKey) transport = "generate"
|
||||||
await pipe(deps.streamGenerate(model, context, options), output)
|
await pipe(deps.streamGenerate(model, context, options), output)
|
||||||
}
|
}
|
||||||
output.end()
|
output.end()
|
||||||
|
|||||||
@@ -159,6 +159,70 @@ describe("Command Code transport router", () => {
|
|||||||
assert.equal(generateCalls, 1)
|
assert.equal(generateCalls, 1)
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it("does not let a stale request overwrite the transport for a new API key", async () => {
|
||||||
|
let releaseGoRequest: (() => void) | undefined
|
||||||
|
const goRequestGate = new Promise<void>((resolve) => {
|
||||||
|
releaseGoRequest = resolve
|
||||||
|
})
|
||||||
|
let providerCalls = 0
|
||||||
|
let generateCalls = 0
|
||||||
|
const upgradeBody = JSON.stringify({ error: { code: "upgrade_required" } })
|
||||||
|
const router = createCommandCodeTransportRouter({
|
||||||
|
createStream: createTestEventStream,
|
||||||
|
streamProvider: (_model, _context, options) => {
|
||||||
|
providerCalls += 1
|
||||||
|
const response =
|
||||||
|
options?.apiKey === "go-key"
|
||||||
|
? new Response(upgradeBody, { status: 403 })
|
||||||
|
: new Response("ok", { status: 200 })
|
||||||
|
const stream = createTestEventStream()
|
||||||
|
const run = async () => {
|
||||||
|
if (options?.apiKey === "go-key") await goRequestGate
|
||||||
|
const received = await (options?.fetch ?? fetch)("https://provider.test", {})
|
||||||
|
await options?.onResponse?.(
|
||||||
|
{ status: received.status, headers: {} },
|
||||||
|
makeModel({ api: "openai-completions" }),
|
||||||
|
)
|
||||||
|
if (response.ok) {
|
||||||
|
for await (const event of completedStream("provider")) stream.push(event)
|
||||||
|
}
|
||||||
|
stream.end()
|
||||||
|
}
|
||||||
|
run().catch(() => stream.end())
|
||||||
|
return stream
|
||||||
|
},
|
||||||
|
streamGenerate: () => {
|
||||||
|
generateCalls += 1
|
||||||
|
return completedStream("generate")
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
const staleGoRequest = collectEvents(
|
||||||
|
router.stream(makeModel(), makeContext(), {
|
||||||
|
apiKey: "go-key",
|
||||||
|
fetch: () => Promise.resolve(new Response(upgradeBody, { status: 403 })),
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
await collectEvents(
|
||||||
|
router.stream(makeModel(), makeContext(), {
|
||||||
|
apiKey: "provider-key",
|
||||||
|
fetch: () => Promise.resolve(new Response("ok", { status: 200 })),
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
releaseGoRequest?.()
|
||||||
|
await staleGoRequest
|
||||||
|
await collectEvents(
|
||||||
|
router.stream(makeModel(), makeContext(), {
|
||||||
|
apiKey: "provider-key",
|
||||||
|
fetch: () => Promise.resolve(new Response("ok", { status: 200 })),
|
||||||
|
}),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert.equal(router.getTransport(), "provider")
|
||||||
|
assert.equal(providerCalls, 3)
|
||||||
|
assert.equal(generateCalls, 1)
|
||||||
|
})
|
||||||
|
|
||||||
it("does not fall back for other 403 errors", async () => {
|
it("does not fall back for other 403 errors", async () => {
|
||||||
let generateCalls = 0
|
let generateCalls = 0
|
||||||
const responseBody = JSON.stringify({ error: { code: "permission_denied" } })
|
const responseBody = JSON.stringify({ error: { code: "permission_denied" } })
|
||||||
|
|||||||
Reference in New Issue
Block a user