refactor(server): use session model in auth service (#9660)

This commit is contained in:
fengmk2
2025-01-14 07:43:26 +00:00
parent ee99b0cc9d
commit b7635c8944
9 changed files with 116 additions and 152 deletions

View File

@@ -6,6 +6,7 @@ import request from 'supertest';
import { AuthModule, CurrentUser, Public, Session } from '../../core/auth'; import { AuthModule, CurrentUser, Public, Session } from '../../core/auth';
import { AuthService } from '../../core/auth/service'; import { AuthService } from '../../core/auth/service';
import { Models } from '../../models';
import { createTestingApp } from '../utils'; import { createTestingApp } from '../utils';
@Controller('/') @Controller('/')
@@ -35,6 +36,8 @@ let server!: any;
let auth!: AuthService; let auth!: AuthService;
let u1!: CurrentUser; let u1!: CurrentUser;
let sessionId = '';
test.before(async t => { test.before(async t => {
const { app } = await createTestingApp({ const { app } = await createTestingApp({
imports: [AuthModule], imports: [AuthModule],
@@ -44,13 +47,10 @@ test.before(async t => {
auth = app.get(AuthService); auth = app.get(AuthService);
u1 = await auth.signUp('u1@affine.pro', '1'); u1 = await auth.signUp('u1@affine.pro', '1');
const db = app.get(PrismaClient); const models = app.get(Models);
await db.session.create({ const session = await models.session.createSession();
data: { sessionId = session.id;
id: '1', await auth.createUserSession(u1.id, sessionId);
},
});
await auth.createUserSession(u1.id, '1');
server = app.getHttpServer(); server = app.getHttpServer();
t.context.app = app; t.context.app = app;
@@ -69,7 +69,7 @@ test('should be able to visit public api if not signed in', async t => {
test('should be able to visit public api if signed in', async t => { test('should be able to visit public api if signed in', async t => {
const res = await request(server) const res = await request(server)
.get('/public') .get('/public')
.set('Cookie', `${AuthService.sessionCookieName}=1`) .set('Cookie', `${AuthService.sessionCookieName}=${sessionId}`)
.expect(HttpStatus.OK); .expect(HttpStatus.OK);
t.is(res.body.user.id, u1.id); t.is(res.body.user.id, u1.id);
@@ -90,7 +90,7 @@ test('should not be able to visit private api if not signed in', async t => {
test('should be able to visit private api if signed in', async t => { test('should be able to visit private api if signed in', async t => {
const res = await request(server) const res = await request(server)
.get('/private') .get('/private')
.set('Cookie', `${AuthService.sessionCookieName}=1`) .set('Cookie', `${AuthService.sessionCookieName}=${sessionId}`)
.expect(HttpStatus.OK); .expect(HttpStatus.OK);
t.is(res.body.user.id, u1.id); t.is(res.body.user.id, u1.id);
@@ -100,10 +100,10 @@ test('should be able to parse session cookie', async t => {
const spy = Sinon.spy(auth, 'getUserSession'); const spy = Sinon.spy(auth, 'getUserSession');
await request(server) await request(server)
.get('/public') .get('/public')
.set('cookie', `${AuthService.sessionCookieName}=1`) .set('cookie', `${AuthService.sessionCookieName}=${sessionId}`)
.expect(200); .expect(200);
t.deepEqual(spy.firstCall.args, ['1', undefined]); t.deepEqual(spy.firstCall.args, [sessionId, undefined]);
spy.restore(); spy.restore();
}); });
@@ -112,17 +112,17 @@ test('should be able to parse bearer token', async t => {
await request(server) await request(server)
.get('/public') .get('/public')
.auth('1', { type: 'bearer' }) .auth(sessionId, { type: 'bearer' })
.expect(200); .expect(200);
t.deepEqual(spy.firstCall.args, ['1', undefined]); t.deepEqual(spy.firstCall.args, [sessionId, undefined]);
spy.restore(); spy.restore();
}); });
test('should be able to refresh session if needed', async t => { test('should be able to refresh session if needed', async t => {
await t.context.app.get(PrismaClient).userSession.updateMany({ await t.context.app.get(PrismaClient).userSession.updateMany({
where: { where: {
sessionId: '1', sessionId,
}, },
data: { data: {
expiresAt: new Date(Date.now() + 1000 * 60 * 60 /* expires in 1 hour */), expiresAt: new Date(Date.now() + 1000 * 60 * 60 /* expires in 1 hour */),
@@ -131,7 +131,7 @@ test('should be able to refresh session if needed', async t => {
const res = await request(server) const res = await request(server)
.get('/session') .get('/session')
.set('cookie', `${AuthService.sessionCookieName}=1`) .set('cookie', `${AuthService.sessionCookieName}=${sessionId}`)
.expect(200); .expect(200);
const cookie = res const cookie = res

View File

@@ -0,0 +1,47 @@
import { ScheduleModule } from '@nestjs/schedule';
import { TestingModule } from '@nestjs/testing';
import { PrismaClient } from '@prisma/client';
import test from 'ava';
import { AuthModule, AuthService } from '../../core/auth';
import { AuthCronJob } from '../../core/auth/job';
import { createTestingModule } from '../utils';
let m: TestingModule;
let db: PrismaClient;
test.before(async () => {
m = await createTestingModule({
imports: [ScheduleModule.forRoot(), AuthModule],
});
db = m.get(PrismaClient);
});
test.after.always(async () => {
await m.close();
});
test('should clean expired user sessions', async t => {
const auth = m.get(AuthService);
const job = m.get(AuthCronJob);
const user1 = await auth.signUp('u1@affine.pro', '1');
const user2 = await auth.signUp('u2@affine.pro', '1');
await auth.createUserSession(user1.id);
await auth.createUserSession(user2.id);
let userSessions = await db.userSession.findMany();
t.is(userSessions.length, 2);
// no expired sessions
await job.cleanExpiredUserSessions();
userSessions = await db.userSession.findMany();
t.is(userSessions.length, 2);
// clean all expired sessions
await db.userSession.updateMany({
data: { expiresAt: new Date(Date.now() - 1000) },
});
await job.cleanExpiredUserSessions();
userSessions = await db.userSession.findMany();
t.is(userSessions.length, 0);
});

View File

@@ -192,8 +192,10 @@ test('should be able to signout multi accounts session', async t => {
const session = await auth.createSession(); const session = await auth.createSession();
await auth.createUserSession(u1.id, session.id); const userSession1 = await auth.createUserSession(u1.id, session.id);
await auth.createUserSession(u2.id, session.id); const userSession2 = await auth.createUserSession(u2.id, session.id);
t.not(userSession1.id, userSession2.id);
t.is(userSession1.sessionId, userSession2.sessionId);
await auth.signOut(session.id, u1.id); await auth.signOut(session.id, u1.id);

View File

@@ -7,6 +7,7 @@ import { QuotaModule } from '../quota';
import { UserModule } from '../user'; import { UserModule } from '../user';
import { AuthController } from './controller'; import { AuthController } from './controller';
import { AuthGuard, AuthWebsocketOptionsProvider } from './guard'; import { AuthGuard, AuthWebsocketOptionsProvider } from './guard';
import { AuthCronJob } from './job';
import { AuthResolver } from './resolver'; import { AuthResolver } from './resolver';
import { AuthService } from './service'; import { AuthService } from './service';
@@ -16,6 +17,7 @@ import { AuthService } from './service';
AuthService, AuthService,
AuthResolver, AuthResolver,
AuthGuard, AuthGuard,
AuthCronJob,
AuthWebsocketOptionsProvider, AuthWebsocketOptionsProvider,
], ],
exports: [AuthService, AuthGuard, AuthWebsocketOptionsProvider], exports: [AuthService, AuthGuard, AuthWebsocketOptionsProvider],

View File

@@ -0,0 +1,14 @@
import { Injectable } from '@nestjs/common';
import { Cron, CronExpression } from '@nestjs/schedule';
import { Models } from '../../models';
@Injectable()
export class AuthCronJob {
constructor(private readonly models: Models) {}
@Cron(CronExpression.EVERY_DAY_AT_MIDNIGHT)
async cleanExpiredUserSessions() {
await this.models.session.cleanExpiredUserSessions();
}
}

View File

@@ -1,11 +1,9 @@
import { Injectable, OnApplicationBootstrap } from '@nestjs/common'; import { Injectable, OnApplicationBootstrap } from '@nestjs/common';
import { Cron, CronExpression } from '@nestjs/schedule';
import type { User, UserSession } from '@prisma/client';
import { PrismaClient } from '@prisma/client';
import type { CookieOptions, Request, Response } from 'express'; import type { CookieOptions, Request, Response } from 'express';
import { assign, pick } from 'lodash-es'; import { assign, pick } from 'lodash-es';
import { Config, MailService, SignUpForbidden } from '../../base'; import { Config, MailService, SignUpForbidden } from '../../base';
import { Models, type User, type UserSession } from '../../models';
import { FeatureManagementService } from '../features/management'; import { FeatureManagementService } from '../features/management';
import { QuotaService } from '../quota/service'; import { QuotaService } from '../quota/service';
import { QuotaType } from '../quota/types'; import { QuotaType } from '../quota/types';
@@ -46,7 +44,7 @@ export class AuthService implements OnApplicationBootstrap {
constructor( constructor(
private readonly config: Config, private readonly config: Config,
private readonly db: PrismaClient, private readonly models: Models,
private readonly mailer: MailService, private readonly mailer: MailService,
private readonly feature: FeatureManagementService, private readonly feature: FeatureManagementService,
private readonly quota: QuotaService, private readonly quota: QuotaService,
@@ -103,18 +101,9 @@ export class AuthService implements OnApplicationBootstrap {
async signOut(sessionId: string, userId?: string) { async signOut(sessionId: string, userId?: string) {
// sign out all users in the session // sign out all users in the session
if (!userId) { if (!userId) {
await this.db.session.deleteMany({ await this.models.session.deleteSession(sessionId);
where: {
id: sessionId,
},
});
} else { } else {
await this.db.userSession.deleteMany({ await this.models.session.deleteUserSession(userId, sessionId);
where: {
sessionId,
userId,
},
});
} }
} }
@@ -138,7 +127,7 @@ export class AuthService implements OnApplicationBootstrap {
// fallback to the first valid session if user provided userId is invalid // fallback to the first valid session if user provided userId is invalid
if (!userSession) { if (!userSession) {
// checked // checked
// eslint-disable-next-line @typescript-eslint/no-non-null-assertion // oxlint-disable-next-line @typescript-eslint/no-non-null-assertion
userSession = sessions.at(-1)!; userSession = sessions.at(-1)!;
} }
@@ -152,127 +141,50 @@ export class AuthService implements OnApplicationBootstrap {
} }
async getUserSessions(sessionId: string) { async getUserSessions(sessionId: string) {
return this.db.userSession.findMany({ return await this.models.session.findUserSessionsBySessionId(sessionId);
where: {
sessionId,
OR: [{ expiresAt: { gt: new Date() } }, { expiresAt: null }],
},
orderBy: {
createdAt: 'asc',
},
});
} }
async createUserSession( async createUserSession(userId: string, sessionId?: string, ttl?: number) {
userId: string, return await this.models.session.createOrRefreshUserSession(
sessionId?: string, userId,
ttl = this.config.auth.session.ttl sessionId,
) { ttl
// check whether given session is valid );
if (sessionId) {
const session = await this.db.session.findFirst({
where: {
id: sessionId,
},
});
if (!session) {
sessionId = undefined;
}
}
if (!sessionId) {
const session = await this.createSession();
sessionId = session.id;
}
const expiresAt = new Date(Date.now() + ttl * 1000);
return this.db.userSession.upsert({
where: {
sessionId_userId: {
sessionId,
userId,
},
},
update: {
expiresAt,
},
create: {
sessionId,
userId,
expiresAt,
},
});
} }
async getUserList(sessionId: string) { async getUserList(sessionId: string) {
const sessions = await this.db.userSession.findMany({ const sessions = await this.models.session.findUserSessionsBySessionId(
where: { sessionId,
sessionId, {
OR: [
{
expiresAt: null,
},
{
expiresAt: {
gt: new Date(),
},
},
],
},
include: {
user: true, user: true,
}, }
orderBy: { );
createdAt: 'asc',
},
});
return sessions.map(({ user }) => sessionUser(user)); return sessions.map(({ user }) => sessionUser(user));
} }
async createSession() { async createSession() {
return this.db.session.create({ return await this.models.session.createSession();
data: {},
});
} }
async getSession(sessionId: string) { async getSession(sessionId: string) {
return this.db.session.findFirst({ return await this.models.session.getSession(sessionId);
where: {
id: sessionId,
},
});
} }
async refreshUserSessionIfNeeded( async refreshUserSessionIfNeeded(
res: Response, res: Response,
session: UserSession, userSession: UserSession,
ttr = this.config.auth.session.ttr ttr?: number
): Promise<boolean> { ): Promise<boolean> {
if ( const newExpiresAt = await this.models.session.refreshUserSessionIfNeeded(
session.expiresAt && userSession,
session.expiresAt.getTime() - Date.now() > ttr * 1000 ttr
) { );
if (!newExpiresAt) {
// no need to refresh // no need to refresh
return false; return false;
} }
const newExpiresAt = new Date( res.cookie(AuthService.sessionCookieName, userSession.sessionId, {
Date.now() + this.config.auth.session.ttl * 1000
);
await this.db.userSession.update({
where: {
id: session.id,
},
data: {
expiresAt: newExpiresAt,
},
});
res.cookie(AuthService.sessionCookieName, session.sessionId, {
expires: newExpiresAt, expires: newExpiresAt,
...this.cookieOptions, ...this.cookieOptions,
}); });
@@ -281,11 +193,7 @@ export class AuthService implements OnApplicationBootstrap {
} }
async revokeUserSessions(userId: string) { async revokeUserSessions(userId: string) {
return this.db.userSession.deleteMany({ return await this.models.session.deleteUserSession(userId);
where: {
userId,
},
});
} }
getSessionOptionsFromRequest(req: Request) { getSessionOptionsFromRequest(req: Request) {
@@ -425,15 +333,4 @@ export class AuthService implements OnApplicationBootstrap {
to: email, to: email,
}); });
} }
@Cron(CronExpression.EVERY_DAY_AT_MIDNIGHT)
async cleanExpiredSessions() {
await this.db.userSession.deleteMany({
where: {
expiresAt: {
lte: new Date(),
},
},
});
}
} }

View File

@@ -1,8 +1,8 @@
import type { ExecutionContext } from '@nestjs/common'; import type { ExecutionContext } from '@nestjs/common';
import { createParamDecorator } from '@nestjs/common'; import { createParamDecorator } from '@nestjs/common';
import { User, UserSession } from '@prisma/client';
import { getRequestResponseFromContext } from '../../base'; import { getRequestResponseFromContext } from '../../base';
import type { User, UserSession } from '../../models';
/** /**
* Used to fetch current user from the request context. * Used to fetch current user from the request context.
@@ -37,7 +37,7 @@ import { getRequestResponseFromContext } from '../../base';
* ``` * ```
*/ */
// interface and variable don't conflict // interface and variable don't conflict
// eslint-disable-next-line no-redeclare // oxlint-disable-next-line no-redeclare
export const CurrentUser = createParamDecorator( export const CurrentUser = createParamDecorator(
(_: unknown, context: ExecutionContext) => { (_: unknown, context: ExecutionContext) => {
return getRequestResponseFromContext(context).req.session?.user; return getRequestResponseFromContext(context).req.session?.user;
@@ -51,7 +51,7 @@ export interface CurrentUser
} }
// interface and variable don't conflict // interface and variable don't conflict
// eslint-disable-next-line no-redeclare // oxlint-disable-next-line no-redeclare
export const Session = createParamDecorator( export const Session = createParamDecorator(
(_: unknown, context: ExecutionContext) => { (_: unknown, context: ExecutionContext) => {
return getRequestResponseFromContext(context).req.session; return getRequestResponseFromContext(context).req.session;

View File

@@ -4,6 +4,8 @@ import { SessionModel } from './session';
import { UserModel } from './user'; import { UserModel } from './user';
import { VerificationTokenModel } from './verification-token'; import { VerificationTokenModel } from './verification-token';
export * from './session';
export * from './user';
export * from './verification-token'; export * from './verification-token';
const models = [UserModel, SessionModel, VerificationTokenModel] as const; const models = [UserModel, SessionModel, VerificationTokenModel] as const;

View File

@@ -14,7 +14,7 @@ import {
} from '../base'; } from '../base';
import type { Payload } from '../base/event/def'; import type { Payload } from '../base/event/def';
import { Permission } from '../core/permission'; import { Permission } from '../core/permission';
import { Quota_FreePlanV1_1 } from '../core/quota'; import { Quota_FreePlanV1_1 } from '../core/quota/schema';
const publicUserSelect = { const publicUserSelect = {
id: true, id: true,