feat(server): support socketio auth field (#8595)

fix AF-1531
This commit is contained in:
forehalo
2024-10-25 03:27:02 +00:00
parent 10963da706
commit 08319bc560
5 changed files with 45 additions and 29 deletions

View File

@@ -7,12 +7,12 @@ import type {
import { Injectable, SetMetadata } from '@nestjs/common'; import { Injectable, SetMetadata } from '@nestjs/common';
import { ModuleRef, Reflector } from '@nestjs/core'; import { ModuleRef, Reflector } from '@nestjs/core';
import type { Request, Response } from 'express'; import type { Request, Response } from 'express';
import { Socket } from 'socket.io';
import { import {
AuthenticationRequired, AuthenticationRequired,
Config, Config,
getRequestResponseFromContext, getRequestResponseFromContext,
mapAnyError,
parseCookies, parseCookies,
} from '../../fundamentals'; } from '../../fundamentals';
import { WEBSOCKET_OPTIONS } from '../../fundamentals/websocket'; import { WEBSOCKET_OPTIONS } from '../../fundamentals/websocket';
@@ -64,9 +64,6 @@ export class AuthGuard implements CanActivate, OnModuleInit {
return req.session; return req.session;
} }
// compatibility with websocket request
parseCookies(req);
// TODO(@forehalo): a cache for user session // TODO(@forehalo): a cache for user session
const userSession = await this.auth.getUserSessionFromRequest(req, res); const userSession = await this.auth.getUserSessionFromRequest(req, res);
@@ -93,27 +90,22 @@ export const AuthWebsocketOptionsProvider: FactoryProvider = {
useFactory: (config: Config, guard: AuthGuard) => { useFactory: (config: Config, guard: AuthGuard) => {
return { return {
...config.websocket, ...config.websocket,
allowRequest: async ( canActivate: async (socket: Socket) => {
req: any, const upgradeReq = socket.client.request as Request;
pass: (err: string | null | undefined, success: boolean) => void const handshake = socket.handshake;
) => {
if (!config.websocket.requireAuthentication) {
return pass(null, true);
}
try { // compatibility with websocket request
const authentication = await guard.signIn(req); parseCookies(upgradeReq);
if (authentication) { upgradeReq.cookies = {
return pass(null, true); [AuthService.sessionCookieName]: handshake.auth.token,
} else { [AuthService.userCookieName]: handshake.auth.userId,
return pass('unauthenticated', false); ...upgradeReq.cookies,
} };
} catch (e) {
const error = mapAnyError(e); const session = await guard.signIn(upgradeReq);
error.log('Websocket');
return pass('unauthenticated', false); return !!session;
}
}, },
}; };
}, },

View File

@@ -298,7 +298,7 @@ export class AuthService implements OnApplicationBootstrap {
const userId: string | undefined = const userId: string | undefined =
req.cookies[AuthService.userCookieName] || req.cookies[AuthService.userCookieName] ||
req.headers[AuthService.userCookieName]; req.headers[AuthService.userCookieName.replaceAll('_', '-')];
return { return {
sessionId, sessionId,

View File

@@ -26,7 +26,7 @@ export function getRequestResponseFromHost(host: ArgumentsHost) {
} }
case 'ws': { case 'ws': {
const ws = host.switchToWs(); const ws = host.switchToWs();
const req = ws.getClient<Socket>().client.conn.request as Request; const req = ws.getClient<Socket>().request as Request;
parseCookies(req); parseCookies(req);
return { req }; return { req };
} }

View File

@@ -1,4 +1,5 @@
import { GatewayMetadata } from '@nestjs/websockets'; import { GatewayMetadata } from '@nestjs/websockets';
import { Socket } from 'socket.io';
import { defineStartupConfig, ModuleConfig } from '../config'; import { defineStartupConfig, ModuleConfig } from '../config';
@@ -6,7 +7,7 @@ declare module '../config' {
interface AppConfig { interface AppConfig {
websocket: ModuleConfig< websocket: ModuleConfig<
GatewayMetadata & { GatewayMetadata & {
requireAuthentication?: boolean; canActivate?: (socket: Socket) => Promise<boolean>;
} }
>; >;
} }
@@ -16,5 +17,4 @@ defineStartupConfig('websocket', {
// see: https://socket.io/docs/v4/server-options/#maxhttpbuffersize // see: https://socket.io/docs/v4/server-options/#maxhttpbuffersize
transports: ['websocket'], transports: ['websocket'],
maxHttpBufferSize: 1e8, // 100 MB maxHttpBufferSize: 1e8, // 100 MB
requireAuthentication: true,
}); });

View File

@@ -10,6 +10,7 @@ import { IoAdapter } from '@nestjs/platform-socket.io';
import { Server } from 'socket.io'; import { Server } from 'socket.io';
import { Config } from '../config'; import { Config } from '../config';
import { AuthenticationRequired } from '../error';
export const SocketIoAdapterImpl = Symbol('SocketIoAdapterImpl'); export const SocketIoAdapterImpl = Symbol('SocketIoAdapterImpl');
@@ -19,8 +20,31 @@ export class SocketIoAdapter extends IoAdapter {
} }
override createIOServer(port: number, options?: any): Server { override createIOServer(port: number, options?: any): Server {
const config = this.app.get(WEBSOCKET_OPTIONS); const config = this.app.get(WEBSOCKET_OPTIONS) as Config['websocket'];
return super.createIOServer(port, { ...config, ...options }); const server: Server = super.createIOServer(port, {
...config,
...options,
});
if (config.canActivate) {
server.use((socket, next) => {
// @ts-expect-error checked
config
.canActivate(socket)
.then(pass => {
if (pass) {
next();
} else {
throw new AuthenticationRequired();
}
})
.catch(e => {
next(e);
});
});
}
return server;
} }
} }