diff --git a/app/server/modules/events/__tests__/events.controller.test.ts b/app/server/modules/events/__tests__/events.controller.test.ts index 48cc2625..20fb4789 100644 --- a/app/server/modules/events/__tests__/events.controller.test.ts +++ b/app/server/modules/events/__tests__/events.controller.test.ts @@ -1,5 +1,6 @@ import { test, describe, expect } from "bun:test"; import { createApp } from "~/server/app"; +import { serverEvents } from "~/server/core/events"; import { createTestSession, getAuthHeaders } from "~/test/helpers/auth"; const app = createApp(); @@ -30,6 +31,32 @@ describe("events security", () => { expect(res.status).toBe(200); expect(res.headers.get("Content-Type")).toBe("text/event-stream"); + await res.body?.cancel(); + }); + + test("should cleanup SSE listeners when client disconnects", async () => { + const { token } = await createTestSession(); + const initialCount = serverEvents.listenerCount("doctor:cancelled"); + + const res = await app.request("/api/v1/events", { + headers: getAuthHeaders(token), + }); + + expect(res.status).toBe(200); + + for (let i = 0; i < 20 && serverEvents.listenerCount("doctor:cancelled") < initialCount + 1; i++) { + await new Promise((resolve) => setTimeout(resolve, 10)); + } + + expect(serverEvents.listenerCount("doctor:cancelled")).toBe(initialCount + 1); + + await res.body?.cancel(); + + for (let i = 0; i < 20 && serverEvents.listenerCount("doctor:cancelled") > initialCount; i++) { + await new Promise((resolve) => setTimeout(resolve, 10)); + } + + expect(serverEvents.listenerCount("doctor:cancelled")).toBe(initialCount); }); describe("unauthenticated access", () => { diff --git a/app/server/modules/events/events.controller.ts b/app/server/modules/events/events.controller.ts index 473255ae..7e913d98 100644 --- a/app/server/modules/events/events.controller.ts +++ b/app/server/modules/events/events.controller.ts @@ -162,10 +162,13 @@ export const eventsController = new Hono().use(requireAuth).get("/", (c) => { serverEvents.on("doctor:cancelled", onDoctorCancelled); let keepAlive = true; + let cleanedUp = false; - stream.onAbort(() => { - logger.info("Client disconnected from SSE endpoint"); - keepAlive = false; + function cleanup() { + if (cleanedUp) return; + cleanedUp = true; + + c.req.raw.signal.removeEventListener("abort", onRequestAbort); serverEvents.off("backup:started", onBackupStarted); serverEvents.off("backup:progress", onBackupProgress); serverEvents.off("backup:completed", onBackupCompleted); @@ -177,14 +180,33 @@ export const eventsController = new Hono().use(requireAuth).get("/", (c) => { serverEvents.off("doctor:started", onDoctorStarted); serverEvents.off("doctor:completed", onDoctorCompleted); serverEvents.off("doctor:cancelled", onDoctorCancelled); - }); + } - while (keepAlive) { - await stream.writeSSE({ - data: JSON.stringify({ timestamp: Date.now() }), - event: "heartbeat", - }); - await stream.sleep(5000); + function handleDisconnect() { + if (!keepAlive) return; + logger.info("Client disconnected from SSE endpoint"); + keepAlive = false; + cleanup(); + } + + function onRequestAbort() { + handleDisconnect(); + stream.abort(); + } + + stream.onAbort(handleDisconnect); + c.req.raw.signal.addEventListener("abort", onRequestAbort, { once: true }); + + try { + while (keepAlive && !c.req.raw.signal.aborted && !stream.aborted) { + await stream.writeSSE({ + data: JSON.stringify({ timestamp: Date.now() }), + event: "heartbeat", + }); + await stream.sleep(5000); + } + } finally { + cleanup(); } }); });