Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
203 changes: 203 additions & 0 deletions src/repositories/refreshTokenRepository.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,203 @@
import { newDb } from 'pg-mem';
import fs from 'fs';
import path from 'path';
import crypto from 'crypto';
import { DatabaseRefreshTokenRepository } from './refreshTokenRepository.js';

describe('DatabaseRefreshTokenRepository', () => {
let db: any;
let pool: any;
let repo: DatabaseRefreshTokenRepository;

const userId1 = crypto.randomUUID();
const userId2 = crypto.randomUUID();

beforeAll(() => {
db = newDb();

// Register Postgres native gen_random_uuid() for schema defaults
db.public.registerFunction({
name: 'gen_random_uuid',
returns: 'uuid',
implementation: () => crypto.randomUUID(),
impure: true,
});

// Seed root foreign key dependency
db.public.none(`CREATE TABLE users (id UUID PRIMARY KEY);`);

// Load and execute production migrations in order
const migrationsDir = path.join(process.cwd(), 'migrations');
const baseMigration = fs.readFileSync(path.join(migrationsDir, 'add_refresh_tokens.sql'), 'utf-8');
const familyMigration = fs.readFileSync(path.join(migrationsDir, 'add_refresh_token_family.sql'), 'utf-8');

db.public.none(baseMigration);
db.public.none(familyMigration);

const { Pool } = db.adapters.createPg();
pool = new Pool();
});

beforeEach(async () => {
await pool.query(`DELETE FROM refresh_tokens;`);
await pool.query(`DELETE FROM users;`);
await pool.query(`INSERT INTO users (id) VALUES ($1), ($2);`, [userId1, userId2]);

repo = new DatabaseRefreshTokenRepository(pool);
});

const createValidToken = (userId: string, overrides: any = {}) => ({
userId,
tokenHash: crypto.randomBytes(32).toString('hex'), // Matches constraints length=64
expiresAt: new Date(Date.now() + 86400000), // +1 Day
createdAt: new Date(),
isRevoked: false,
familyId: crypto.randomUUID(),
...overrides
});

it('createRefreshToken: should securely store a new refresh token and return the model', async () => {
const tokenData = createValidToken(userId1);
const token = await repo.createRefreshToken(tokenData);

expect(token.id).toBeDefined();
expect(token.userId).toBe(userId1);
expect(token.tokenHash).toBe(tokenData.tokenHash);
expect(token.isRevoked).toBe(false);
});

it('schema: enforces token_hash length boundary constraints', async () => {
const invalidToken = createValidToken(userId1, { tokenHash: 'short_hash' });
await expect(repo.createRefreshToken(invalidToken)).rejects.toThrow();
});

it('findRefreshTokenById: should retrieve active token and strict-exclude revoked tokens', async () => {
const activeTokenData = createValidToken(userId1);
const activeToken = await repo.createRefreshToken(activeTokenData);

const found = await repo.findRefreshTokenById(activeToken.id as string, userId1);
expect(found).not.toBeNull();
expect(found?.id).toBe(activeToken.id);

await repo.revokeRefreshToken(activeToken.id as string, userId1);

const notFound = await repo.findRefreshTokenById(activeToken.id as string, userId1);
expect(notFound).toBeNull();
});

it('findRefreshTokenByHash: should retrieve active token by hash and strict-exclude revoked tokens', async () => {
const activeTokenData = createValidToken(userId1);
const activeToken = await repo.createRefreshToken(activeTokenData);

const found = await repo.findRefreshTokenByHash(activeTokenData.tokenHash, userId1);
expect(found).not.toBeNull();
expect(found?.id).toBe(activeToken.id);

await repo.revokeRefreshToken(activeToken.id as string, userId1);

const notFound = await repo.findRefreshTokenByHash(activeTokenData.tokenHash, userId1);
expect(notFound).toBeNull();
});

it('updateLastUsed: should update the last_used_at timestamp directly', async () => {
const token = await repo.createRefreshToken(createValidToken(userId1));
expect(token.lastUsedAt).toBeUndefined();

await repo.updateLastUsed(token.id as string, userId1);

const updated = await repo.findRefreshTokenById(token.id as string, userId1);
expect(updated?.lastUsedAt).toBeDefined();
expect(updated?.lastUsedAt?.getTime()).toBeLessThanOrEqual(Date.now());
});

it('revokeRefreshToken: should atomically mark the target token as revoked', async () => {
const token = await repo.createRefreshToken(createValidToken(userId1));
await repo.revokeRefreshToken(token.id as string, userId1);

const result = await pool.query(`SELECT is_revoked FROM refresh_tokens WHERE id = $1`, [token.id]);
expect(result.rows[0].is_revoked).toBe(true);
});

it('revokeFamily: should cascade invalidation to all tokens within a specific token family', async () => {
const familyId = crypto.randomUUID();
const t1 = await repo.createRefreshToken(createValidToken(userId1, { familyId }));
const t2 = await repo.createRefreshToken(createValidToken(userId1, { familyId }));
const otherFamilyToken = await repo.createRefreshToken(createValidToken(userId1));

await repo.revokeFamily(familyId, userId1);

const count = await repo.countActiveTokens(userId1);
expect(count).toBe(1); // Only otherFamilyToken survives

const result = await pool.query(`SELECT is_revoked FROM refresh_tokens WHERE id IN ($1, $2) ORDER BY id`, [t1.id, t2.id]);
expect(result.rows[0].is_revoked).toBe(true);
expect(result.rows[1].is_revoked).toBe(true);
});

it('revokeAllUserTokens: affects only the structurally targeted user', async () => {
const t1 = await repo.createRefreshToken(createValidToken(userId1));
const t2 = await repo.createRefreshToken(createValidToken(userId1));
const target2 = await repo.createRefreshToken(createValidToken(userId2));

await repo.revokeAllUserTokens(userId1);

const count1 = await repo.countActiveTokens(userId1);
const count2 = await repo.countActiveTokens(userId2);

expect(count1).toBe(0);
expect(count2).toBe(1);

const res = await pool.query(`SELECT is_revoked FROM refresh_tokens WHERE id = $1`, [target2.id]);
expect(res.rows[0].is_revoked).toBe(false);
});

it('cleanupExpiredTokens: drops explicitly revoked and chronologically expired tokens', async () => {
const yesterday = new Date(Date.now() - 86400000);
const tomorrow = new Date(Date.now() + 86400000);

await repo.createRefreshToken(createValidToken(userId1, { expiresAt: yesterday })); // Expired
await repo.createRefreshToken(createValidToken(userId1, { isRevoked: true })); // Revoked
const active = await repo.createRefreshToken(createValidToken(userId1, { expiresAt: tomorrow })); // Keeps

const deletedCount = await repo.cleanupExpiredTokens();
expect(deletedCount).toBe(2);

const res = await pool.query(`SELECT id FROM refresh_tokens`);
expect(res.rowCount).toBe(1);
expect(res.rows[0].id).toBe(active.id);
});

it('countActiveTokens: ignores expired and explicitly revoked rows', async () => {
const yesterday = new Date(Date.now() - 86400000);
const tomorrow = new Date(Date.now() + 86400000);

await repo.createRefreshToken(createValidToken(userId1, { expiresAt: tomorrow })); // Active
await repo.createRefreshToken(createValidToken(userId1, { expiresAt: tomorrow })); // Active
await repo.createRefreshToken(createValidToken(userId1, { isRevoked: true })); // Revoked
await repo.createRefreshToken(createValidToken(userId1, { expiresAt: yesterday })); // Expired
await repo.createRefreshToken(createValidToken(userId2, { expiresAt: tomorrow })); // Wrong User

const count = await repo.countActiveTokens(userId1);
expect(count).toBe(2);
});

it('listRefreshTokens: supports cursor keyset pagination over structural ranges', async () => {
const baseDate = Date.now();
for (let i = 0; i < 3; i++) {
await repo.createRefreshToken(createValidToken(userId1, {
createdAt: new Date(baseDate + (i * 1000)) // Force sequential timestamps for deterministic sort
}));
}

const firstPage = await repo.listRefreshTokens(userId1, 2);
expect(firstPage.tokens.length).toBe(2);
expect(firstPage.hasMore).toBe(true);

const lastToken = firstPage.tokens[1];
const cursor = { timestamp: lastToken.createdAt, id: lastToken.id as string };

const secondPage = await repo.listRefreshTokens(userId1, 2, cursor);
expect(secondPage.tokens.length).toBe(1);
expect(secondPage.hasMore).toBe(false);
});
});
80 changes: 6 additions & 74 deletions src/repositories/refreshTokenRepository.ts
Original file line number Diff line number Diff line change
Expand Up @@ -3,86 +3,27 @@ import type { RefreshToken } from '../types/auth.js';
import type { CursorPayload } from '../lib/cursorPagination.js';
import { readQuery, writeQuery } from '../db.js';

