WebSocketServer.java 6.78 KB
package com.infoloop.tianting.server;


import com.infoloop.tianting.constant.CommonConstants;
import com.infoloop.tianting.server.message.SocketMessage;
import com.infoloop.tianting.server.session.UserSessionData;
import com.infoloop.tianting.server.session.UserSessionKey;
import com.infoloop.tianting.server.session.UserTypeEnum;
import lombok.extern.slf4j.Slf4j;
import org.springframework.web.socket.CloseStatus;
import org.springframework.web.socket.TextMessage;
import org.springframework.web.socket.WebSocketSession;
import org.springframework.web.socket.handler.TextWebSocketHandler;

import java.io.IOException;
import java.util.List;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.Executors;
import java.util.concurrent.ScheduledExecutorService;
import java.util.concurrent.TimeUnit;
import java.util.stream.Collectors;

@Slf4j
public class WebSocketServer extends TextWebSocketHandler {

    private static final ConcurrentHashMap<UserSessionKey, UserSessionData> userSessions = new ConcurrentHashMap<>();
    private static final long INACTIVITY_TIMEOUT = 60 * 60 * 1000;
    private static final ScheduledExecutorService scheduler = Executors.newScheduledThreadPool(1);

    static {
        scheduler.scheduleAtFixedRate(WebSocketServer::cleanInactiveSessions, 1, 1, TimeUnit.HOURS);
    }

    private static void cleanInactiveSessions() {
        final var currentTime = System.currentTimeMillis();
        final var iterator = userSessions.entrySet().iterator();
        log.info("cleanInactiveSessions; sessions:{}", userSessions.keySet());
        while (iterator.hasNext()) {
            final var entry = iterator.next();
            final var key = entry.getKey();
            final var userData = entry.getValue();
            final var lastActiveTime = userData.getLastActiveTime();
            if (currentTime - lastActiveTime > INACTIVITY_TIMEOUT) {
                try {
                    userData.closeSession();
                    log.info("closeSession,userId : {}, userType : {}", key.getUserId(), key.getUserType());
                } catch (IOException e) {
                    log.error("closeSession error,userId : {}, userType : {}", key.getUserId(), key.getUserType(), e);
                }
                iterator.remove();
            }
        }
    }

    public static <T> void sendMessageToUser(UserSessionKey key, SocketMessage<T> socketMessage) throws IOException {
        sendMessageToUsers(List.of(key), socketMessage);
    }

    public static <T> void sendMessageToUsers(List<UserSessionKey> keys, SocketMessage<T> socketMessage) throws IOException {
        for (final var key : keys) {
            final var userData = userSessions.get(key);
            if (userData != null && userData.getSession().isOpen()) {
                userData.getSession().sendMessage(new TextMessage(socketMessage.toJsonString()));
                userData.updateLastActiveTime();
                log.info("send message to userId : {}, userType : {}", key.getUserId(), key.getUserType());
            } else {
                log.info("userId : {}, userType : {} WebSocket connect closed; ", key.getUserId(), key.getUserType());
            }
        }
    }

    public static <T> void sendMessageByUserType(UserTypeEnum userType, SocketMessage<T> socketMessage) throws IOException {
       sendMessageByUserTypes(List.of(userType), socketMessage);
    }

    public static <T> void sendMessageByUserTypes(List<UserTypeEnum> userTypes, SocketMessage<T> socketMessage) throws IOException {
        final var keys = userSessions.keySet().stream()
                .filter(userSessionData -> userTypes.contains(userSessionData.getUserType()))
                .collect(Collectors.toList());
        sendMessageToUsers(keys, socketMessage);
    }

    public static <T> void sendMessageByUserTypes(List<UserTypeEnum> userTypes, List<String> authUserIds, SocketMessage<T> socketMessage) throws IOException {
        final var keys = userSessions.keySet().stream()
                .filter(userSessionData -> userTypes.contains(userSessionData.getUserType()))
                .filter(userSessionData -> authUserIds.isEmpty() || authUserIds.contains(userSessionData.getUserId()))
                .collect(Collectors.toList());
        sendMessageToUsers(keys, socketMessage);
    }

    public static void closeConnection(UserSessionKey key) {
        final var userData = userSessions.get(key);
        if (userData != null) {
            try {
                userData.closeSession();
                userSessions.remove(key);
                log.info("WebSocket close; userId = {}, userType = {}", key.getUserId(), key.getUserType());
            } catch (IOException e) {
                log.error("WebSocket close failed: userId = {}, userType = {}", key.getUserId(), key.getUserType(), e);
            }
        } else {
            log.warn("Not Found userId : {}, userType : {} WebSocket Connect ", key.getUserId(), key.getUserType());
        }
    }

    private UserSessionKey getUserSessionKey(WebSocketSession session) {
        final var userId = (String) session.getAttributes().get(CommonConstants.USER_ID);
        final var userTypeStr = (String) session.getAttributes().get(CommonConstants.USER_TYPE);
        final var userType = UserTypeEnum.fromString(userTypeStr);
        if (userId != null && userType != null) {
            return UserSessionKey.builder().userId(userId).userType(userType).build();
        }
        return null;
    }

    @Override
    public void afterConnectionEstablished(WebSocketSession session) {
        final var key = getUserSessionKey(session);
        if (key != null) {
            final var existingSession = userSessions.remove(key);
            if (existingSession != null) {
                try {
                    existingSession.closeSession();
                    log.info("WebSocket old connect closed,userId : {}, userType : {}", key.getUserId(), key.getUserType());
                } catch (IOException e) {
                    log.error("old WebSocket close failed,userId : {}, userType : {}", key.getUserId(), key.getUserType(), e);
                }
            }
            userSessions.put(key, new UserSessionData(session));
            log.info("WebSocket connect success, userId : {}, userType : {}", key.getUserId(), key.getUserType());
        } else {
            log.warn("WebSocket connect failed,unable to get valid user information");
        }
    }

    @Override
    public void afterConnectionClosed(WebSocketSession session, CloseStatus status) {
        final var key = getUserSessionKey(session);
        if (key != null) {
            userSessions.remove(key);
            log.info("WebSocket closed: userId : {}, userType : {}", key.getUserId(), key.getUserType());
        } else {
            log.warn("WebSocket closed failed,unable to get valid user information");
        }
    }
}