diff --git a/apps/server/src/integrations/throttle/user-throttler.guard.spec.ts b/apps/server/src/integrations/throttle/user-throttler.guard.spec.ts new file mode 100644 index 000000000..1110d8ad8 --- /dev/null +++ b/apps/server/src/integrations/throttle/user-throttler.guard.spec.ts @@ -0,0 +1,51 @@ +import { Reflector } from '@nestjs/core'; +import { ThrottlerStorageService } from '@nestjs/throttler'; +import { JwtType } from '../../core/auth/dto/jwt-payload'; +import { UserThrottlerGuard } from './user-throttler.guard'; + +const guard = new UserThrottlerGuard( + [], + new ThrottlerStorageService(), + new Reflector(), +); + +function getTracker(req: Record): Promise { + return guard['getTracker'](req); +} + +describe('UserThrottlerGuard.getTracker', () => { + it('tracks a signed-in request by its user id', async () => { + const req = { + ip: '203.0.113.7', + user: { + user: { id: 'user_1' }, + workspace: { id: 'ws_1' }, + authType: JwtType.ACCESS, + }, + }; + + await expect(getTracker(req)).resolves.toBe('user:user_1'); + }); + + it('falls back to the client ip when nobody is signed in', async () => { + const req = { ip: '203.0.113.7', user: null }; + + await expect(getTracker(req)).resolves.toBe('203.0.113.7'); + }); + + it('falls back to the socket address when the ip is empty', async () => { + const req = { + ip: '', + user: null, + socket: { remoteAddress: '198.51.100.4' }, + }; + + await expect(getTracker(req)).resolves.toBe('198.51.100.4'); + }); + + it('uses a constant tracker for an unidentifiable client', async () => { + const req = { ip: '', user: null, socket: {} }; + + await expect(getTracker(req)).resolves.toBe('unknown'); + }); +}); diff --git a/apps/server/src/integrations/throttle/user-throttler.guard.ts b/apps/server/src/integrations/throttle/user-throttler.guard.ts index 35744c094..379cf37d5 100644 --- a/apps/server/src/integrations/throttle/user-throttler.guard.ts +++ b/apps/server/src/integrations/throttle/user-throttler.guard.ts @@ -1,13 +1,19 @@ import { Injectable } from '@nestjs/common'; import { ThrottlerGuard } from '@nestjs/throttler'; -type AuthedRequest = { user?: { id?: string } }; +type AuthedRequest = { + user?: { user?: { id?: string } } | null; + socket?: { remoteAddress?: string }; +}; @Injectable() export class UserThrottlerGuard extends ThrottlerGuard { protected async getTracker(req: AuthedRequest): Promise { - const userId = req.user?.id; + const userId = req.user?.user?.id; if (userId) return `user:${userId}`; - return super.getTracker(req as Parameters[0]); + const ip = await super.getTracker( + req as Parameters[0], + ); + return ip || req.socket?.remoteAddress || 'unknown'; } }