/** Injectable queryable for tests. */
export interface RefreshTokenRepositoryQueryable {
query<T = unknown>(text: string, params?: unknown[]): Promise<{ rows: T[]; rowCount?: number | null }>;
}

export interface RefreshTokenRepository {
/**
* Store a new refresh token in the database
*/
createRefreshToken(token: Omit<RefreshToken, 'id'> & { id?: string }): Promise<RefreshToken>;

/**
* Find refresh token by ID and user ID
*/
findRefreshTokenById(tokenId: string, userId: string): Promise<RefreshToken | null>;

/**
* Find refresh token by hash (for verification)
*/
findRefreshTokenByHash(tokenHash: string, userId: string): Promise<RefreshToken | null>;

/**
* Update the last used timestamp for a refresh token
*/
updateLastUsed(tokenId: string, userId: string): Promise<void>;

/**
* Revoke a refresh token
*/
revokeRefreshToken(tokenId: string, userId: string): Promise<void>;

/**
* Revoke all refresh tokens belonging to a token family atomically
*/
revokeFamily(familyId: string, userId: string): Promise<void>;

/**
* Revoke all refresh tokens for a user
*/
revokeAllUserTokens(userId: string): Promise<void>;

/**
* Clean up expired and revoked tokens
*/
cleanupExpiredTokens(): Promise<number>;

/**
* Count active refresh tokens for a user
*/
countActiveTokens(userId: string): Promise<number>;

/**
* List refresh tokens for a user with cursor-based pagination.
* Uses stable keyset ordering over (created_at DESC, id DESC) to
* guarantee consistent results under concurrent writes.
*
* @param userId - The user whose tokens to list
* @param limit - Maximum number of tokens to return (clamped to 1..100)
* @param afterCursor - Optional cursor encoding the last seen (created_at, id)
* @returns - Array of refresh tokens and a hasMore flag
*/
listRefreshTokens(
userId: string,
limit: number,
afterCursor?: CursorPayload,
): Promise<{ tokens: RefreshToken[]; hasMore: boolean }>;
listRefreshTokens(userId: string, limit: number, afterCursor?: CursorPayload): Promise<{ tokens: RefreshToken[]; hasMore: boolean }>;
}

