288 lines
8.8 KiB
TypeScript
288 lines
8.8 KiB
TypeScript
/**
|
|
* Notifications WebSocket Gateway
|
|
*
|
|
* Handles real-time notification delivery via WebSocket
|
|
*/
|
|
|
|
import {
|
|
WebSocketGateway,
|
|
WebSocketServer,
|
|
SubscribeMessage,
|
|
OnGatewayConnection,
|
|
OnGatewayDisconnect,
|
|
ConnectedSocket,
|
|
MessageBody,
|
|
} from '@nestjs/websockets';
|
|
import { Server, Socket } from 'socket.io';
|
|
import { Logger, UseGuards, UnauthorizedException } from '@nestjs/common';
|
|
import { JwtService } from '@nestjs/jwt';
|
|
import { JwtStrategy, JwtPayload } from '../auth/jwt.strategy';
|
|
import { NotificationService } from '../services/notification.service';
|
|
import { Notification } from '@domain/entities/notification.entity';
|
|
import { notificationTarget } from '@domain/services/notification-target';
|
|
|
|
/**
|
|
* WebSocket authentication guard
|
|
*/
|
|
@UseGuards()
|
|
@WebSocketGateway({
|
|
cors: {
|
|
origin: process.env.FRONTEND_URL || ['http://localhost:3000', 'http://localhost:3001'],
|
|
credentials: true,
|
|
},
|
|
namespace: '/notifications',
|
|
})
|
|
export class NotificationsGateway implements OnGatewayConnection, OnGatewayDisconnect {
|
|
@WebSocketServer()
|
|
server: Server;
|
|
|
|
private readonly logger = new Logger(NotificationsGateway.name);
|
|
private readonly connections = new Map<string, Socket>();
|
|
private userSockets: Map<string, Set<string>> = new Map(); // userId -> Set of socket IDs
|
|
|
|
constructor(
|
|
private readonly jwtService: JwtService,
|
|
private readonly notificationService: NotificationService,
|
|
private readonly jwtStrategy: JwtStrategy
|
|
) {}
|
|
|
|
/**
|
|
* Handle client connection
|
|
*/
|
|
async handleConnection(client: Socket) {
|
|
try {
|
|
// Extract JWT token from handshake
|
|
const token = this.extractToken(client);
|
|
if (!token) {
|
|
this.logger.warn(`Client ${client.id} connection rejected: No token provided`);
|
|
client.disconnect();
|
|
return;
|
|
}
|
|
|
|
// Verify JWT token
|
|
const user = await this.authenticate(client);
|
|
const userId = user.id;
|
|
this.connections.set(client.id, client);
|
|
|
|
// Store socket connection for user
|
|
if (!this.userSockets.has(userId)) {
|
|
this.userSockets.set(userId, new Set());
|
|
}
|
|
this.userSockets.get(userId)!.add(client.id);
|
|
|
|
// Store user ID in socket data for later use
|
|
client.data.userId = userId;
|
|
client.data.organizationId = user.organizationId;
|
|
|
|
// Join user-specific room
|
|
client.join(`user:${userId}`);
|
|
|
|
this.logger.log(`Client ${client.id} connected for user ${userId}`);
|
|
|
|
// Send unread count on connection
|
|
const unreadCount = await this.notificationService.getUnreadCount(userId);
|
|
client.emit('unread_count', { count: unreadCount });
|
|
|
|
// Send recent notifications on connection
|
|
const recentNotifications = await this.notificationService.getRecentNotifications(userId, 10);
|
|
client.emit('recent_notifications', {
|
|
notifications: recentNotifications.map(n => this.mapNotificationToDto(n)),
|
|
});
|
|
} catch (error: any) {
|
|
this.logger.error(
|
|
`Error during client connection: ${error?.message || 'Unknown error'}`,
|
|
error?.stack
|
|
);
|
|
client.disconnect();
|
|
}
|
|
}
|
|
|
|
/**
|
|
* Handle client disconnection
|
|
*/
|
|
handleDisconnect(client: Socket) {
|
|
this.connections.delete(client.id);
|
|
const userId = client.data.userId;
|
|
if (userId && this.userSockets.has(userId)) {
|
|
this.userSockets.get(userId)!.delete(client.id);
|
|
if (this.userSockets.get(userId)!.size === 0) {
|
|
this.userSockets.delete(userId);
|
|
}
|
|
}
|
|
this.logger.log(`Client ${client.id} disconnected`);
|
|
}
|
|
|
|
/**
|
|
* Handle mark notification as read
|
|
*/
|
|
@SubscribeMessage('mark_as_read')
|
|
async handleMarkAsRead(
|
|
@ConnectedSocket() client: Socket,
|
|
@MessageBody() data: { notificationId: string }
|
|
) {
|
|
try {
|
|
const userId = (await this.authenticate(client)).id;
|
|
await this.notificationService.markAsRead(data.notificationId, userId);
|
|
|
|
// Send updated unread count
|
|
const unreadCount = await this.notificationService.getUnreadCount(userId);
|
|
await this.emitToUser(userId, 'unread_count', { count: unreadCount });
|
|
|
|
return { success: true };
|
|
} catch (error: any) {
|
|
this.logger.error(`Error marking notification as read: ${error?.message}`);
|
|
return { success: false, error: error?.message };
|
|
}
|
|
}
|
|
|
|
/**
|
|
* Handle mark all notifications as read
|
|
*/
|
|
@SubscribeMessage('mark_all_as_read')
|
|
async handleMarkAllAsRead(@ConnectedSocket() client: Socket) {
|
|
try {
|
|
const userId = (await this.authenticate(client)).id;
|
|
await this.notificationService.markAllAsRead(userId);
|
|
|
|
// Send updated unread count (should be 0)
|
|
await this.emitToUser(userId, 'unread_count', { count: 0 });
|
|
|
|
return { success: true };
|
|
} catch (error: any) {
|
|
this.logger.error(`Error marking all notifications as read: ${error?.message}`);
|
|
return { success: false, error: error?.message };
|
|
}
|
|
}
|
|
|
|
/**
|
|
* Handle get unread count
|
|
*/
|
|
@SubscribeMessage('get_unread_count')
|
|
async handleGetUnreadCount(@ConnectedSocket() client: Socket) {
|
|
try {
|
|
const userId = (await this.authenticate(client)).id;
|
|
const unreadCount = await this.notificationService.getUnreadCount(userId);
|
|
return { count: unreadCount };
|
|
} catch (error: any) {
|
|
this.logger.error(`Error getting unread count: ${error?.message}`);
|
|
return { count: 0 };
|
|
}
|
|
}
|
|
|
|
/**
|
|
* Send notification to a specific user
|
|
*/
|
|
async sendNotificationToUser(userId: string, notification: Notification) {
|
|
const notificationDto = this.mapNotificationToDto(notification);
|
|
|
|
// Emit to all connected sockets for this user
|
|
await this.emitToUser(userId, 'new_notification', { notification: notificationDto });
|
|
|
|
// Update unread count
|
|
const unreadCount = await this.notificationService.getUnreadCount(userId);
|
|
await this.emitToUser(userId, 'unread_count', { count: unreadCount });
|
|
|
|
this.logger.log(`Notification sent to user ${userId}: ${notification.title}`);
|
|
}
|
|
|
|
/**
|
|
* Broadcast notification to organization
|
|
*/
|
|
async broadcastToOrganization(organizationId: string, notification: Notification) {
|
|
const notificationDto = this.mapNotificationToDto(notification);
|
|
for (const client of this.connections.values()) {
|
|
try {
|
|
const user = await this.authenticate(client);
|
|
if (user.organizationId === organizationId) {
|
|
client.emit('new_notification', { notification: notificationDto });
|
|
}
|
|
} catch {
|
|
client.disconnect();
|
|
}
|
|
}
|
|
|
|
this.logger.log(`Notification broadcasted to organization ${organizationId}`);
|
|
}
|
|
|
|
/**
|
|
* Helper: Emit event to all sockets of a user
|
|
*/
|
|
private async emitToUser(userId: string, event: string, data: unknown) {
|
|
for (const socketId of this.userSockets.get(userId) ?? []) {
|
|
const client = this.connections.get(socketId);
|
|
if (!client) continue;
|
|
try {
|
|
const user = await this.authenticate(client);
|
|
if (user.id === userId) client.emit(event, data);
|
|
} catch {
|
|
client.disconnect();
|
|
}
|
|
}
|
|
}
|
|
|
|
private async authenticate(client: Socket) {
|
|
try {
|
|
const token = this.extractToken(client);
|
|
if (!token) throw new UnauthorizedException();
|
|
const payload = await this.jwtService.verifyAsync<JwtPayload>(token);
|
|
if (
|
|
typeof payload.sub !== 'string' ||
|
|
!Number.isFinite(payload.exp) ||
|
|
payload.exp! * 1000 <= Date.now()
|
|
) {
|
|
throw new UnauthorizedException();
|
|
}
|
|
return await this.jwtStrategy.validate(payload);
|
|
} catch {
|
|
client.disconnect();
|
|
throw new UnauthorizedException('Invalid or expired session');
|
|
}
|
|
}
|
|
|
|
/**
|
|
* Helper: Extract JWT token from socket handshake
|
|
*/
|
|
private extractToken(client: Socket): string | null {
|
|
// Check Authorization header
|
|
const authHeader = client.handshake.headers.authorization;
|
|
if (authHeader && authHeader.startsWith('Bearer ')) {
|
|
return authHeader.substring(7);
|
|
}
|
|
|
|
// Check query parameter
|
|
const token = client.handshake.query.token;
|
|
if (typeof token === 'string') {
|
|
return token;
|
|
}
|
|
|
|
// Check auth object (socket.io-client way)
|
|
const auth = client.handshake.auth;
|
|
if (auth && typeof auth.token === 'string') {
|
|
return auth.token;
|
|
}
|
|
|
|
return null;
|
|
}
|
|
|
|
/**
|
|
* Helper: Map notification entity to DTO
|
|
*/
|
|
private mapNotificationToDto(notification: Notification) {
|
|
return {
|
|
id: notification.id,
|
|
type: notification.type,
|
|
priority: notification.priority,
|
|
title: notification.title,
|
|
message: notification.message,
|
|
metadata: notification.metadata,
|
|
read: notification.read,
|
|
readAt: notification.readAt?.toISOString(),
|
|
actionUrl:
|
|
notification.actionUrl ??
|
|
notificationTarget(notification.type, notification.metadata) ??
|
|
undefined,
|
|
createdAt: notification.createdAt.toISOString(),
|
|
};
|
|
}
|
|
}
|