diff --git a/libs/sdk-js/src/tests/sse.test.ts b/libs/sdk-js/src/tests/sse.test.ts index 51ebbbd0c..b4f7ca7d8 100644 --- a/libs/sdk-js/src/tests/sse.test.ts +++ b/libs/sdk-js/src/tests/sse.test.ts @@ -2,26 +2,21 @@ import { Readable } from "node:stream"; import { IterableReadableStream } from "../utils/stream.js"; import { BytesLineDecoder, SSEDecoder } from "../utils/sse.js"; +const gather = async (stream: ReadableStream): Promise => { + const results: T[] = []; + const iterator = IterableReadableStream.fromReadableStream(stream); + for await (const chunk of iterator) results.push(chunk); + return results; +}; + +const textEncoder = new TextEncoder(); +const textDecoder = new TextDecoder(); + describe("BytesLineDecoder", () => { const createStream = (chunks: Uint8Array[]) => { return Readable.toWeb(Readable.from(chunks)) as ReadableStream; }; - const gather = async ( - stream: ReadableStream, - ): Promise => { - const results: Uint8Array[] = []; - for await (const chunk of IterableReadableStream.fromReadableStream( - stream, - )) { - results.push(chunk); - } - return results; - }; - - const textEncoder = new TextEncoder(); - const textDecoder = new TextDecoder(); - test("handles single line with newline", async () => { const input = createStream([textEncoder.encode("hello\n")]); const decoded = input.pipeThrough(new BytesLineDecoder()); @@ -78,25 +73,22 @@ describe("BytesLineDecoder", () => { expect(textDecoder.decode(results[0])).toBe("line1"); expect(textDecoder.decode(results[1])).toBe("line2"); }); + + test("handles stale line", async () => { + const input = createStream([textEncoder.encode("hello")]); + const decoded = input.pipeThrough(new BytesLineDecoder()); + const results = await gather(decoded); + + expect(results.length).toBe(1); + expect(textDecoder.decode(results[0])).toBe("hello"); + }); }); describe("SSEDecoder", () => { const createStream = (lines: string[]) => { - const encoder = new TextEncoder(); - const chunks = lines.map((line) => encoder.encode(line)); - return Readable.toWeb(Readable.from(chunks)) as ReadableStream; - }; - - const collectResults = async ( - stream: ReadableStream, - ): Promise => { - const results: any[] = []; - for await (const chunk of IterableReadableStream.fromReadableStream( - stream, - )) { - results.push(chunk); - } - return results; + return Readable.toWeb( + Readable.from(lines.map((line) => textEncoder.encode(line))), + ) as ReadableStream; }; test("decodes simple event", async () => { @@ -109,7 +101,7 @@ describe("SSEDecoder", () => { .pipeThrough(new BytesLineDecoder()) .pipeThrough(new SSEDecoder()); - const results = await collectResults(decoded); + const results = await gather(decoded); expect(results.length).toBe(1); expect(results[0]).toEqual({ event: "test", @@ -127,7 +119,7 @@ describe("SSEDecoder", () => { .pipeThrough(new BytesLineDecoder()) .pipeThrough(new SSEDecoder()); - const results = await collectResults(decoded); + const results = await gather(decoded); expect(results.length).toBe(1); expect(results[0]).toEqual({ event: "test", @@ -148,7 +140,7 @@ describe("SSEDecoder", () => { .pipeThrough(new BytesLineDecoder()) .pipeThrough(new SSEDecoder()); - const results = await collectResults(decoded); + const results = await gather(decoded); expect(results.length).toBe(2); expect(results[0]).toEqual({ event: "test1", @@ -166,7 +158,7 @@ describe("SSEDecoder", () => { .pipeThrough(new BytesLineDecoder()) .pipeThrough(new SSEDecoder()); - const results = await collectResults(decoded); + const results = await gather(decoded); expect(results.length).toBe(1); expect(results[0]).toEqual({ event: "test", @@ -180,7 +172,7 @@ describe("SSEDecoder", () => { .pipeThrough(new BytesLineDecoder()) .pipeThrough(new SSEDecoder()); - const results = await collectResults(decoded); + const results = await gather(decoded); expect(results.length).toBe(1); expect(results[0]).toEqual({ event: "end", diff --git a/libs/sdk-js/src/utils/sse.ts b/libs/sdk-js/src/utils/sse.ts index fb1ff0f59..c3c525113 100644 --- a/libs/sdk-js/src/utils/sse.ts +++ b/libs/sdk-js/src/utils/sse.ts @@ -1,10 +1,3 @@ -const mergeArrays = (a: ArrayLike, b: ArrayLike) => { - const mergedArray = new Uint8Array(a.length + b.length); - mergedArray.set(a); - mergedArray.set(b, a.length); - return mergedArray; -}; - const CR = "\r".charCodeAt(0); const LF = "\n".charCodeAt(0); const NULL = "\0".charCodeAt(0); @@ -15,12 +8,12 @@ const TRAILING_NEWLINE = [CR, LF]; export class BytesLineDecoder extends TransformStream { constructor() { - let buffer = new Uint8Array(); + let buffer: Uint8Array[] = []; let trailingCr = false; super({ start() { - buffer = new Uint8Array(); + buffer = []; trailingCr = false; }, @@ -30,7 +23,7 @@ export class BytesLineDecoder extends TransformStream { // Handle trailing CR from previous chunk if (trailingCr) { - text = mergeArrays([CR], text); + text = joinArrays([[CR], text]); trailingCr = false; } @@ -43,44 +36,45 @@ export class BytesLineDecoder extends TransformStream { if (!text.length) return; const trailingNewline = TRAILING_NEWLINE.includes(text.at(-1)!); - // Pre-allocate lines array with estimated capacity - let lines: Uint8Array[] = []; + const lastIdx = text.length - 1; + const { lines } = text.reduce<{ lines: Uint8Array[]; from: number }>( + (acc, cur, idx) => { + if (acc.from > idx) return acc; - for (let offset = 0; offset < text.byteLength; ) { - let idx = text.indexOf(CR, offset); - if (idx === -1) idx = text.indexOf(LF, offset); - if (idx === -1) { - lines.push(text.subarray(offset)); - break; - } + if (cur === CR || cur === LF) { + acc.lines.push(text.subarray(acc.from, idx)); + if (cur === CR && text[idx + 1] === LF) { + acc.from = idx + 2; + } else { + acc.from = idx + 1; + } + } - lines.push(text.subarray(offset, idx)); - if (text[idx] === CR && text[idx + 1] === LF) { - offset = idx + 2; - } else { - offset = idx + 1; - } - } + if (idx === lastIdx && acc.from <= lastIdx) { + acc.lines.push(text.subarray(acc.from)); + } + + return acc; + }, + { lines: [], from: 0 }, + ); if (lines.length === 1 && !trailingNewline) { - buffer = mergeArrays(buffer, lines[0]); + buffer.push(lines[0]); return; } if (buffer.length) { // Include existing buffer in first line - buffer = mergeArrays(buffer, lines[0]); - - lines = lines.slice(1); - lines.unshift(buffer); - - buffer = new Uint8Array(); + buffer.push(lines[0]); + lines[0] = joinArrays(buffer); + buffer = []; } if (!trailingNewline) { // If the last segment is not newline terminated, // buffer it for the next chunk - if (lines.length) buffer = lines.pop()!; + if (lines.length) buffer = [lines.pop()!]; } // Enqueue complete lines @@ -91,7 +85,7 @@ export class BytesLineDecoder extends TransformStream { flush(controller) { if (buffer.length) { - controller.enqueue(buffer); + controller.enqueue(joinArrays(buffer)); } }, }); @@ -106,7 +100,7 @@ interface StreamPart { export class SSEDecoder extends TransformStream { constructor() { let event = ""; - let data: Uint8Array = new Uint8Array(); + let data: Uint8Array[] = []; let lastEventId = ""; let retry: number | null = null; @@ -120,12 +114,12 @@ export class SSEDecoder extends TransformStream { const sse = { event, - data: data.length ? JSON.parse(decoder.decode(data)) : null, + data: data.length ? decodeArraysToJson(decoder, data) : null, }; // NOTE: as per the SSE spec, do not reset lastEventId event = ""; - data = new Uint8Array(); + data = []; retry = null; controller.enqueue(sse); @@ -145,7 +139,7 @@ export class SSEDecoder extends TransformStream { if (fieldName === "event") { event = decoder.decode(value); } else if (fieldName === "data") { - data = mergeArrays(data, value); + data.push(value); } else if (fieldName === "id") { if (value.indexOf(NULL) === -1) lastEventId = decoder.decode(value); } else if (fieldName === "retry") { @@ -158,10 +152,25 @@ export class SSEDecoder extends TransformStream { if (event) { controller.enqueue({ event, - data: data.length ? JSON.parse(decoder.decode(data)) : null, + data: data.length ? decodeArraysToJson(decoder, data) : null, }); } }, }); } } + +function joinArrays(data: ArrayLike[]) { + const totalLength = data.reduce((acc, curr) => acc + curr.length, 0); + let merged = new Uint8Array(totalLength); + let offset = 0; + for (const c of data) { + merged.set(c, offset); + offset += c.length; + } + return merged; +} + +function decodeArraysToJson(decoder: TextDecoder, data: ArrayLike[]) { + return JSON.parse(decoder.decode(joinArrays(data))); +}