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
157 changes: 156 additions & 1 deletion packages/client/src/plugins/batch.test.ts
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import type { StandardLazyResponse, StandardRequest } from '@standard-server/core'
import type { StandardLinkCodec, StandardLinkTransport } from '../adapters/standard'
import { AsyncLocalStorage } from 'node:async_hooks'
import { sleep } from '@orpc/shared'
import { promiseWithResolvers, sleep } from '@orpc/shared'
import { encodePeerMessage } from '@standard-server/peer'
import { StandardLink } from '../adapters/standard'
import { BatchLinkPlugin } from './batch'
Expand Down Expand Up @@ -93,6 +93,37 @@ async function toLengthPrefixedBytes(messages: any[]): Promise<Uint8Array<ArrayB
return output
}

/**
* Sends calls `a` and `b` as one streaming batch whose response messages the test pushes one by one.
*/
async function startStreamingBatch() {
const batch = promiseWithResolvers<{ ids: string[], signal: AbortSignal, push: (message: any) => Promise<void> }>()

const link = new StandardLink(makeCodec(), {
send: async (request) => {
const stream = new ReadableStream<Uint8Array>({
start(controller) {
batch.resolve({
ids: extractBatchMessagesFromRequest(request).map(message => message.id),
signal: request.signal!,
push: async message => controller.enqueue(await toLengthPrefixedBytes([message])),
})
},
})

return { status: 207, headers: {}, resolveBody: async () => stream }
},
}, {
plugins: [new BatchLinkPlugin({ groups: [{ condition: () => true, context: {} }], mode: 'streaming' })],
})

const outputA = link.call(['a'], {}, { context: {} })
const outputB = link.call(['b'], {}, { context: {} })
const { ids: [idA, idB], signal, push } = await batch.promise

return { outputA, outputB, idA, idB, signal, push }
}

