WebSocketServer.java 7.2 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.apache.commons.lang3.StringUtils;
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 javax.annotation.Nullable;
import java.io.IOException;
import java.util.List;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.CopyOnWriteArrayList;
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, List<UserSessionData>> userSessions = new ConcurrentHashMap<>();
    private static final long INACTIVITY_TIMEOUT = 8 * 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();
        userSessions.forEach((key, sessionList) ->
                userSessions.computeIfPresent(key, (k, sessions) -> {
                    sessions.removeIf(userData -> {
                        if (currentTime - userData.getLastActiveTime() > 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);
                            }
                            return true;
                        }
                        return false;
                    });
                    return sessions.isEmpty() ? null : sessions; // 为空时移除
                })
        );
    }


    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 sessionList = userSessions.get(key);
            if (sessionList != null) {
                sessionList.forEach(userData -> {
                    try {
                        if (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("WebSocket closed; userId: {}, userType: {}", key.getUserId(), key.getUserType());
                        }
                    } catch (IOException e) {
                        log.error("Failed to send message; userId: {}, userType: {}", key.getUserId(), key.getUserType(), e);
                    }
                });
            }
        }
    }

    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(userKey -> userTypes.contains(userKey.getUserType()))
                .collect(Collectors.toList());
        sendMessageToUsers(keys, socketMessage);
    }

    public static <T> void sendMessageByUserTypes(List<UserTypeEnum> userTypes, List<UserAuthData> userAuthData, SocketMessage<T> socketMessage) throws IOException {
        final var keys = userSessions.keySet().stream()
                .filter(userKey -> userTypes.contains(userKey.getUserType()))
                .filter(userKey -> userAuthData.isEmpty()
                        || userAuthData.stream().anyMatch(data -> String.valueOf(data.getOperatorId()).equals(userKey.getUserId()) && data.getUserType().equals(userKey.getUserType())))
                .collect(Collectors.toList());
        sendMessageToUsers(keys, socketMessage);
    }

    public static void closeConnection(UserSessionKey key) {
        userSessions.computeIfPresent(key, (k, sessionList) -> {
            sessionList.forEach(userData -> {
                try {
                    userData.closeSession();
                } catch (IOException e) {
                    log.error("WebSocket close failed: userId: {}, userType: {}", key.getUserId(), key.getUserType(), e);
                }
            });
            return null; // 直接返回 null,移除 key
        });
        log.info("WebSocket closed; userId: {}, userType: {}", key.getUserId(), key.getUserType());
    }

    @Nullable
    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 (StringUtils.isNotEmpty(userId) && 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) {
            userSessions.computeIfAbsent(key, k -> new CopyOnWriteArrayList<>()).add(new UserSessionData(session));
            log.info("WebSocket connection established; userId: {}, userType: {}", key.getUserId(), key.getUserType());
        } else {
            log.warn("WebSocket connection failed; unable to retrieve valid user information");
        }
    }


    @Override
    public void afterConnectionClosed(WebSocketSession session, CloseStatus status) {
        final var key = getUserSessionKey(session);
        if (key != null) {
            userSessions.computeIfPresent(key, (k, sessionList) -> {
                sessionList.removeIf(userData -> userData.getSession().getId().equals(session.getId()));
                return sessionList.isEmpty() ? null : sessionList;
            });
            log.info("WebSocket closed; userId: {}, userType: {}", key.getUserId(), key.getUserType());
        } else {
            log.warn("WebSocket closed; unable to retrieve valid user information");
        }
    }
}