From e4edf7195519112e0d3fa659ded5580aa699f29e Mon Sep 17 00:00:00 2001 From: diyaa Date: Sun, 28 Jun 2026 02:35:15 +0200 Subject: [PATCH] Add refresh token foundation --- .../migration.sql | 3 + backend/prisma/schema.prisma | 3 +- backend/src/modules/auth/auth.module.ts | 3 + .../auth/refresh-token.service.spec.ts | 316 ++++++++++++++++++ .../src/modules/auth/refresh-token.service.ts | 220 ++++++++++++ backend/src/modules/auth/session.service.ts | 1 + 6 files changed, 545 insertions(+), 1 deletion(-) create mode 100644 backend/prisma/migrations/20260628140000_milestone1173_refresh_token_foundation/migration.sql create mode 100644 backend/src/modules/auth/refresh-token.service.spec.ts create mode 100644 backend/src/modules/auth/refresh-token.service.ts diff --git a/backend/prisma/migrations/20260628140000_milestone1173_refresh_token_foundation/migration.sql b/backend/prisma/migrations/20260628140000_milestone1173_refresh_token_foundation/migration.sql new file mode 100644 index 0000000..df53ed3 --- /dev/null +++ b/backend/prisma/migrations/20260628140000_milestone1173_refresh_token_foundation/migration.sql @@ -0,0 +1,3 @@ +ALTER TABLE "account_sessions" +ALTER COLUMN "refresh_token_hash" DROP NOT NULL, +ADD COLUMN "last_seen_at" TIMESTAMP(3); diff --git a/backend/prisma/schema.prisma b/backend/prisma/schema.prisma index 9e263bb..04f029a 100644 --- a/backend/prisma/schema.prisma +++ b/backend/prisma/schema.prisma @@ -286,11 +286,12 @@ model OwnershipTransfer { model AccountSession { id String @id @default(uuid()) @db.Uuid userId String @map("user_id") @db.Uuid - refreshTokenHash String @unique @map("refresh_token_hash") + refreshTokenHash String? @unique @map("refresh_token_hash") status AccountSessionStatus @default(ACTIVE) accessTokenVersion Int @default(1) @map("access_token_version") expiresAt DateTime @map("expires_at") lastRefreshedAt DateTime? @map("last_refreshed_at") + lastSeenAt DateTime? @map("last_seen_at") revokedAt DateTime? @map("revoked_at") createdAt DateTime @default(now()) @map("created_at") updatedAt DateTime @updatedAt @map("updated_at") diff --git a/backend/src/modules/auth/auth.module.ts b/backend/src/modules/auth/auth.module.ts index ccdf99b..629b729 100644 --- a/backend/src/modules/auth/auth.module.ts +++ b/backend/src/modules/auth/auth.module.ts @@ -7,6 +7,7 @@ import { DeviceAuthGuard } from './device-auth.guard'; import { DeviceAuthService } from './device-auth.service'; import { OptionalDeviceAuthGuard } from './optional-device-auth.guard'; import { ProtectedDeviceAuthMiddleware } from './protected-device-auth.middleware'; +import { RefreshTokenService } from './refresh-token.service'; import { SessionService } from './session.service'; @Module({ @@ -17,6 +18,7 @@ import { SessionService } from './session.service'; DeviceAuthGuard, OptionalDeviceAuthGuard, ProtectedDeviceAuthMiddleware, + RefreshTokenService, SessionService, ], exports: [ @@ -25,6 +27,7 @@ import { SessionService } from './session.service'; DeviceAuthGuard, OptionalDeviceAuthGuard, ProtectedDeviceAuthMiddleware, + RefreshTokenService, SessionService, ], }) diff --git a/backend/src/modules/auth/refresh-token.service.spec.ts b/backend/src/modules/auth/refresh-token.service.spec.ts new file mode 100644 index 0000000..3eddc52 --- /dev/null +++ b/backend/src/modules/auth/refresh-token.service.spec.ts @@ -0,0 +1,316 @@ +import { createHash, randomUUID } from 'node:crypto'; +import { UnauthorizedException } from '@nestjs/common'; +import { + AccountSessionStatus, + UserAccountStatus, +} from '@prisma/client'; +import { RefreshTokenService } from './refresh-token.service'; +import { SessionService } from './session.service'; + +describe('RefreshTokenService', () => { + const issuedAt = new Date('2026-06-28T12:00:00.000Z'); + const rotatedAt = new Date('2026-06-28T13:00:00.000Z'); + const expiresAt = new Date('2026-07-28T12:00:00.000Z'); + const activeUserId = randomUUID(); + const lockedUserId = randomUUID(); + const deletedUserId = randomUUID(); + + const hashToken = (token: string): string => + createHash('sha256').update(token, 'utf8').digest('hex'); + + const buildService = () => { + const users = new Map([ + [ + activeUserId, + { + id: activeUserId, + accountStatus: UserAccountStatus.ACTIVE, + }, + ], + [ + lockedUserId, + { + id: lockedUserId, + accountStatus: UserAccountStatus.LOCKED, + }, + ], + [ + deletedUserId, + { + id: deletedUserId, + accountStatus: UserAccountStatus.DELETED, + }, + ], + ]); + const sessions = new Map(); + + const accountSession = { + create: jest.fn().mockImplementation(async ({ data }) => { + const session = { + id: randomUUID(), + userId: data.userId, + refreshTokenHash: data.refreshTokenHash, + status: data.status ?? AccountSessionStatus.ACTIVE, + accessTokenVersion: data.accessTokenVersion ?? 1, + expiresAt: data.expiresAt, + lastRefreshedAt: null, + lastSeenAt: null, + revokedAt: null, + createdAt: issuedAt, + updatedAt: issuedAt, + }; + sessions.set(session.id, session); + return { ...session }; + }), + findUnique: jest.fn().mockImplementation(async ({ where, include }) => { + const session = + (where.id ? sessions.get(where.id) : undefined) ?? + Array.from(sessions.values()).find( + (candidate) => + where.refreshTokenHash !== undefined && + candidate.refreshTokenHash === where.refreshTokenHash, + ) ?? + null; + + if (!session) { + return null; + } + + if (include?.user) { + return { + ...session, + user: users.get(session.userId), + }; + } + + return { ...session }; + }), + updateMany: jest.fn().mockImplementation(async ({ where, data }) => { + let count = 0; + + for (const session of sessions.values()) { + const user = users.get(session.userId); + const matches = + (where.id === undefined || session.id === where.id) && + (where.refreshTokenHash === undefined || + session.refreshTokenHash === where.refreshTokenHash) && + (where.status === undefined || session.status === where.status) && + (where.expiresAt?.gt === undefined || + session.expiresAt.getTime() > where.expiresAt.gt.getTime()) && + (where.user?.accountStatus === undefined || + user?.accountStatus === where.user.accountStatus); + + if (!matches) { + continue; + } + + for (const [key, value] of Object.entries(data)) { + if (value && typeof value === 'object' && 'increment' in value) { + session[key] += (value as { increment: number }).increment; + } else { + session[key] = value; + } + } + session.updatedAt = rotatedAt; + count += 1; + } + + return { count }; + }), + }; + + const prismaService = { + user: { + findUnique: jest.fn().mockImplementation(async ({ where }) => + users.get(where.id) ?? null, + ), + }, + accountSession, + } as any; + const sessionService = new SessionService(prismaService); + const service = new RefreshTokenService( + prismaService, + sessionService, + ); + + const addSession = ( + token: string, + overrides: Record = {}, + ) => { + const session = { + id: randomUUID(), + userId: activeUserId, + refreshTokenHash: hashToken(token), + status: AccountSessionStatus.ACTIVE, + accessTokenVersion: 1, + expiresAt, + lastRefreshedAt: null, + lastSeenAt: null, + revokedAt: null, + createdAt: issuedAt, + updatedAt: issuedAt, + ...overrides, + }; + sessions.set(session.id, session); + return session; + }; + + return { + service, + prismaService, + sessions, + addSession, + }; + }; + + it('issues an opaque random refresh token and stores only its hash', async () => { + const { service, sessions } = buildService(); + + const refreshToken = await service.issueRefreshToken( + { + userId: activeUserId, + expiresAt, + }, + issuedAt, + ); + const [session] = Array.from(sessions.values()); + + expect(refreshToken).toMatch(/^[A-Za-z0-9_-]{43}$/); + expect(session.refreshTokenHash).toBe(hashToken(refreshToken)); + expect(session.refreshTokenHash).not.toContain(refreshToken); + }); + + it('verifies a refresh token without exposing its hash', async () => { + const { service } = buildService(); + const refreshToken = await service.issueRefreshToken( + { + userId: activeUserId, + expiresAt, + }, + issuedAt, + ); + + const session = await service.verifyRefreshToken( + refreshToken, + issuedAt, + ); + + expect(session).toMatchObject({ + userId: activeUserId, + status: AccountSessionStatus.ACTIVE, + accessTokenVersion: 1, + }); + expect(session).not.toHaveProperty('refreshTokenHash'); + expect(session).not.toHaveProperty('user'); + }); + + it('rotates the token, invalidates the previous token, and advances the session', async () => { + const { service, sessions } = buildService(); + const refreshToken = await service.issueRefreshToken( + { + userId: activeUserId, + expiresAt, + }, + issuedAt, + ); + + const nextRefreshToken = await service.rotateRefreshToken( + refreshToken, + rotatedAt, + ); + const [session] = Array.from(sessions.values()); + + expect(nextRefreshToken).not.toBe(refreshToken); + expect(session.refreshTokenHash).toBe(hashToken(nextRefreshToken)); + expect(session.accessTokenVersion).toBe(2); + expect(session.lastRefreshedAt).toEqual(rotatedAt); + expect(session.lastSeenAt).toEqual(rotatedAt); + await expect( + service.verifyRefreshToken(refreshToken, rotatedAt), + ).rejects.toThrow(UnauthorizedException); + await expect( + service.verifyRefreshToken(nextRefreshToken, rotatedAt), + ).resolves.toMatchObject({ id: session.id }); + }); + + it('rejects a refresh token for a revoked session', async () => { + const { service, addSession } = buildService(); + const refreshToken = 'revoked-refresh-token'; + addSession(refreshToken, { + status: AccountSessionStatus.REVOKED, + revokedAt: issuedAt, + }); + + await expect( + service.verifyRefreshToken(refreshToken, issuedAt), + ).rejects.toThrow(UnauthorizedException); + }); + + it('rejects a refresh token for an expired session', async () => { + const { service, addSession } = buildService(); + const refreshToken = 'expired-refresh-token'; + addSession(refreshToken, { + expiresAt: new Date(issuedAt.getTime() - 1), + }); + + await expect( + service.verifyRefreshToken(refreshToken, issuedAt), + ).rejects.toThrow(UnauthorizedException); + }); + + it.each([ + ['locked', lockedUserId], + ['deleted', deletedUserId], + ])('rejects a refresh token for a %s user', async (_label, userId) => { + const { service, addSession } = buildService(); + const refreshToken = `${_label}-refresh-token`; + addSession(refreshToken, { userId }); + + await expect( + service.verifyRefreshToken(refreshToken, issuedAt), + ).rejects.toThrow(UnauthorizedException); + }); + + it('rejects invalid tokens, hash mismatches, and missing sessions', async () => { + const { service, addSession } = buildService(); + const refreshToken = 'valid-refresh-token'; + const session = addSession(refreshToken); + + await expect( + service.verifyRefreshToken('invalid-refresh-token', issuedAt), + ).rejects.toThrow(UnauthorizedException); + await expect( + service.assertRefreshTokenSession( + session.id, + 'different-refresh-token', + issuedAt, + ), + ).rejects.toThrow(UnauthorizedException); + await expect( + service.assertRefreshTokenSession( + randomUUID(), + refreshToken, + issuedAt, + ), + ).rejects.toThrow(UnauthorizedException); + }); + + it('clears the hash and revokes the session without deleting history', async () => { + const { service, sessions, addSession } = buildService(); + const refreshToken = 'refresh-token-to-revoke'; + const session = addSession(refreshToken); + + await service.revokeRefreshToken(refreshToken, rotatedAt); + + expect(sessions.get(session.id)).toMatchObject({ + id: session.id, + refreshTokenHash: null, + status: AccountSessionStatus.REVOKED, + revokedAt: rotatedAt, + }); + expect(sessions.size).toBe(1); + await expect( + service.verifyRefreshToken(refreshToken, rotatedAt), + ).rejects.toThrow(UnauthorizedException); + }); +}); diff --git a/backend/src/modules/auth/refresh-token.service.ts b/backend/src/modules/auth/refresh-token.service.ts new file mode 100644 index 0000000..ba550b1 --- /dev/null +++ b/backend/src/modules/auth/refresh-token.service.ts @@ -0,0 +1,220 @@ +import { Injectable, UnauthorizedException } from '@nestjs/common'; +import { + AccountSession, + AccountSessionStatus, + UserAccountStatus, +} from '@prisma/client'; +import { createHash, randomBytes, timingSafeEqual } from 'node:crypto'; +import { PrismaService } from '../../infrastructure/database/prisma.service'; +import { SessionService } from './session.service'; + +const REFRESH_TOKEN_BYTES = 32; + +export interface IssueRefreshTokenInput { + userId: string; + expiresAt: Date; +} + +export type VerifiedRefreshTokenSession = Omit< + AccountSession, + 'refreshTokenHash' +>; + +type RefreshTokenSessionRecord = AccountSession & { + user: { + accountStatus: UserAccountStatus; + }; +}; + +@Injectable() +export class RefreshTokenService { + constructor( + private readonly prismaService: PrismaService, + private readonly sessionService: SessionService, + ) {} + + async issueRefreshToken( + input: IssueRefreshTokenInput, + now = new Date(), + ): Promise { + const refreshToken = this.generateRefreshToken(); + + await this.sessionService.createSession( + { + userId: input.userId, + refreshTokenHash: this.hashRefreshToken(refreshToken), + expiresAt: input.expiresAt, + }, + now, + ); + + return refreshToken; + } + + async verifyRefreshToken( + refreshToken: string, + now = new Date(), + ): Promise { + const refreshTokenHash = this.hashRefreshToken(refreshToken); + const session = await this.prismaService.accountSession.findUnique({ + where: { + refreshTokenHash, + }, + include: { + user: { + select: { + accountStatus: true, + }, + }, + }, + }); + + this.assertSessionRecord(session, refreshTokenHash, now); + return this.withoutRefreshTokenHash(session); + } + + async rotateRefreshToken( + refreshToken: string, + now = new Date(), + ): Promise { + const session = await this.verifyRefreshToken(refreshToken, now); + const nextRefreshToken = this.generateRefreshToken(); + + await this.sessionService.refreshSession( + { + sessionId: session.id, + refreshTokenHash: this.hashRefreshToken(refreshToken), + nextRefreshTokenHash: this.hashRefreshToken(nextRefreshToken), + }, + now, + ); + + return nextRefreshToken; + } + + async revokeRefreshToken( + refreshToken: string, + now = new Date(), + ): Promise { + const refreshTokenHash = this.hashRefreshToken(refreshToken); + const session = await this.prismaService.accountSession.findUnique({ + where: { + refreshTokenHash, + }, + }); + + if ( + !session || + !this.refreshTokenHashesMatch( + refreshTokenHash, + session.refreshTokenHash, + ) + ) { + throw this.invalidRefreshToken(); + } + + const revocation = await this.prismaService.accountSession.updateMany({ + where: { + id: session.id, + refreshTokenHash, + }, + data: { + refreshTokenHash: null, + ...(session.status === AccountSessionStatus.ACTIVE + ? { + status: AccountSessionStatus.REVOKED, + revokedAt: now, + } + : {}), + }, + }); + + if (revocation.count !== 1) { + throw this.invalidRefreshToken(); + } + } + + async assertRefreshTokenSession( + sessionId: string, + refreshToken: string, + now = new Date(), + ): Promise { + const refreshTokenHash = this.hashRefreshToken(refreshToken); + const session = await this.prismaService.accountSession.findUnique({ + where: { + id: sessionId, + }, + include: { + user: { + select: { + accountStatus: true, + }, + }, + }, + }); + + this.assertSessionRecord(session, refreshTokenHash, now); + return this.withoutRefreshTokenHash(session); + } + + private generateRefreshToken(): string { + return randomBytes(REFRESH_TOKEN_BYTES).toString('base64url'); + } + + private hashRefreshToken(refreshToken: string): string { + if (typeof refreshToken !== 'string' || refreshToken.length === 0) { + throw this.invalidRefreshToken(); + } + + return createHash('sha256').update(refreshToken, 'utf8').digest('hex'); + } + + private assertSessionRecord( + session: RefreshTokenSessionRecord | null, + refreshTokenHash: string, + now: Date, + ): asserts session is RefreshTokenSessionRecord { + if ( + !session || + session.status !== AccountSessionStatus.ACTIVE || + session.expiresAt.getTime() <= now.getTime() || + session.user.accountStatus !== UserAccountStatus.ACTIVE || + !this.refreshTokenHashesMatch( + refreshTokenHash, + session.refreshTokenHash, + ) + ) { + throw this.invalidRefreshToken(); + } + } + + private refreshTokenHashesMatch( + suppliedHash: string, + storedHash: string | null, + ): boolean { + if (!storedHash) { + return false; + } + + const supplied = Buffer.from(suppliedHash, 'utf8'); + const stored = Buffer.from(storedHash, 'utf8'); + + return supplied.length === stored.length && timingSafeEqual(supplied, stored); + } + + private withoutRefreshTokenHash( + session: RefreshTokenSessionRecord, + ): VerifiedRefreshTokenSession { + const { + refreshTokenHash: _refreshTokenHash, + user: _user, + ...safeSession + } = session; + + return safeSession; + } + + private invalidRefreshToken(): UnauthorizedException { + return new UnauthorizedException('Invalid refresh token'); + } +} diff --git a/backend/src/modules/auth/session.service.ts b/backend/src/modules/auth/session.service.ts index a610556..12ece0b 100644 --- a/backend/src/modules/auth/session.service.ts +++ b/backend/src/modules/auth/session.service.ts @@ -122,6 +122,7 @@ export class SessionService { increment: 1, }, lastRefreshedAt: now, + lastSeenAt: now, ...(input.expiresAt ? { expiresAt: input.expiresAt } : {}), }, });