beforeEach(() => {
vi.clearAllMocks()
vi.useRealTimers()
Expand Down Expand Up @@ -720,6 +751,43 @@ describe('batchLinkPlugin', () => {
expect(transport.send).not.toHaveBeenCalled()
})

it('sends the batch with the cancel of a subrequest aborted after its request message but before the batch', async () => {
const codec = makeCodec()
const transport = makeTransport()
const controller = new AbortController()

const link = new StandardLink(codec, transport, {
plugins: [new BatchLinkPlugin({
groups: [defaultGroup],
mode: 'buffered',
mapSubrequest: ({ request }) => {
if (request.url !== '/b') {
return request
}

// Read when the peer starts sending b, after a's request message
return {
...request,
get body() {
controller.abort(new Error('TEST_ABORT'))
return request.body
},
}
},
})],
})

const abortedPromise = expect(link.call(['a'], {}, { context: {}, signal: controller.signal })).rejects.toThrow('TEST_ABORT')
const promise2 = link.call(['b'], {}, { context: {} })

await abortedPromise
await expect(promise2).resolves.toBe('result-2')

const batchRequest = vi.mocked(transport.send).mock.calls[0]![0]
expect(batchRequest.signal?.aborted).toBe(false)
expect(extractBatchMessagesFromRequest(batchRequest).map(m => m.kind)).toEqual(['request', 'cancel', 'request'])
})

it('aborts the batch request once every subrequest is aborted, including ones aborted before sending', async () => {
const codec = makeCodec()
const transport = makeTransport()
Expand Down Expand Up @@ -761,6 +829,93 @@ describe('batchLinkPlugin', () => {

await promise
})

it('aborts the batch request when a cancelled stream is all the server is still running', async () => {
const { outputA, outputB, idA, idB, signal, push } = await startStreamingBatch()

await push({ kind: 'response', id: idA, json: { body: 'a' } })
await push({ kind: 'response', id: idB, json: { headers: { 'standard-server': 'event-stream' } } })
await expect(outputA).resolves.toBe('a')
const iteratorB = await outputB as AsyncIteratorObject<unknown>

await iteratorB.return?.()
expect(signal.aborted).toBe(true)
})

it('keeps the batch request open after a cancel until every other stream finishes', async () => {
const { outputA, outputB, idA, idB, signal, push } = await startStreamingBatch()

await push({ kind: 'response', id: idA, json: { headers: { 'standard-server': 'event-stream' } } })
await push({ kind: 'response', id: idB, json: { headers: { 'content-type': 'application/octet-stream' } } })
const iteratorA = await outputA as AsyncIteratorObject<unknown>
const streamB = await outputB as ReadableStream<Uint8Array>

await streamB.cancel()
await push({ kind: 'event-stream', id: idA, json: { data: 'a1' } })
await expect(iteratorA.next()).resolves.toEqual({ value: 'a1', done: false })
expect(signal.aborted).toBe(false)

await push({ kind: 'event-stream', id: idA, json: { event: 'close' } })
await expect(iteratorA.next()).resolves.toEqual({ value: undefined, done: true })
await sleep(0)
expect(signal.aborted).toBe(true)
})

it('keeps the batch request open when a subrequest is cancelled only after the server finished it', async () => {
const { outputA, outputB, idA, idB, signal, push } = await startStreamingBatch()

await push({ kind: 'response', id: idA, json: { headers: { 'standard-server': 'event-stream' } } })
await push({ kind: 'response', id: idB, json: { headers: { 'standard-server': 'event-stream' } } })
const iteratorA = await outputA as AsyncIteratorObject<unknown>
const iteratorB = await outputB as AsyncIteratorObject<unknown>

// B's event arriving proves A's close was received, though A never reads it
await push({ kind: 'event-stream', id: idA, json: { event: 'close' } })
await push({ kind: 'event-stream', id: idB, json: { data: 'b1' } })
await expect(iteratorB.next()).resolves.toEqual({ value: 'b1', done: false })

await iteratorA.return?.()
await push({ kind: 'event-stream', id: idB, json: { event: 'close' } })
await expect(iteratorB.next()).resolves.toEqual({ value: undefined, done: true })
await sleep(0)
expect(signal.aborted).toBe(false)
})

it('keeps the batch request open when a cancelled stream also finishes on the server before the others', async () => {
const { outputA, outputB, idA, idB, signal, push } = await startStreamingBatch()

await push({ kind: 'response', id: idA, json: { headers: { 'standard-server': 'event-stream' } } })
await push({ kind: 'response', id: idB, json: { headers: { 'standard-server': 'event-stream' } } })
const iteratorA = await outputA as AsyncIteratorObject<unknown>
const iteratorB = await outputB as AsyncIteratorObject<unknown>

await iteratorA.return?.()
await push({ kind: 'event-stream', id: idA, json: { event: 'close' } })
await push({ kind: 'event-stream', id: idB, json: { event: 'close' } })
await expect(iteratorB.next()).resolves.toEqual({ value: undefined, done: true })
await sleep(0)
expect(signal.aborted).toBe(false)
})

it('treats a server cancel as the end of a subrequest, but not a stream/cancel', async () => {
const { outputA, outputB, idA, idB, signal, push } = await startStreamingBatch()

await push({ kind: 'response', id: idA, json: { headers: { 'standard-server': 'event-stream' } } })
await push({ kind: 'response', id: idB, json: { headers: { 'standard-server': 'event-stream' } } })
const iteratorA = await outputA as AsyncIteratorObject<unknown>
const iteratorB = await outputB as AsyncIteratorObject<unknown>

await iteratorB.return?.()
await push({ kind: 'stream/cancel', id: idA })
await push({ kind: 'event-stream', id: idA, json: { data: 'a1' } })
await expect(iteratorA.next()).resolves.toEqual({ value: 'a1', done: false })
expect(signal.aborted).toBe(false)

await push({ kind: 'cancel', id: idA })
await expect(iteratorA.next()).rejects.toThrow('Server canceled the request')
await sleep(0)
expect(signal.aborted).toBe(true)
})
})

describe('batch response decoding', () => {
Expand Down
70 changes: 54 additions & 16 deletions packages/client/src/plugins/batch.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
import type { InterceptorOptions, Promisable, Value } from '@orpc/shared'
import type { StandardHeaders, StandardLazyResponse, StandardRequest, StandardUrl } from '@standard-server/core'
import type { ClientPeerSendMessage } from '@standard-server/peer'
import type { ClientPeerSendMessage, ServerPeerSendMessage } from '@standard-server/peer'
import type { StandardLinkOptions, StandardLinkPlugin, StandardLinkTransportInterceptor, StandardLinkTransportInterceptorOptions } from '../adapters/standard'
import type { ClientContext } from '../types'
import { captureAsyncContext, defer, isAsyncIteratorObject, loadBytes, once, promiseWithResolvers, safeEncodeURIComponent, splitInHalf, stringifyJSON, toArray, value } from '@orpc/shared'
Expand Down Expand Up @@ -307,29 +307,38 @@ export class BatchLinkPlugin<T extends ClientContext> implements StandardLinkPlu
const pendingMessages: ClientPeerSendMessage[] = []
const subrequests = groupItems.map(([subOptions]) => this.mapSubrequest(subOptions, { url, headers }))
let batchResponse: StandardLazyResponse
let activeCount = subrequests.length
let markRequestSent: (() => void) | undefined
const openRequestIds = new Set<string>()
const cancelledRunningRequestIds = new Set<string>()
let isBatchSent = false

const abortIfOnlyCancelledRemain = () => {
if (openRequestIds.size === 0 && cancelledRunningRequestIds.size > 0) {
controller.abort()
}
}

const peer = new ClientPeer(async (message) => {
pendingMessages.push(message)

if (message.kind === 'request') {
openRequestIds.add(message.id)
markRequestSent?.()
}

if (message.kind === 'cancel' && --activeCount === 0) {
controller.abort()
else if (message.kind === 'cancel' && openRequestIds.delete(message.id) && isBatchSent) {
cancelledRunningRequestIds.add(message.id)
abortIfOnlyCancelledRemain()
}
})

/**
* Subrequests go to the peer one at a time, so a request message always belongs to the current one.
* A subrequest aborted before the peer starts sending its request message sends nothing, not even a
* cancel, so one that settles without a request message is no longer active.
* cancel, so settling also ends the wait for it.
*/
for (const [index, [subOptions, resolve, reject]] of groupItems.entries()) {
const sent = promiseWithResolvers<boolean>()
markRequestSent = () => sent.resolve(true)
const sent = promiseWithResolvers<void>()
markRequestSent = sent.resolve

peer
.request(subrequests[index]!)
Expand All @@ -339,17 +348,17 @@ export class BatchLinkPlugin<T extends ClientContext> implements StandardLinkPlu
reject(error)
}
})
.then(() => sent.resolve(false))
.then(sent.resolve)

if (!await sent.promise) {
activeCount--
}
await sent.promise
}

if (activeCount === 0) {
if (openRequestIds.size === 0) {
return
}

isBatchSent = true

try {
const request: StandardRequest = {
url,
Expand Down Expand Up @@ -422,7 +431,15 @@ export class BatchLinkPlugin<T extends ClientContext> implements StandardLinkPlu
await decodeLengthPrefixedBlob(body, peer)
}
else if (body instanceof ReadableStream) {
await decodeLengthPrefixedStream(body, peer)
await decodeLengthPrefixedStream(body, async (message) => {
await peer.message(message)

if (isLastServerMessage(message)) {
openRequestIds.delete(message.id)
cancelledRunningRequestIds.delete(message.id)
abortIfOnlyCancelledRemain()
}
})
}
else {
throw new TypeError('Invalid batch response format.')
Expand All @@ -447,6 +464,27 @@ type BatchLinkPluginItem<T extends ClientContext> = [
runInOwnContext: ReturnType<typeof captureAsyncContext>,
]

/**
* Whether the server sends nothing more for the subrequest after this message.
*/
function isLastServerMessage(message: ServerPeerSendMessage): boolean {
switch (message.kind) {
case 'response':
// A body-less response with a content type or body hint is followed by stream messages
return message.json.body !== undefined
|| message.binary !== undefined
|| (message.json.headers?.['content-type'] === undefined && message.json.headers?.['standard-server'] === undefined)
case 'event-stream':
return message.json.event === 'close' || message.json.event === 'error'
case 'octet-stream':
return message.json.close === true
case 'cancel':
return true
default:
return false
}
}

async function decodeLengthPrefixedBlob(blob: Blob, peer: ClientPeer): Promise<void> {
const buffer = await loadBytes(blob)
let offset = 0
Expand Down Expand Up @@ -476,7 +514,7 @@ async function decodeLengthPrefixedBlob(blob: Blob, peer: ClientPeer): Promise<v
}
}

async function decodeLengthPrefixedStream(stream: ReadableStream<Uint8Array>, peer: ClientPeer): Promise<void> {
async function decodeLengthPrefixedStream(stream: ReadableStream<Uint8Array>, receive: (message: ServerPeerSendMessage) => Promise<void>): Promise<void> {
const reader = stream.getReader()
let buffer = new Uint8Array(0)

Expand Down Expand Up @@ -514,7 +552,7 @@ async function decodeLengthPrefixedStream(stream: ReadableStream<Uint8Array>, pe
throw new TypeError('Invalid batch response: invalid message.')
}

await peer.message(result.message)
await receive(result.message)
}

if (done) {
Expand Down
56 changes: 56 additions & 0 deletions tests/batch/batch-plugin.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -226,6 +226,62 @@ describe.each([

expect(fetchSpy).toHaveBeenCalledTimes(2) // the upload plus one batch for both echoes
})

it('stops a cancelled stream once every other subrequest in the batch has finished', async () => {
const release = promiseWithResolvers<void>()
const stopped = promiseWithResolvers<void>()

const router = {
json: os.handler(() => 'ok'),
blob: os.handler(() => new Blob(['blob'])),
stream: os.handler(() => new ReadableStream<Uint8Array>({
async start(controller) {
await release.promise
controller.enqueue(new TextEncoder().encode('stream'))
controller.close()
},
})),
events: os.handler(async function* () {
await release.promise
yield 'event'
}),
endless: os.handler(async function* () {
try {
for (let i = 0; ; i++) {
yield i
await sleep(10)
}
}
finally {
stopped.resolve()
}
}),
}

const { client, fetchSpy } = createClientServer(router, { mode: 'streaming' })

const [info, file, stream, events, endless] = await Promise.all([
client.json(),
client.blob(),
client.stream() as Promise<ReadableStream<Uint8Array>>,
client.events() as Promise<AsyncIteratorObject<unknown>>,
client.endless() as Promise<AsyncIteratorObject<unknown>>,
])

expect(info).toBe('ok')
await expect((file as Blob).text()).resolves.toBe('blob')
await expect(endless.next()).resolves.toEqual({ value: 0, done: false })
await endless.return?.()

// The still open streams keep the batch alive, then only the cancelled one is left running.
release.resolve()
await expect(new Response(stream).text()).resolves.toBe('stream')
await expect(events.next()).resolves.toEqual({ value: 'event', done: false })
await expect(events.next()).resolves.toEqual({ value: undefined, done: true })

await stopped.promise
expect(fetchSpy).toHaveBeenCalledTimes(1) // ensure batch was used
})
})

describe('batch plugin: QUERY over node-http', () => {
Expand Down
Loading