From 7668a6a32845807857277f3f49b867eb2bf238a3 Mon Sep 17 00:00:00 2001 From: Matthias Erll Date: Wed, 7 Oct 2026 15:52:57 +0200 Subject: [PATCH] feat: apply rate limit for jwt token verification --- src/middleware/rate-limit.test.ts | 134 ++++++++++++++++++++++++++++++ src/middleware/rate-limit.ts | 61 +++++++++++++- src/middleware/session.test.ts | 27 ++++++ src/middleware/session.ts | 13 +-- 4 files changed, 227 insertions(+), 8 deletions(-) create mode 100644 src/middleware/rate-limit.test.ts diff --git a/src/middleware/rate-limit.test.ts b/src/middleware/rate-limit.test.ts new file mode 100644 index 000000000..464f6f997 --- /dev/null +++ b/src/middleware/rate-limit.test.ts @@ -0,0 +1,134 @@ +import express from 'express' +import request from 'supertest' + +let mockTrustProxy = 1 + +jest.mock('src/validators', () => ({ + cleanEnv: () => ({ + RATE_LIMIT_WINDOW_MS: 60_000, + RATE_LIMIT_MAX_REQUESTS: 100, + RATE_LIMIT_AUTH_MAX_ATTEMPTS: 2, + TRUST_PROXY: mockTrustProxy, + }), +})) + +describe('socket authentication rate limiting', () => { + let rateLimits: typeof import('./rate-limit') + + beforeEach(async () => { + jest.resetModules() + jest.useFakeTimers() + mockTrustProxy = 1 + rateLimits = await import('./rate-limit') + }) + + afterEach(() => { + jest.clearAllTimers() + jest.useRealTimers() + }) + + it('blocks attempts above the configured limit and reports retry metadata', async () => { + await rateLimits.beginSocketAuthAttempt('192.0.2.1', {}) + await rateLimits.beginSocketAuthAttempt('192.0.2.1', {}) + + await expect(rateLimits.beginSocketAuthAttempt('192.0.2.1', {})).rejects.toMatchObject({ + message: 'Too many authentication attempts from this IP, please try again later.', + data: { status: 429, retryAfter: 60 }, + }) + }) + + it('does not count successful authentication attempts', async () => { + for (let attempt = 0; attempt < 5; attempt++) { + const complete = await rateLimits.beginSocketAuthAttempt('192.0.2.1', {}) + await complete() + } + + await expect(rateLimits.beginSocketAuthAttempt('192.0.2.1', {})).resolves.toBeInstanceOf(Function) + }) + + it('reserves concurrent attempts before authentication completes', async () => { + const results = await Promise.allSettled( + Array.from({ length: 3 }, () => rateLimits.beginSocketAuthAttempt('192.0.2.1', {})), + ) + + expect(results.filter(({ status }) => status === 'fulfilled').length).toBeLessThanOrEqual(2) + expect(results.some(({ status }) => status === 'rejected')).toBe(true) + await expect(rateLimits.beginSocketAuthAttempt('192.0.2.1', {})).rejects.toBeInstanceOf( + rateLimits.SocketAuthRateLimitError, + ) + }) + + it('allows attempts again after the configured window', async () => { + await rateLimits.beginSocketAuthAttempt('192.0.2.1', {}) + await rateLimits.beginSocketAuthAttempt('192.0.2.1', {}) + jest.advanceTimersByTime(60_000) + + await expect(rateLimits.beginSocketAuthAttempt('192.0.2.1', {})).resolves.toBeInstanceOf(Function) + }) + + it('keeps separate counters for different client IPs behind a trusted proxy', async () => { + const headers = { 'x-forwarded-for': '192.0.2.1' } + await rateLimits.beginSocketAuthAttempt('10.0.0.1', headers) + await rateLimits.beginSocketAuthAttempt('10.0.0.1', headers) + + await expect( + rateLimits.beginSocketAuthAttempt('10.0.0.1', { 'x-forwarded-for': '192.0.2.2' }), + ).resolves.toBeInstanceOf(Function) + await expect( + rateLimits.beginSocketAuthAttempt('10.0.0.1', { 'x-forwarded-for': 'spoofed, 192.0.2.1' }), + ).rejects.toBeInstanceOf(rateLimits.SocketAuthRateLimitError) + }) + + it('groups IPv6 clients by subnet like the REST limiter', async () => { + await rateLimits.beginSocketAuthAttempt('2001:db8::1', {}) + await rateLimits.beginSocketAuthAttempt('2001:db8::2', {}) + + await expect(rateLimits.beginSocketAuthAttempt('2001:db8::3', {})).rejects.toBeInstanceOf( + rateLimits.SocketAuthRateLimitError, + ) + }) + + it('ignores forwarded addresses when proxy trust is disabled', async () => { + jest.resetModules() + mockTrustProxy = 0 + rateLimits = await import('./rate-limit') + await rateLimits.beginSocketAuthAttempt('192.0.2.1', { 'x-forwarded-for': '192.0.2.2' }) + await rateLimits.beginSocketAuthAttempt('192.0.2.1', { 'x-forwarded-for': '192.0.2.3' }) + + await expect( + rateLimits.beginSocketAuthAttempt('192.0.2.1', { 'x-forwarded-for': '192.0.2.4' }), + ).rejects.toBeInstanceOf(rateLimits.SocketAuthRateLimitError) + }) + + it('uses the configured number of trusted proxy hops', async () => { + jest.resetModules() + mockTrustProxy = 2 + rateLimits = await import('./rate-limit') + const headers = { 'x-forwarded-for': '192.0.2.1, 10.0.0.2' } + await rateLimits.beginSocketAuthAttempt('10.0.0.1', headers) + await rateLimits.beginSocketAuthAttempt('10.0.0.1', headers) + + await expect(rateLimits.beginSocketAuthAttempt('192.0.2.1', {})).rejects.toBeInstanceOf( + rateLimits.SocketAuthRateLimitError, + ) + await expect(rateLimits.beginSocketAuthAttempt('10.0.0.2', {})).resolves.toBeInstanceOf(Function) + }) + + it('rejects invalid client addresses', async () => { + await expect(rateLimits.beginSocketAuthAttempt('invalid', {})).rejects.toThrow('Invalid socket client IP address') + }) + + it('shares the failed-attempt budget with REST authentication', async () => { + jest.useRealTimers() + const app = express() + app.set('trust proxy', 1) + app.use(rateLimits.authRateLimiter) + app.get('/', (_req, res) => { + res.sendStatus(401) + }) + await request(app).get('/').set('Authorization', 'invalid').set('X-Forwarded-For', '192.0.2.1').expect(401) + await rateLimits.beginSocketAuthAttempt('192.0.2.1', {}) + + await request(app).get('/').set('Authorization', 'invalid').set('X-Forwarded-For', '192.0.2.1').expect(429) + }) +}) diff --git a/src/middleware/rate-limit.ts b/src/middleware/rate-limit.ts index 78e40950d..3185f87fd 100644 --- a/src/middleware/rate-limit.ts +++ b/src/middleware/rate-limit.ts @@ -1,12 +1,66 @@ -import rateLimit from 'express-rate-limit' -import { cleanEnv, RATE_LIMIT_AUTH_MAX_ATTEMPTS, RATE_LIMIT_MAX_REQUESTS, RATE_LIMIT_WINDOW_MS } from 'src/validators' +import rateLimit, { ipKeyGenerator, MemoryStore } from 'express-rate-limit' +import { IncomingHttpHeaders } from 'http' +import { isIP } from 'net' +import { + cleanEnv, + RATE_LIMIT_AUTH_MAX_ATTEMPTS, + RATE_LIMIT_MAX_REQUESTS, + RATE_LIMIT_WINDOW_MS, + TRUST_PROXY, +} from 'src/validators' const env = cleanEnv({ RATE_LIMIT_WINDOW_MS, RATE_LIMIT_MAX_REQUESTS, RATE_LIMIT_AUTH_MAX_ATTEMPTS, + TRUST_PROXY, }) +const authLimitMessage = 'Too many authentication attempts from this IP, please try again later.' + +export class SocketAuthRateLimitError extends Error { + readonly data: { status: number; retryAfter: number } + + constructor(retryAfter: number) { + super(authLimitMessage) + this.data = { status: 429, retryAfter } + } +} + +const authStore = new MemoryStore() + +export async function beginSocketAuthAttempt( + address: string, + headers: IncomingHttpHeaders, +): Promise<() => Promise> { + // Match Express's numeric trust-proxy setting: walk from the nearest hop. + const forwarded = headers['x-forwarded-for'] + const addresses = [ + address, + ...(typeof forwarded === 'string' + ? forwarded + .split(',') + .map((ip) => ip.trim()) + .reverse() + : []), + ] + const ip = addresses[Math.min(Math.max(0, Math.ceil(env.TRUST_PROXY)), addresses.length - 1)] + if (!isIP(ip)) throw new Error('Invalid socket client IP address') + + const key = ipKeyGenerator(ip) + const { totalHits, resetTime } = await authStore.increment(key) + if (totalHits > env.RATE_LIMIT_AUTH_MAX_ATTEMPTS) { + const retryAfter = resetTime + ? Math.ceil((resetTime.getTime() - Date.now()) / 1000) + : env.RATE_LIMIT_WINDOW_MS / 1000 + throw new SocketAuthRateLimitError(Math.max(0, Math.ceil(retryAfter))) + } + + // Reserve before verification so concurrent attempts cannot bypass the limit. + // Successful authentication releases the reservation, just as for REST. + return () => authStore.decrement(key) +} + /** * General API rate limiter * @@ -65,12 +119,13 @@ export const apiRateLimiter = rateLimit({ * */ export const authRateLimiter = rateLimit({ + store: authStore, windowMs: env.RATE_LIMIT_WINDOW_MS, max: env.RATE_LIMIT_AUTH_MAX_ATTEMPTS, standardHeaders: true, legacyHeaders: false, message: { - message: 'Too many authentication attempts from this IP, please try again later.', + message: authLimitMessage, retryAfter: Math.floor(env.RATE_LIMIT_WINDOW_MS / 1000), }, // Only count requests with Authorization header diff --git a/src/middleware/session.test.ts b/src/middleware/session.test.ts index 247e54730..72f73c8eb 100644 --- a/src/middleware/session.test.ts +++ b/src/middleware/session.test.ts @@ -11,6 +11,8 @@ const mockVerifyJwt = jest.fn() const mockIoUse = jest.fn() const mockIoOn = jest.fn() const mockIoOf = jest.fn() +const mockBeginSocketAuthAttempt = jest.fn() +const mockCompleteAuthAttempt = jest.fn() const mockReadOnlyInit = jest.fn() const mockReadOnlySetLocked = jest.fn() @@ -98,6 +100,13 @@ jest.mock('src/jwt-verification', () => ({ verifyJwt: (...args: unknown[]) => mockVerifyJwt(...args), })) +jest.mock('./rate-limit', () => ({ + beginSocketAuthAttempt: (...args: unknown[]) => mockBeginSocketAuthAttempt(...args), + SocketAuthRateLimitError: class extends Error { + data = { status: 429, retryAfter: 60 } + }, +})) + jest.mock('socket.io', () => ({ Server: jest.fn().mockImplementation(() => ({ use: mockIoUse, @@ -213,6 +222,8 @@ describe('session middleware', () => { sessionModule.sessionMiddleware({} as http.Server) ;[[authenticate]] = mockIoUse.mock.calls mockVerifyJwt.mockResolvedValue({ email: 'verified@example.com' }) + mockBeginSocketAuthAttempt.mockResolvedValue(mockCompleteAuthAttempt) + mockCompleteAuthAttempt.mockResolvedValue(undefined) }) it.each([undefined, '', ' ', 123, {}, ['token']])( @@ -224,6 +235,7 @@ describe('session middleware', () => { expect(next).toHaveBeenCalledWith(expect.objectContaining({ message: 'Unauthorized' })) expect(mockVerifyJwt).not.toHaveBeenCalled() + expect(mockCompleteAuthAttempt).not.toHaveBeenCalled() }, ) @@ -235,6 +247,7 @@ describe('session middleware', () => { expect(mockVerifyJwt).toHaveBeenCalledWith(token) expect(socket.data.email).toBe('verified@example.com') + expect(mockCompleteAuthAttempt).toHaveBeenCalledTimes(1) expect(next).toHaveBeenCalledWith() }) @@ -270,6 +283,20 @@ describe('session middleware', () => { expect(next).toHaveBeenCalledTimes(1) expect(next).toHaveBeenCalledWith(expect.objectContaining({ message: 'Unauthorized' })) expect(socket.data.email).toBeUndefined() + expect(mockCompleteAuthAttempt).not.toHaveBeenCalled() + }) + + it('rejects rate-limited connections before verifying JWTs', async () => { + const { SocketAuthRateLimitError } = await import('./rate-limit') + const error = new SocketAuthRateLimitError(60) + mockBeginSocketAuthAttempt.mockRejectedValueOnce(error) + const next = jest.fn() + + await authenticate(createSocket('valid-token'), next) + + expect(next).toHaveBeenCalledWith(error) + expect(mockVerifyJwt).not.toHaveBeenCalled() + expect(mockCompleteAuthAttempt).not.toHaveBeenCalled() }) it('uses the verified identity for user events', async () => { diff --git a/src/middleware/session.ts b/src/middleware/session.ts index 68ea346bb..f58213696 100644 --- a/src/middleware/session.ts +++ b/src/middleware/session.ts @@ -13,6 +13,7 @@ import { API_NAMESPACE, cleanEnv, EDITOR_INACTIVITY_TIMEOUT } from 'src/validato import { v4 as uuidv4 } from 'uuid' import { setApiStatusInConfigMap } from '../k8s-operations' import { getSanitizedErrorMessage } from '../utils' +import { beginSocketAuthAttempt, SocketAuthRateLimitError } from './rate-limit' const debug = Debug('otomi:session') const env = cleanEnv({ @@ -96,19 +97,21 @@ export function sessionMiddleware(server: http.Server): RequestHandler { if (!env.isTest && server) { io = new Server(server, { path: '/ws' }) io.use(async (socket, next) => { - const token: unknown = socket.handshake.auth.token ?? socket.handshake.headers.authorization - if (typeof token !== 'string' || !token.trim()) { - return next(new Error('Unauthorized')) - } - try { + const completeAuthAttempt = await beginSocketAuthAttempt(socket.handshake.address, socket.handshake.headers) + const token: unknown = socket.handshake.auth.token ?? socket.handshake.headers.authorization + if (typeof token !== 'string' || !token.trim()) { + return next(new Error('Unauthorized')) + } const { email, sub, groups, roles } = await verifyJwt(token) + await completeAuthAttempt() const { data } = socket data.email = email data.sub = sub data.groups = groups data.roles = roles } catch (error) { + if (error instanceof SocketAuthRateLimitError) return next(error) debug(`Socket JWT verification failed: ${getSanitizedErrorMessage(error)}`) return next(new Error('Unauthorized')) }