@@ -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;
|
||||||
}
|
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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 };
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
|
||||||
});
|
});
|
||||||
|
|||||||
@@ -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;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user