Add refresh token foundation

This commit is contained in:
diyaa 2026-06-28 02:35:15 +02:00
parent 4e743df7f1
commit e4edf71955
6 changed files with 545 additions and 1 deletions

View File

@ -0,0 +1,3 @@
ALTER TABLE "account_sessions"
ALTER COLUMN "refresh_token_hash" DROP NOT NULL,
ADD COLUMN "last_seen_at" TIMESTAMP(3);

View File

@ -286,11 +286,12 @@ model OwnershipTransfer {
model AccountSession { model AccountSession {
id String @id @default(uuid()) @db.Uuid id String @id @default(uuid()) @db.Uuid
userId String @map("user_id") @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) status AccountSessionStatus @default(ACTIVE)
accessTokenVersion Int @default(1) @map("access_token_version") accessTokenVersion Int @default(1) @map("access_token_version")
expiresAt DateTime @map("expires_at") expiresAt DateTime @map("expires_at")
lastRefreshedAt DateTime? @map("last_refreshed_at") lastRefreshedAt DateTime? @map("last_refreshed_at")
lastSeenAt DateTime? @map("last_seen_at")
revokedAt DateTime? @map("revoked_at") revokedAt DateTime? @map("revoked_at")
createdAt DateTime @default(now()) @map("created_at") createdAt DateTime @default(now()) @map("created_at")
updatedAt DateTime @updatedAt @map("updated_at") updatedAt DateTime @updatedAt @map("updated_at")

View File

@ -7,6 +7,7 @@ import { DeviceAuthGuard } from './device-auth.guard';
import { DeviceAuthService } from './device-auth.service'; import { DeviceAuthService } from './device-auth.service';
import { OptionalDeviceAuthGuard } from './optional-device-auth.guard'; import { OptionalDeviceAuthGuard } from './optional-device-auth.guard';
import { ProtectedDeviceAuthMiddleware } from './protected-device-auth.middleware'; import { ProtectedDeviceAuthMiddleware } from './protected-device-auth.middleware';
import { RefreshTokenService } from './refresh-token.service';
import { SessionService } from './session.service'; import { SessionService } from './session.service';
@Module({ @Module({
@ -17,6 +18,7 @@ import { SessionService } from './session.service';
DeviceAuthGuard, DeviceAuthGuard,
OptionalDeviceAuthGuard, OptionalDeviceAuthGuard,
ProtectedDeviceAuthMiddleware, ProtectedDeviceAuthMiddleware,
RefreshTokenService,
SessionService, SessionService,
], ],
exports: [ exports: [
@ -25,6 +27,7 @@ import { SessionService } from './session.service';
DeviceAuthGuard, DeviceAuthGuard,
OptionalDeviceAuthGuard, OptionalDeviceAuthGuard,
ProtectedDeviceAuthMiddleware, ProtectedDeviceAuthMiddleware,
RefreshTokenService,
SessionService, SessionService,
], ],
}) })

View File

@ -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<string, any>([
[
activeUserId,
{
id: activeUserId,
accountStatus: UserAccountStatus.ACTIVE,
},
],
[
lockedUserId,
{
id: lockedUserId,
accountStatus: UserAccountStatus.LOCKED,
},
],
[
deletedUserId,
{
id: deletedUserId,
accountStatus: UserAccountStatus.DELETED,
},
],
]);
const sessions = new Map<string, any>();
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<string, unknown> = {},
) => {
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);
});
});

View File

@ -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<string> {
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<VerifiedRefreshTokenSession> {
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<string> {
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<void> {
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<VerifiedRefreshTokenSession> {
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');
}
}

View File

@ -122,6 +122,7 @@ export class SessionService {
increment: 1, increment: 1,
}, },
lastRefreshedAt: now, lastRefreshedAt: now,
lastSeenAt: now,
...(input.expiresAt ? { expiresAt: input.expiresAt } : {}), ...(input.expiresAt ? { expiresAt: input.expiresAt } : {}),
}, },
}); });