diff --git a/.changeset/ws-drain-on-shutdown.md b/.changeset/ws-drain-on-shutdown.md new file mode 100644 index 00000000..8a38f617 --- /dev/null +++ b/.changeset/ws-drain-on-shutdown.md @@ -0,0 +1,7 @@ +--- +"nostream": minor +--- + +feat(shutdown): drain WebSocket clients on SIGTERM + +On SIGTERM, `/readyz` returns 503, new WebSocket connections are rejected, and existing clients receive Nostr CLOSED messages before the socket closes. Drain is bounded by `WS_DRAIN_TIMEOUT_MS` (default 30s). diff --git a/deploy/README.md b/deploy/README.md index 937878db..a6983692 100644 --- a/deploy/README.md +++ b/deploy/README.md @@ -116,8 +116,10 @@ Each dependency ping uses the default 3s timeout (`ADMIN_DEPENDENCY_PING_TIMEOUT Set your load balancer check timeout above that (for example HAProxy `timeout check 5s`) so slow-but-healthy backends do not flap during probes. Responses are cached in-process for 1s to absorb polling without hammering the DB pool. -Use readiness before routing traffic to a new instance during deploys; graceful -WebSocket draining on shutdown is planned as a follow-up. +Use readiness before routing traffic to a new instance during deploys. On +SIGTERM the relay returns `503` on `/readyz`, stops accepting new WebSocket +connections, and drains existing clients before exit (`WS_DRAIN_TIMEOUT_MS`, +default 30s). ## Image delivery on restricted networks diff --git a/src/@types/adapters.ts b/src/@types/adapters.ts index 7a2ddc7e..f84f542c 100644 --- a/src/@types/adapters.ts +++ b/src/@types/adapters.ts @@ -19,6 +19,7 @@ export type IWebSocketAdapter = EventEmitter & { getAuthenticatedPubkeys(): ReadonlySet /** Returns false if this AUTH event id was already accepted on this socket. */ addAuthenticatedPubkey(pubkey: string, authEventId: string): boolean + drainAndClose(reason?: string): void } export interface ICacheAdapter { diff --git a/src/adapters/web-socket-adapter.ts b/src/adapters/web-socket-adapter.ts index 6a279e32..37bb229c 100644 --- a/src/adapters/web-socket-adapter.ts +++ b/src/adapters/web-socket-adapter.ts @@ -5,7 +5,12 @@ import { WebSocket } from 'ws' import { ZodError } from 'zod' import { ContextMetadata, Factory } from '../@types/base' -import { createAuthChallengeMessage, createNoticeMessage, createOutgoingEventMessage } from '../utils/messages' +import { + createAuthChallengeMessage, + createClosedMessage, + createNoticeMessage, + createOutgoingEventMessage, +} from '../utils/messages' import { IAbortable, IMessageHandler } from '../@types/message-handlers' import { IncomingMessage, OutgoingMessage } from '../@types/messages' import { IWebSocketAdapter, IWebSocketServerAdapter } from '../@types/adapters' @@ -156,6 +161,16 @@ export class WebSocketAdapter extends EventEmitter implements IWebSocketAdapter return new Map(this.subscriptions) } + public drainAndClose(reason = 'relay shutting down'): void { + this.subscriptions.forEach((_filters, subscriptionId) => { + this.sendMessage(createClosedMessage(subscriptionId, `closed: ${reason}`)) + }) + + if (this.client.readyState === WebSocket.OPEN) { + this.client.close(1001, reason) + } + } + // NIP-42 public getChallenge(): string { return this.session.getChallenge() diff --git a/src/adapters/web-socket-server-adapter.ts b/src/adapters/web-socket-server-adapter.ts index bc93a754..adb12925 100644 --- a/src/adapters/web-socket-server-adapter.ts +++ b/src/adapters/web-socket-server-adapter.ts @@ -1,5 +1,5 @@ import { IncomingMessage, Server } from 'http' -import WebSocket, { OPEN, WebSocketServer } from 'ws' +import WebSocket, { CLOSED, CLOSING, OPEN, WebSocketServer } from 'ws' import { propEq } from 'ramda' import { IWebSocketAdapter, IWebSocketServerAdapter } from '../@types/adapters' @@ -10,6 +10,7 @@ import { Factory } from '../@types/base' import { getRemoteAddress } from '../utils/http' import { isRateLimited } from '../handlers/request-handlers/rate-limiter-middleware' import { Settings } from '../@types/settings' +import { getWsDrainTimeoutMs, isDraining } from '../utils/shutdown-state' import { WebServerAdapter } from './web-server-adapter' const logger = createLogger('web-socket-server-adapter') @@ -49,25 +50,61 @@ export class WebSocketServerAdapter extends WebServerAdapter implements IWebSock super.close(() => { logger('closing') clearInterval(this.heartbeatInterval) - this.webSocketServer.clients.forEach((webSocket: WebSocket) => { - const webSocketAdapter = this.webSocketsAdapters.get(webSocket) - if (webSocketAdapter) { - logger('terminating client %s: %s', webSocketAdapter.getClientId(), webSocketAdapter.getClientAddress()) - } - webSocket.terminate() - }) - logger('closing web socket server') - this.webSocketServer.close(() => { - this.webSocketServer.removeAllListeners() - if (typeof callback !== 'undefined') { - callback() - } - logger('closed') + void this.drainClients(getWsDrainTimeoutMs()).finally(() => { + logger('closing web socket server') + this.webSocketServer.close(() => { + this.webSocketServer.removeAllListeners() + if (typeof callback !== 'undefined') { + callback() + } + logger('closed') + }) }) }) this.removeAllListeners() } + private async drainClients(timeoutMs: number): Promise { + const clients = [...this.webSocketServer.clients] as WebSocket[] + if (clients.length === 0) { + return + } + + logger('draining %d websocket client(s)', clients.length) + + for (const webSocket of clients) { + const webSocketAdapter = this.webSocketsAdapters.get(webSocket) + if (webSocketAdapter) { + logger('closing client %s: %s', webSocketAdapter.getClientId(), webSocketAdapter.getClientAddress()) + webSocketAdapter.drainAndClose() + } else if (webSocket.readyState === OPEN) { + webSocket.close(1001, 'relay shutting down') + } + } + + await Promise.race([ + Promise.all(clients.map((webSocket) => this.waitForWebSocketClose(webSocket))), + new Promise((resolve) => setTimeout(resolve, timeoutMs)), + ]) + + for (const webSocket of this.webSocketServer.clients) { + if (webSocket.readyState === OPEN || webSocket.readyState === CLOSING) { + logger('terminating client after drain timeout') + webSocket.terminate() + } + } + } + + private waitForWebSocketClose(webSocket: WebSocket): Promise { + if (webSocket.readyState === CLOSED) { + return Promise.resolve() + } + + return new Promise((resolve) => { + webSocket.once('close', () => resolve()) + }) + } + private onBroadcast(event: Event) { this.webSocketServer.clients.forEach((webSocket: WebSocket) => { if (!propEq('readyState', OPEN)(webSocket)) { @@ -86,6 +123,12 @@ export class WebSocketServerAdapter extends WebServerAdapter implements IWebSock } private async onConnection(client: WebSocket, req: IncomingMessage) { + if (isDraining()) { + logger('client rejected: draining') + client.close(1001, 'relay shutting down') + return + } + const currentSettings = this.settings() const remoteAddress = getRemoteAddress(req, currentSettings) diff --git a/src/app/app.ts b/src/app/app.ts index 6ae8c6a2..5b947028 100644 --- a/src/app/app.ts +++ b/src/app/app.ts @@ -17,6 +17,7 @@ const logger = createLogger('app-primary') export class App implements IRunnable { private workers: WeakMap> private watchers: FSWatcher[] | undefined + private shuttingDown = false public constructor( private readonly process: NodeJS.Process, @@ -155,7 +156,7 @@ export class App implements IRunnable { private onClusterExit(deadWorker: Worker, code: number, signal: string) { logger('worker %s died', deadWorker.process.pid) - if (code === 0 || signal === 'SIGINT') { + if (this.shuttingDown || code === 0 || signal === 'SIGINT') { return } setTimeout(() => { @@ -172,7 +173,33 @@ export class App implements IRunnable { } private onExit() { + if (this.shuttingDown) { + return + } + this.shuttingDown = true logger.info('exiting') + + const workers = Object.values(this.cluster.workers ?? {}) as Worker[] + if (workers.length === 0) { + this.finishExit() + return + } + + let remaining = workers.length + const onWorkerDone = () => { + remaining -= 1 + if (remaining <= 0) { + this.finishExit() + } + } + + for (const worker of workers) { + worker.once('exit', onWorkerDone) + worker.process.kill('SIGTERM') + } + } + + private finishExit() { void shutdownMetricsTelemetry().finally(() => { this.close(() => { this.process.exit(0) diff --git a/src/app/worker.ts b/src/app/worker.ts index 6512104a..4c239afc 100644 --- a/src/app/worker.ts +++ b/src/app/worker.ts @@ -6,9 +6,11 @@ import { createLogger } from '../factories/logger-factory' import { FSWatcher } from 'fs' import { SettingsStatic } from '../utils/settings' import { shutdownMetricsTelemetry } from '../telemetry/metrics' +import { beginDraining } from '../utils/shutdown-state' const logger = createLogger('app-worker') export class AppWorker implements IRunnable { + private exiting = false private watchers: FSWatcher[] | undefined public constructor( @@ -46,6 +48,11 @@ export class AppWorker implements IRunnable { } private onExit() { + if (this.exiting) { + return + } + this.exiting = true + beginDraining() logger('exiting') void shutdownMetricsTelemetry().finally(() => { this.close(() => { diff --git a/src/handlers/request-handlers/get-readyz-request-handler.ts b/src/handlers/request-handlers/get-readyz-request-handler.ts index 0ff9ff1e..34c0c586 100644 --- a/src/handlers/request-handlers/get-readyz-request-handler.ts +++ b/src/handlers/request-handlers/get-readyz-request-handler.ts @@ -1,6 +1,7 @@ import { NextFunction, Request, Response } from 'express' import { AdminDependencyHealth, collectAdminHealthSnapshot } from '../../utils/admin-health' +import { isDraining } from '../../utils/shutdown-state' // Public readiness probe for load balancers (e.g. HAProxy blue/green). Unlike /healthz // (liveness), /readyz returns non-200 when Postgres or Redis is unavailable. @@ -72,6 +73,16 @@ const sendReadyzResponse = (res: Response, statusCode: number, snapshot: ReadyzS } export const getReadyzRequestHandler = async (_req: Request, res: Response, next: NextFunction) => { + if (isDraining()) { + sendReadyzResponse(res, 503, { + status: 'unavailable', + database: { ok: false }, + redis: { ok: false }, + }) + next() + return + } + try { const snapshot = await collectReadyzSnapshot() const statusCode = snapshot.status === 'ok' ? 200 : 503 diff --git a/src/utils/shutdown-state.ts b/src/utils/shutdown-state.ts new file mode 100644 index 00000000..d740e1cf --- /dev/null +++ b/src/utils/shutdown-state.ts @@ -0,0 +1,27 @@ +const DEFAULT_WS_DRAIN_TIMEOUT_MS = 30_000 + +let draining = false + +export const beginDraining = (): void => { + draining = true +} + +export const isDraining = (): boolean => draining + +export const resetDrainingState = (): void => { + draining = false +} + +export const getWsDrainTimeoutMs = (): number => { + const raw = process.env.WS_DRAIN_TIMEOUT_MS + if (raw === undefined || raw === '') { + return DEFAULT_WS_DRAIN_TIMEOUT_MS + } + + const parsed = Number(raw) + if (!Number.isFinite(parsed) || parsed < 0) { + return DEFAULT_WS_DRAIN_TIMEOUT_MS + } + + return parsed +} diff --git a/test/unit/adapters/web-socket-server-adapter.spec.ts b/test/unit/adapters/web-socket-server-adapter.spec.ts index adf85373..cd28f6ed 100644 --- a/test/unit/adapters/web-socket-server-adapter.spec.ts +++ b/test/unit/adapters/web-socket-server-adapter.spec.ts @@ -12,6 +12,7 @@ const { expect } = chai import { WebSocketAdapterEvent, WebSocketServerAdapterEvent } from '../../../src/constants/adapter' import { WebSocketServerAdapter } from '../../../src/adapters/web-socket-server-adapter' +import * as shutdownState from '../../../src/utils/shutdown-state' describe('WebSocketServerAdapter', () => { let sandbox: Sinon.SinonSandbox @@ -23,10 +24,19 @@ describe('WebSocketServerAdapter', () => { let isRateLimitedStub: Sinon.SinonStub let originalConsoleError: typeof console.error + let originalDrainTimeout: string | undefined + + const flushClose = async (callback?: () => void) => { + adapter.close(callback) + await sandbox.clock.runAllAsync() + } beforeEach(() => { sandbox = Sinon.createSandbox() sandbox.useFakeTimers() + originalDrainTimeout = process.env.WS_DRAIN_TIMEOUT_MS + process.env.WS_DRAIN_TIMEOUT_MS = '0' + shutdownState.resetDrainingState() originalConsoleError = console.error console.error = () => undefined @@ -71,6 +81,12 @@ describe('WebSocketServerAdapter', () => { webSocketServer.close.callsFake((cb: () => void) => cb()) adapter.close() sandbox.restore() + shutdownState.resetDrainingState() + if (originalDrainTimeout === undefined) { + delete process.env.WS_DRAIN_TIMEOUT_MS + } else { + process.env.WS_DRAIN_TIMEOUT_MS = originalDrainTimeout + } }) describe('constructor', () => { @@ -111,48 +127,99 @@ describe('WebSocketServerAdapter', () => { expect(webServer.close).to.have.been.calledOnce }) - it('terminates all connected WebSocket clients', () => { - const terminateStub1 = sandbox.stub() - const terminateStub2 = sandbox.stub() + it('drains connected WebSocket clients before closing the server', async () => { + const drainAndCloseStub1 = sandbox.stub() + const drainAndCloseStub2 = sandbox.stub() + const client1 = { readyState: 1, once: sandbox.stub().callsFake((_event: string, cb: () => void) => cb()) } + const client2 = { readyState: 1, once: sandbox.stub().callsFake((_event: string, cb: () => void) => cb()) } + + const mockAdapter1 = { + drainAndClose: drainAndCloseStub1, + getClientId: () => 'client-1', + getClientAddress: () => '127.0.0.1', + } + const mockAdapter2 = { + drainAndClose: drainAndCloseStub2, + getClientId: () => 'client-2', + getClientAddress: () => '127.0.0.2', + } + + const connectionCall = webSocketServer.on + .getCalls() + .find((call: any) => call.args[0] === WebSocketServerAdapterEvent.Connection) + const onConnection = connectionCall.args[1] + + webSocketServer.clients = new Set([client1, client2] as any) + createWebSocketAdapter.callsFake(([client]: [typeof client1, unknown, unknown]) => { + if (client === client1) { + return mockAdapter1 + } + if (client === client2) { + return mockAdapter2 + } + }) + + await onConnection(client1, { headers: {}, socket: { remoteAddress: '127.0.0.1' } }) + await onConnection(client2, { headers: {}, socket: { remoteAddress: '127.0.0.2' } }) + + webServer.close.callsFake((cb: () => void) => cb()) + webSocketServer.close.callsFake((cb: () => void) => cb()) + + await flushClose() + + expect(drainAndCloseStub1).to.have.been.calledOnce + expect(drainAndCloseStub2).to.have.been.calledOnce + expect(webSocketServer.close).to.have.been.calledOnce + }) + + it('terminates clients that remain open after the drain timeout', async () => { + process.env.WS_DRAIN_TIMEOUT_MS = '1000' - webSocketServer.clients = new Set([{ terminate: terminateStub1 }, { terminate: terminateStub2 }] as any) + const terminateStub = sandbox.stub() + const client = { + readyState: 1, + close: sandbox.stub(), + terminate: terminateStub, + once: sandbox.stub(), + } + webSocketServer.clients = new Set([client] as any) webServer.close.callsFake((cb: () => void) => cb()) webSocketServer.close.callsFake((cb: () => void) => cb()) adapter.close() + await sandbox.clock.tickAsync(1000) - expect(terminateStub1).to.have.been.calledOnce - expect(terminateStub2).to.have.been.calledOnce + expect(terminateStub).to.have.been.calledOnce }) - it('closes the webSocketServer after terminating clients', () => { + it('closes the webSocketServer after draining clients', async () => { webSocketServer.clients = new Set() webServer.close.callsFake((cb: () => void) => cb()) webSocketServer.close.callsFake((cb: () => void) => cb()) - adapter.close() + await flushClose() expect(webSocketServer.close).to.have.been.calledOnce }) - it('invokes callback after full close', () => { + it('invokes callback after full close', async () => { const callback = sandbox.stub() webSocketServer.clients = new Set() webServer.close.callsFake((cb: () => void) => cb()) webSocketServer.close.callsFake((cb: () => void) => cb()) - adapter.close(callback) + await flushClose(callback) expect(callback).to.have.been.calledOnce }) - it('removes all listeners from webSocketServer after close', () => { + it('removes all listeners from webSocketServer after close', async () => { webSocketServer.clients = new Set() webServer.close.callsFake((cb: () => void) => cb()) webSocketServer.close.callsFake((cb: () => void) => cb()) - adapter.close() + await flushClose() expect(webSocketServer.removeAllListeners).to.have.been.calledOnce }) @@ -277,5 +344,20 @@ describe('WebSocketServerAdapter', () => { expect(terminateStub).to.have.been.calledOnce expect(createWebSocketAdapter).not.to.have.been.called }) + + it('rejects new connections while draining', async () => { + const closeStub = sandbox.stub() + shutdownState.beginDraining() + + const connectionCall = webSocketServer.on + .getCalls() + .find((call: any) => call.args[0] === WebSocketServerAdapterEvent.Connection) + const onConnection = connectionCall.args[1] + + await onConnection({ close: closeStub }, { headers: {}, socket: { remoteAddress: '127.0.0.1' } }) + + expect(closeStub).to.have.been.calledOnceWithExactly(1001, 'relay shutting down') + expect(createWebSocketAdapter).not.to.have.been.called + }) }) }) diff --git a/test/unit/handlers/request-handlers/get-readyz-request-handler.spec.ts b/test/unit/handlers/request-handlers/get-readyz-request-handler.spec.ts index 25912fd6..84dda718 100644 --- a/test/unit/handlers/request-handlers/get-readyz-request-handler.spec.ts +++ b/test/unit/handlers/request-handlers/get-readyz-request-handler.spec.ts @@ -8,6 +8,7 @@ import { getReadyzRequestHandler, resetReadyzSnapshotCache, } from '../../../../src/handlers/request-handlers/get-readyz-request-handler' +import { beginDraining, resetDrainingState } from '../../../../src/utils/shutdown-state' chai.use(sinonChai) const { expect } = chai @@ -58,6 +59,25 @@ describe('getReadyzRequestHandler', () => { afterEach(() => { sandbox.restore() resetReadyzSnapshotCache() + resetDrainingState() + }) + + it('responds with 503 JSON while the relay is draining', async () => { + beginDraining() + + const res = createResponse() + const next = sinon.stub() + + await getReadyzRequestHandler({} as any, res, next) + + expect(collectAdminHealthSnapshotStub).not.to.have.been.called + expect(res.status).to.have.been.calledOnceWithExactly(503) + expect(res.send).to.have.been.calledOnceWithExactly({ + status: 'unavailable', + database: { ok: false }, + redis: { ok: false }, + }) + expect(next).to.have.been.calledOnce }) it('responds with 200 JSON when dependencies are ready', async () => { diff --git a/test/unit/utils/shutdown-state.spec.ts b/test/unit/utils/shutdown-state.spec.ts new file mode 100644 index 00000000..0925c821 --- /dev/null +++ b/test/unit/utils/shutdown-state.spec.ts @@ -0,0 +1,44 @@ +import chai from 'chai' + +import { + beginDraining, + getWsDrainTimeoutMs, + isDraining, + resetDrainingState, +} from '../../../src/utils/shutdown-state' + +const { expect } = chai + +describe('shutdown-state', () => { + const originalTimeout = process.env.WS_DRAIN_TIMEOUT_MS + + afterEach(() => { + resetDrainingState() + if (originalTimeout === undefined) { + delete process.env.WS_DRAIN_TIMEOUT_MS + } else { + process.env.WS_DRAIN_TIMEOUT_MS = originalTimeout + } + }) + + it('tracks draining state', () => { + expect(isDraining()).to.equal(false) + beginDraining() + expect(isDraining()).to.equal(true) + }) + + it('defaults WS drain timeout to 30s', () => { + delete process.env.WS_DRAIN_TIMEOUT_MS + expect(getWsDrainTimeoutMs()).to.equal(30_000) + }) + + it('reads WS drain timeout from env', () => { + process.env.WS_DRAIN_TIMEOUT_MS = '5000' + expect(getWsDrainTimeoutMs()).to.equal(5000) + }) + + it('falls back to default for invalid WS drain timeout', () => { + process.env.WS_DRAIN_TIMEOUT_MS = 'invalid' + expect(getWsDrainTimeoutMs()).to.equal(30_000) + }) +})