diff --git a/src/rules/requests/request-rule.ts b/src/rules/requests/request-rule.ts index e32c8ad7a..ae0d5dde2 100644 --- a/src/rules/requests/request-rule.ts +++ b/src/rules/requests/request-rule.ts @@ -79,6 +79,7 @@ export class RequestRule implements RequestRule { handle(req: OngoingRequest, res: OngoingResponse, options: { record?: boolean, + bufferRequestBody?: boolean, debug: boolean, keyLogStream?: Writable, emitEventCallback?: (type: string, event: unknown) => void @@ -88,6 +89,7 @@ export class RequestRule implements RequestRule { const result = await step.handle(req, res, { emitEventCallback: options.emitEventCallback, keyLogStream: options.keyLogStream, + bufferRequestBody: options.record || options.bufferRequestBody, debug: options.debug }); diff --git a/src/rules/requests/request-step-impls.ts b/src/rules/requests/request-step-impls.ts index 4ac122875..c59f9d49d 100644 --- a/src/rules/requests/request-step-impls.ts +++ b/src/rules/requests/request-step-impls.ts @@ -29,6 +29,7 @@ import { AbortError } from '../../util/abort-error'; import { isAbsoluteUrl, getEffectivePort } from '../../util/url'; import { waitForCompletedRequest, + streamBodyWithoutBuffering, buildBodyReader, isHttp2, writeHead, @@ -164,6 +165,7 @@ export interface RequestStepOptions { emitEventCallback?: (type: string, event: unknown) => void; keyLogStream?: Writable; debug: boolean; + bufferRequestBody?: boolean; } export class FixedResponseStepImpl extends FixedResponseStep { @@ -484,9 +486,19 @@ export class PassThroughStepImpl extends PassThroughStep { `); } - // We have to capture the request stream immediately, to make sure nothing is lost if it - // goes past its max length (truncating the data) before we start sending upstream. - const clientReqBody = clientReq.body.asStream(); + const needsRequestBody = !!( + options.bufferRequestBody || + this.beforeRequest || + this.beforeResponse || + this.transformRequest?.updateJsonBody || + this.transformRequest?.patchJsonBody || + this.transformRequest?.matchReplaceBody + ); + // Capture replay streams before async setup so truncation cannot discard + // bytes before forwarding starts. Otherwise retain normal stream backpressure. + const clientReqBody = needsRequestBody + ? clientReq.body.asStream() + : streamBodyWithoutBuffering(clientReq.body); const isH2Downstream = isHttp2(clientReq); diff --git a/src/server/mockttp-server.ts b/src/server/mockttp-server.ts index 693d82212..086397178 100644 --- a/src/server/mockttp-server.ts +++ b/src/server/mockttp-server.ts @@ -768,6 +768,8 @@ export class MockttpServer extends AbstractMockttp implements Mockttp { if (this.debug) console.log(`Request matched rule: ${nextRule.explain()}`); await nextRule.handle(request, response, { record: this.recordTraffic, + bufferRequestBody: emitter.listenerCount('request') > 0 || + emitter.listenerCount('request-body-data') > 0, debug: this.debug, keyLogStream: this.keyLogStream, emitEventCallback: (emitter.listenerCount('rule-event') !== 0) diff --git a/src/util/request-utils.ts b/src/util/request-utils.ts index ba3b9b328..dd9ee3318 100644 --- a/src/util/request-utils.ts +++ b/src/util/request-utils.ts @@ -142,11 +142,24 @@ export async function decodeBodyBuffer(buffer: Buffer, headers: Headers) { ) } +const unbufferedBodyReaders = new WeakMap stream.Readable>(); + // Parse an in-progress request or response stream, i.e. where the body or possibly even the headers have // not been fully received/sent yet. const parseBodyStream = (bodyStream: stream.Readable, maxSize: number, getHeaders: () => Headers): OngoingBody => { let bufferPromise: BufferInProgress | null = null; let completedBuffer: Buffer | null = null; + let streamTaken = false; + + const capture = (): BufferInProgress => { + if (!bufferPromise) { + bufferPromise = streamToBuffer(bodyStream, maxSize); + bufferPromise + .then((buffer) => completedBuffer = buffer) + .catch(() => {}); // If we get no body, completedBuffer stays null + } + return bufferPromise; + }; let body = { // Returns a stream for the full body, not the live streaming body. @@ -154,24 +167,19 @@ const parseBodyStream = (bodyStream: stream.Readable, maxSize: number, getHeader // and buffered data, and then continues with the live stream, if active. // Listeners to this stream *must* be attached synchronously after this call. asStream() { + if (streamTaken) throw new Error('Cannot replay an unbuffered body stream'); // If we've already buffered the whole body, just stream it out: if (completedBuffer) return bufferToStream(completedBuffer); // Otherwise, we want to start buffering now, and wrap that with // a stream that can live-stream the buffered data on demand: - const buffer = body.asBuffer(); + const buffer = capture(); buffer.catch(() => {}); // Errors will be handled via the stream, so silence unhandled rejections here. return bufferThenStream(buffer, bodyStream); }, asBuffer() { - if (!bufferPromise) { - bufferPromise = streamToBuffer(bodyStream, maxSize); - - bufferPromise - .then((buffer) => completedBuffer = buffer) - .catch(() => {}); // If we get no body, completedBuffer stays null - } - return bufferPromise; + if (streamTaken) return Promise.reject(new Error('Cannot replay an unbuffered body stream')); + return capture(); }, async asDecodedBuffer() { const buffer = await body.asBuffer(); @@ -188,9 +196,24 @@ const parseBodyStream = (bodyStream: stream.Readable, maxSize: number, getHeader }, }; + unbufferedBodyReaders.set(body, () => { + // Matchers or earlier steps may already have read some of the body. + // Reuse their replay stream instead of losing that prefix. + if (bufferPromise) return body.asStream(); + if (streamTaken) throw new Error('Cannot replay an unbuffered body stream'); + streamTaken = true; + return bodyStream; + }); + return body; } +/** @internal Consume a body once unless an existing reader requires replay. */ +export function streamBodyWithoutBuffering(body: OngoingBody): stream.Readable { + const read = unbufferedBodyReaders.get(body); + return read ? read() : body.asStream(); +} + async function runAsyncOrUndefined(func: () => Promise): Promise { try { return await func(); diff --git a/test/integration/proxying/request-buffering.spec.ts b/test/integration/proxying/request-buffering.spec.ts new file mode 100644 index 000000000..87929716a --- /dev/null +++ b/test/integration/proxying/request-buffering.spec.ts @@ -0,0 +1,241 @@ +import { Buffer } from 'buffer'; +import * as http from 'http'; +import * as http2 from 'http2'; + +import { getLocal, Mockttp, MockttpOptions, CompletedRequest, BodyData } from '../../..'; +import { expect, nodeOnly, getDeferred, Deferred, makeDestroyable, DestroyableServer } from '../../test-utils'; +import type { BufferInProgress } from '../../../src/util/buffer-utils'; + +nodeOnly(() => { + // Package-root imports run the built server, so instrument its body reader. + const bufferUtils: typeof import('../../../src/util/buffer-utils') = require('../../../dist/util/buffer-utils'); + describe("Passthrough request buffering", () => { + let proxy: Mockttp | undefined; + let target: DestroyableServer; + let targetUrl: string; + let received: Deferred<{ body: Buffer, trailers: http.IncomingHttpHeaders }>; + let releaseResponse: Deferred; + let firstChunk: Deferred; + let upstreamAborted: Deferred; + let captures: BufferInProgress[]; + let bufferDescriptor: PropertyDescriptor; + + beforeEach(async () => { + received = getDeferred(); + releaseResponse = getDeferred(); + firstChunk = getDeferred(); + upstreamAborted = getDeferred(); + captures = []; + bufferDescriptor = Object.getOwnPropertyDescriptor(bufferUtils, 'streamToBuffer')!; + const original = bufferUtils.streamToBuffer; + Object.defineProperty(bufferUtils, 'streamToBuffer', { + ...bufferDescriptor, + value: (...args: Parameters) => { + const result = original(...args); + if ((args[0] as http.IncomingMessage).method === 'POST') captures.push(result); + return result; + } + }); + target = makeDestroyable(http.createServer((request, response) => { + const chunks: Buffer[] = []; + request.on('data', chunk => { + chunks.push(chunk); + firstChunk.resolve(); + }); + request.on('aborted', () => upstreamAborted.resolve()); + request.on('end', () => { + received.resolve({ body: Buffer.concat(chunks), trailers: request.trailers }); + releaseResponse.then(() => response.end('ok')); + }); + })); + await new Promise(resolve => target.listen(0, '127.0.0.1', resolve)); + targetUrl = `http://127.0.0.1:${(target.address() as { port: number }).port}`; + }); + + afterEach(async () => { + releaseResponse.resolve(); + await proxy?.stop(); + proxy = undefined; + await target.destroy(); + Object.defineProperty(bufferUtils, 'streamToBuffer', bufferDescriptor); + }); + + async function start(options: MockttpOptions = { recordTraffic: false }) { + proxy = getLocal(options); + await proxy.start(); + return proxy; + } + + function upload(body: Buffer | string, trailers?: Record) { + const finished = getDeferred(); + const request = http.request(proxy!.urlFor('/upload'), { + method: 'POST', + path: `${targetUrl}/upload`, + headers: { + Host: new URL(targetUrl).host, + ...(trailers ? { Trailer: Object.keys(trailers).join(', ') } : {}) + } + }, response => { + response.resume(); + response.on('end', () => finished.resolve()); + response.on('error', finished.reject); + }); + request.on('error', finished.reject); + request.write(body); + if (trailers) request.addTrailers(trailers); + request.end(); + // A failed assertion can close the proxy before this request finishes. + finished.catch(() => {}); + return finished; + } + + const retainedBytes = () => captures.reduce((total, buffer) => total + + buffer.currentChunks.reduce((size, chunk) => size + chunk.length, 0), 0); + + it("forwards an unobserved upload without retaining its body", async () => { + const server = await start(); + const endpoint = await server.forAnyRequest().thenForwardTo(targetUrl); + const body = Buffer.alloc(1024 * 1024, 'a'); + const finished = upload(body, { 'X-Upload-End': 'complete' }); + const actual = await received; + + expect(actual.body).to.deep.equal(body); + expect(actual.trailers).to.deep.equal({ 'x-upload-end': 'complete' }); + expect(retainedBytes()).to.equal(0); + releaseResponse.resolve(); + await finished; + expect(await endpoint.getSeenRequests()).to.deep.equal([]); + }); + + it("preserves bodies for traffic recording", async () => { + const server = await start({ recordTraffic: true }); + const endpoint = await server.forAnyRequest().thenForwardTo(targetUrl); + const finished = upload('recorded upload'); + expect((await received).body.toString()).to.equal('recorded upload'); + releaseResponse.resolve(); + await finished; + const requests = await endpoint.getSeenRequests(); + expect(await requests[0].body.getText()).to.equal('recorded upload'); + }); + + it("streams HTTP/2 uploads without retaining a replay buffer", async () => { + const server = await start({ + recordTraffic: false, + http2: true, + https: { + keyPath: './test/fixtures/test-ca.key', + certPath: './test/fixtures/test-ca.pem' + } + }); + await server.forAnyRequest().thenForwardTo(targetUrl); + const client = http2.connect(server.url); + try { + const request = client.request({ ':method': 'POST', ':path': '/upload' }); + const finished = getDeferred(); + request.on('error', finished.reject); + request.on('end', () => finished.resolve()); + finished.catch(() => {}); + request.resume(); + const body = Buffer.alloc(256 * 1024, 'h'); + request.end(body); + expect((await received).body).to.deep.equal(body); + expect(retainedBytes()).to.equal(0); + releaseResponse.resolve(); + await finished; + } finally { + client.destroy(); + } + }); + + it("cancels an unbuffered upstream upload when the client disconnects", async () => { + const server = await start(); + await server.forAnyRequest().thenForwardTo(targetUrl); + const request = http.request(server.urlFor('/upload'), { method: 'POST' }); + request.on('error', () => {}); + request.write(Buffer.alloc(64 * 1024, 'a')); + await firstChunk; + request.destroy(); + await upstreamAborted; + }); + + it("preserves bodies for complete request subscriptions", async () => { + const server = await start(); + const observed = getDeferred(); + await server.on('request', request => observed.resolve(request)); + await server.forAnyRequest().thenForwardTo(targetUrl); + const finished = upload('observed upload'); + expect((await received).body.toString()).to.equal('observed upload'); + expect(await (await observed).body.getText()).to.equal('observed upload'); + releaseResponse.resolve(); + await finished; + }); + + it("preserves streaming request data subscriptions", async () => { + const server = await start(); + const events: BodyData[] = []; + const ended = getDeferred(); + await server.on('request-body-data', event => { + events.push(event); + if (event.isEnded) ended.resolve(); + }); + await server.forAnyRequest().thenForwardTo(targetUrl); + const body = Buffer.alloc(256 * 1024, 'b'); + const finished = upload(body); + expect((await received).body).to.deep.equal(body); + await ended; + expect(Buffer.concat(events.map(event => Buffer.from(event.content)))).to.deep.equal(body); + releaseResponse.resolve(); + await finished; + }); + + it("replays bodies already read by a matcher", async () => { + const server = await start(); + await server.forAnyRequest().withBody('matched upload').thenForwardTo(targetUrl); + const finished = upload('matched upload'); + expect((await received).body.toString()).to.equal('matched upload'); + releaseResponse.resolve(); + await finished; + }); + + for (const callback of ['beforeRequest', 'beforeResponse'] as const) { + it(`preserves bodies needed by ${callback}`, async () => { + const server = await start(); + let observed: string | undefined; + await server.forAnyRequest().thenPassThrough(callback === 'beforeRequest' ? { + beforeRequest: async request => { observed = await request.body.getText(); } + } : { + beforeResponse: async (_response, request) => { observed = await request.body.getText(); } + }); + const finished = upload('callback upload'); + expect((await received).body.toString()).to.equal('callback upload'); + releaseResponse.resolve(); + await finished; + expect(observed).to.equal('callback upload'); + }); + } + + it("preserves request body transforms", async () => { + const server = await start(); + await server.forAnyRequest().thenForwardTo(targetUrl, { + transformRequest: { updateJsonBody: { extra: true } } + }); + const finished = upload(JSON.stringify({ original: true })); + expect(JSON.parse((await received).body.toString())).to.deep.equal({ original: true, extra: true }); + releaseResponse.resolve(); + await finished; + }); + + it("forwards oversized observed uploads without losing their prefix", async () => { + const server = await start({ recordTraffic: false, maxBodySize: 4 }); + const observed = getDeferred(); + await server.on('request', request => observed.resolve(request)); + await server.forAnyRequest().thenForwardTo(targetUrl); + const body = Buffer.alloc(256 * 1024, 'c'); + const finished = upload(body); + expect((await received).body).to.deep.equal(body); + expect((await observed).body.buffer.length).to.equal(0); + releaseResponse.resolve(); + await finished; + }); + }); +});