diff --git a/packages/peer/src/client.test.ts b/packages/peer/src/client.test.ts index 1c710e9..cde6322 100644 --- a/packages/peer/src/client.test.ts +++ b/packages/peer/src/client.test.ts @@ -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(async () => { @@ -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({ + 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>({ diff --git a/packages/peer/src/client.ts b/packages/peer/src/client.ts index a93cbbe..2bf7c95 100644 --- a/packages/peer/src/client.ts +++ b/packages/peer/src/client.ts @@ -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 { @@ -104,6 +105,10 @@ export class ClientPeer { return } + if (state.streamCancelled) { + return + } + untransmittedBody = undefined if (isAsyncIteratorObject(request.body)) { @@ -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(), diff --git a/packages/peer/tests/peer.test.ts b/packages/peer/tests/peer.test.ts index 04f8131..b060690 100644 --- a/packages/peer/tests/peer.test.ts +++ b/packages/peer/tests/peer.test.ts @@ -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): { client: ClientPeer, server: ServerPeer } { +function connect( + handler: (request: StandardLazyRequest) => Promise, + { waitForRemote = false } = {}, +): { client: ClientPeer, server: ServerPeer } { const prefix = 'peer:' const wire = {} as { client: ClientPeer, server: ServerPeer } @@ -19,7 +25,13 @@ function connect(handler: (request: StandardLazyRequest) => Promise {}) + const handled = wire.server.message(decoded.message as ClientPeerSendMessage, handler) + if (waitForRemote) { + await handled + } + else { + void handled.catch(() => {}) + } }) wire.server = new ServerPeerClass(async (message) => { @@ -27,7 +39,10 @@ function connect(handler: (request: StandardLazyRequest) => Promise 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 => ({ + 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' }) }) @@ -323,6 +317,27 @@ describe('peer integration (client <-> server over encoded wire)', () => { expect(await readAll(body as ReadableStream)).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({ + 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