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
49 changes: 49 additions & 0 deletions packages/peer/src/client.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -482,6 +482,28 @@ describe('clientPeer', () => {
await promise
})

it('does not transmit the event-stream request body when stream/cancel arrives while the request is being sent', async () => {
const iter = makeAsyncIter(['a', 'b'])
const returnSpy = vi.spyOn(iter, 'return')

send.mockImplementation(async (message) => {
if (message.kind === 'request') {
await peer.message(makeStreamCancelMessage(message.id))
}
})

const { id, promise } = await requestAndGetId(
makeRequest({ method: 'POST', headers: {}, body: iter }),
)

await vi.waitFor(() => expect(returnSpy).toHaveBeenCalled())
expect(send.mock.calls.map(([m]) => m.kind)).toEqual(['request'])

await peer.message(makeResponseMessage(id, 'ok'))
const response = await promise
expect(await response.resolveBody()).toBe('ok')
})

it('send cancel message and reject request on non-protocol error', async () => {
const nonProtocolError = new Error('non-protocol error')
const iter = new AsyncIteratorClass<unknown>(async () => {
Expand Down Expand Up @@ -709,6 +731,33 @@ describe('clientPeer', () => {
expect(cancel).toHaveBeenCalled()
})

it('does not transmit the octet-stream request body when stream/cancel arrives while the request is being sent', async () => {
const cancel = vi.fn()
const stream = new ReadableStream<Uint8Array>({
start(controller) {
controller.enqueue(new Uint8Array([1, 2]))
},
cancel,
})

send.mockImplementation(async (message) => {
if (message.kind === 'request') {
await peer.message(makeStreamCancelMessage(message.id))
}
})

const { id, promise } = await requestAndGetId(
makeRequest({ method: 'POST', headers: {}, body: stream }),
)

await vi.waitFor(() => expect(cancel).toHaveBeenCalled())
expect(send.mock.calls.map(([m]) => m.kind)).toEqual(['request'])

await peer.message(makeResponseMessage(id, 'ok'))
const response = await promise
expect(await response.resolveBody()).toBe('ok')
})

it('send cancel message and reject request on stream error', async () => {
const error = new Error('stream error')
const stream = new ReadableStream<Uint8Array<ArrayBuffer>>({
Expand Down
6 changes: 6 additions & 0 deletions packages/peer/src/client.ts
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ interface ClientPeerRequestStateInternal {
* so until the request message is sent, transmitRequest sends the cancel instead of abortById.
*/
requestSent?: boolean | undefined
streamCancelled?: boolean | undefined
}

export class ClientPeer {
Expand Down Expand Up @@ -104,6 +105,10 @@ export class ClientPeer {
return
}

if (state.streamCancelled) {
return
}

untransmittedBody = undefined

if (isAsyncIteratorObject(request.body)) {
Expand Down Expand Up @@ -150,6 +155,7 @@ export class ClientPeer {
}

if (message.kind === 'stream/cancel') {
state.streamCancelled = true
const promise = Promise.all([
state.eventStreamTransmitter?.cancel(),
state.octetStreamTransmitter?.cancel(),
Expand Down
69 changes: 42 additions & 27 deletions packages/peer/tests/peer.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,14 @@ import { ClientPeer as ClientPeerClass, decodePeerMessage, encodePeerMessage, Se
* Wires a ClientPeer and a ServerPeer together through the real codec,
* simulating a full-duplex connection (e.g. a WebSocket) where every
* message crosses the wire encoded.
*
* With `waitForRemote`, each send resolves only after the remote side handled the message,
* so the remote's replies arrive while the send is still in flight.
*/
function connect(handler: (request: StandardLazyRequest) => Promise<StandardResponse>): { client: ClientPeer, server: ServerPeer } {
function connect(
handler: (request: StandardLazyRequest) => Promise<StandardResponse>,
{ waitForRemote = false } = {},
): { client: ClientPeer, server: ServerPeer } {
const prefix = 'peer:'
const wire = {} as { client: ClientPeer, server: ServerPeer }

Expand All @@ -19,15 +25,24 @@ function connect(handler: (request: StandardLazyRequest) => Promise<StandardResp
if (!decoded.matched) {
throw new Error('Failed to decode message on the wire')
}
void wire.server.message(decoded.message as ClientPeerSendMessage, handler).catch(() => {})
const handled = wire.server.message(decoded.message as ClientPeerSendMessage, handler)
if (waitForRemote) {
await handled
}
else {
void handled.catch(() => {})
}
})

wire.server = new ServerPeerClass(async (message) => {
const decoded = decodePeerMessage(await encodePeerMessage(message, { prefix }), { prefix })
if (!decoded.matched) {
throw new Error('Failed to decode message on the wire')
}
void wire.client.message(decoded.message as ServerPeerSendMessage)
const handled = wire.client.message(decoded.message as ServerPeerSendMessage)
if (waitForRemote) {
await handled
}
})

return wire
Expand All @@ -54,34 +69,13 @@ describe('peer integration (client <-> server over encoded wire)', () => {
})

it('completes when the transport waits for full remote processing before send resolves', async () => {
const prefix = 'peer:'
const wire = {} as { client: ClientPeer, server: ServerPeer }

const handler = async (request: StandardLazyRequest): Promise<StandardResponse> => ({
const { client } = connect(async request => ({
status: 200,
headers: {},
body: { pong: request.url },
})

// unlike `connect()`, each send awaits the remote side handling the message,
// so the response arrives while the request `send` is still in flight
wire.client = new ClientPeerClass(async (message) => {
const decoded = decodePeerMessage(await encodePeerMessage(message, { prefix }), { prefix })
if (!decoded.matched) {
throw new Error('Failed to decode message on the wire')
}
await wire.server.message(decoded.message as ClientPeerSendMessage, handler)
})

wire.server = new ServerPeerClass(async (message) => {
const decoded = decodePeerMessage(await encodePeerMessage(message, { prefix }), { prefix })
if (!decoded.matched) {
throw new Error('Failed to decode message on the wire')
}
await wire.client.message(decoded.message as ServerPeerSendMessage)
})
}), { waitForRemote: true })

const response = await wire.client.request({ url: '/ping', method: 'GET', headers: {} })
const response = await client.request({ url: '/ping', method: 'GET', headers: {} })
expect(response.status).toBe(200)
expect(await response.resolveBody()).toEqual({ pong: '/ping' })
})
Expand Down Expand Up @@ -323,6 +317,27 @@ describe('peer integration (client <-> server over encoded wire)', () => {
expect(await readAll(body as ReadableStream<Uint8Array>)).toEqual([4, 5])
})

it('does not upload a request body the server cancelled while the request message was still being sent', async () => {
const { client } = connect(async (request) => {
await (await request.resolveBody() as ReadableStream).cancel()
// a streamed response keeps the request open after the request `send` resolves
return { status: 200, headers: {}, body: (async function* () {})() }
}, { waitForRemote: true })

const cancel = vi.fn()
await client.request({
url: '/upload',
method: 'POST',
headers: {},
body: new ReadableStream<Uint8Array>({
start: controller => controller.enqueue(new Uint8Array([1, 2])),
cancel,
}),
})

await vi.waitFor(() => expect(cancel).toHaveBeenCalled())
})

it('propagates client aborts to the server handler signal', async () => {
let serverSignal: AbortSignal | undefined

Expand Down
Loading