/**
* Database implementation of RefreshTokenRepository
* This should be adapted to your specific database setup
*/
export class DatabaseRefreshTokenRepository implements RefreshTokenRepository {
private readonly readDb: RefreshTokenRepositoryQueryable;
private readonly writeDb: RefreshTokenRepositoryQueryable;

/**
* @param db - Optional injectable queryable (test helper).
* When omitted, reads route to replicas and writes route to the primary.
*/
constructor(db?: RefreshTokenRepositoryQueryable) {
if (db) {
this.readDb = db;
Expand Down Expand Up @@ -123,13 +64,11 @@ export class DatabaseRefreshTokenRepository implements RefreshTokenRepository {
const result = await this.readDb.query(
`SELECT id, user_id, token_hash, expires_at, created_at, last_used_at, is_revoked, family_id
FROM refresh_tokens
WHERE id = $1 AND user_id = $2`,
WHERE id = $1 AND user_id = $2 AND is_revoked = false`,
[tokenId, userId]
);

if (result.rows.length === 0) {
return null;
}
if (result.rows.length === 0) return null;

const row = result.rows[0] as Record<string, unknown>;
return {
Expand All @@ -148,13 +87,11 @@ export class DatabaseRefreshTokenRepository implements RefreshTokenRepository {
const result = await this.readDb.query(
`SELECT id, user_id, token_hash, expires_at, created_at, last_used_at, is_revoked, family_id
FROM refresh_tokens
WHERE token_hash = $1 AND user_id = $2`,
WHERE token_hash = $1 AND user_id = $2 AND is_revoked = false`,
[tokenHash, userId]
);

if (result.rows.length === 0) {
return null;
}
if (result.rows.length === 0) return null;

const row = result.rows[0] as Record<string, unknown>;
return {
Expand Down Expand Up @@ -229,16 +166,11 @@ export class DatabaseRefreshTokenRepository implements RefreshTokenRepository {
limit: number,
afterCursor?: CursorPayload,
): Promise<{ tokens: RefreshToken[]; hasMore: boolean }> {
// Fetch one extra row to determine if there are more results beyond the limit.
const fetchLimit = limit + 1;

let query: string;
let params: unknown[];

if (afterCursor) {
// Keyset pagination: rows strictly after the cursor position.
// Ordering is (created_at DESC, id DESC) so we fetch rows that are
// either created earlier, or same timestamp with a smaller id.
query = `
SELECT id, user_id, token_hash, expires_at, created_at, last_used_at, is_revoked, family_id
FROM refresh_tokens
Expand Down Expand Up @@ -280,4 +212,4 @@ export class DatabaseRefreshTokenRepository implements RefreshTokenRepository {

return { tokens, hasMore };
}
}
}
Loading