WebSocket实时通信架构:从连接到百万并发的实战指南

深入讲解WebSocket实时通信架构的设计与实现,涵盖连接管理、心跳机制、消息队列集成、水平扩展、负载均衡等核心主题,提供Node.js和Go的完整实现代码。

WebSocket vs 其他实时方案

实时通信方案对比:
┌─────────────────────────────────────────┐
│ WebSocket                                │
│ ✓ 全双工通信,低延迟                     │
│ ✓ 持久连接,减少握手开销                 │
│ ✓ 适合高频双向交互(聊天、游戏)         │
│ ✗ 需要维护连接状态                       │
│ ✗ 扩展性需要考虑连接亲和性               │
│                                         │
│ Server-Sent Events (SSE)                │
│ ✓ 基于HTTP,简单实现                     │
│ ✓ 自动重连,易于调试                     │
│ ✓ 适合单向推送(通知、股票行情)         │
│ ✗ 只能服务器到客户端                     │
│ ✗ 连接数受限(HTTP/1.1每域6个)          │
│                                         │
│ 长轮询(Long Polling)                   │
│ ✓ 兼容性好,无需特殊支持                 │
│ ✓ 实现简单                               │
│ ✗ 延迟较高                               │
│ ✗ 服务器资源消耗大                       │
│                                         │
│ 选择建议:                               │
│ - 双向高频交互 → WebSocket               │
│ - 单向实时推送 → SSE                     │
│ - 兼容性要求高 → 长轮询                  │
└─────────────────────────────────────────┘

WebSocket服务器实现

Node.js + Socket.IO

// server.js
const express = require('express');
const http = require('http');
const { Server } = require('socket.io');
const redis = require('redis');
const jwt = require('jsonwebtoken');

const app = express();
const server = http.createServer(app);

// Socket.IO配置
const io = new Server(server, {
  cors: {
    origin: process.env.CLIENT_URL,
    methods: ['GET', 'POST']
  },
  pingTimeout: 60000,
  pingInterval: 25000
});

// Redis适配器(多进程支持)
const pubClient = redis.createClient({ url: process.env.REDIS_URL });
const subClient = pubClient.duplicate();

io.adapter(createAdapter(pubClient, subClient));

// 连接认证中间件
io.use(async (socket, next) => {
  try {
    const token = socket.handshake.auth.token;
    if (!token) {
      return next(new Error('Authentication error'));
    }
    
    const decoded = jwt.verify(token, process.env.JWT_SECRET);
    socket.userId = decoded.userId;
    socket.username = decoded.username;
    
    next();
  } catch (err) {
    next(new Error('Authentication error'));
  }
});

// 连接管理
const userSockets = new Map(); // userId -> Set<socketId>

io.on('connection', (socket) => {
  console.log(`User ${socket.userId} connected: ${socket.id}`);
  
  // 记录用户的socket连接
  if (!userSockets.has(socket.userId)) {
    userSockets.set(socket.userId, new Set());
  }
  userSockets.get(socket.userId).add(socket.id);
  
  // 加入用户的个人房间
  socket.join(`user:${socket.userId}`);
  
  // 心跳检测
  socket.on('pong', () => {
    socket.isAlive = true;
  });
  
  // 加入聊天室
  socket.on('join:room', async (roomId) => {
    try {
      // 验证用户是否有权限加入
      const hasAccess = await checkRoomAccess(socket.userId, roomId);
      if (!hasAccess) {
        socket.emit('error', { message: 'Access denied' });
        return;
      }
      
      socket.join(`room:${roomId}`);
      socket.emit('room:joined', { roomId });
      
      // 通知房间内其他用户
      socket.to(`room:${roomId}`).emit('user:joined', {
        userId: socket.userId,
        username: socket.username
      });
    } catch (err) {
      socket.emit('error', { message: 'Failed to join room' });
    }
  });
  
  // 发送消息
  socket.on('message:send', async (data) => {
    try {
      const { roomId, content, type = 'text' } = data;
      
      // 创建消息记录
      const message = await createMessage({
        roomId,
        userId: socket.userId,
        content,
        type,
        timestamp: Date.now()
      });
      
      // 广播给房间内所有用户
      io.to(`room:${roomId}`).emit('message:new', {
        id: message.id,
        roomId,
        userId: socket.userId,
        username: socket.username,
        content,
        type,
        timestamp: message.timestamp
      });
      
      // 发送确认给发送者
      socket.emit('message:sent', { messageId: message.id });
    } catch (err) {
      socket.emit('error', { message: 'Failed to send message' });
    }
  });
  
  // 输入状态
  socket.on('typing:start', (roomId) => {
    socket.to(`room:${roomId}`).emit('user:typing', {
      userId: socket.userId,
      username: socket.username
    });
  });
  
  socket.on('typing:stop', (roomId) => {
    socket.to(`room:${roomId}`).emit('user:stop-typing', {
      userId: socket.userId
    });
  });
  
  // 断开连接
  socket.on('disconnect', (reason) => {
    console.log(`User ${socket.userId} disconnected: ${reason}`);
    
    // 清理用户socket记录
    const sockets = userSockets.get(socket.userId);
    if (sockets) {
      sockets.delete(socket.id);
      if (sockets.size === 0) {
        userSockets.delete(socket.userId);
      }
    }
    
    // 通知相关房间
    socket.rooms.forEach(room => {
      if (room.startsWith('room:')) {
        socket.to(room).emit('user:left', {
          userId: socket.userId,
          username: socket.username
        });
      }
    });
  });
});

// 心跳检测定时器
const heartbeatInterval = setInterval(() => {
  io.sockets.sockets.forEach((socket) => {
    if (socket.isAlive === false) {
      return socket.disconnect(true);
    }
    
    socket.isAlive = false;
    socket.emit('ping');
  });
}, 30000);

// 清理
server.on('close', () => {
  clearInterval(heartbeatInterval);
  pubClient.quit();
  subClient.quit();
});

server.listen(3000, () => {
  console.log('WebSocket server running on port 3000');
});

Go实现(高并发场景)

// main.go
package main

import (
    "context"
    "encoding/json"
    "log"
    "net/http"
    "sync"
    "time"
    
    "github.com/gorilla/websocket"
    "github.com/redis/go-redis/v9"
)

var upgrader = websocket.Upgrader{
    ReadBufferSize:  1024,
    WriteBufferSize: 1024,
    CheckOrigin: func(r *http.Request) bool {
        return true // 生产环境需要验证origin
    },
}

type Message struct {
    Type    string      `json:"type"`
    RoomID  string      `json:"roomId,omitempty"`
    Content interface{} `json:"content"`
    UserID  string      `json:"userId"`
    Time    int64       `json:"time"`
}

type Client struct {
    ID       string
    UserID   string
    Conn     *websocket.Conn
    Send     chan []byte
    Rooms    map[string]bool
    mu       sync.RWMutex
}

type Hub struct {
    Clients    map[string]*Client
    Rooms      map[string]map[string]*Client
    Register   chan *Client
    Unregister chan *Client
    Broadcast  chan *BroadcastMessage
    mu         sync.RWMutex
}

type BroadcastMessage struct {
    RoomID  string
    Message []byte
    Exclude string
}

func NewHub() *Hub {
    return &Hub{
        Clients:    make(map[string]*Client),
        Rooms:      make(map[string]map[string]*Client),
        Register:   make(chan *Client),
        Unregister: make(chan *Client),
        Broadcast:  make(chan *BroadcastMessage, 256),
    }
}

func (h *Hub) Run(ctx context.Context) {
    for {
        select {
        case client := <-h.Register:
            h.mu.Lock()
            h.Clients[client.ID] = client
            h.mu.Unlock()
            log.Printf("Client connected: %s (user: %s)", client.ID, client.UserID)
            
        case client := <-h.Unregister:
            h.mu.Lock()
            if _, ok := h.Clients[client.ID]; ok {
                delete(h.Clients, client.ID)
                close(client.Send)
                
                // 从所有房间移除
                client.mu.RLock()
                for roomID := range client.Rooms {
                    if room, ok := h.Rooms[roomID]; ok {
                        delete(room, client.ID)
                        if len(room) == 0 {
                            delete(h.Rooms, roomID)
                        }
                    }
                }
                client.mu.RUnlock()
            }
            h.mu.Unlock()
            log.Printf("Client disconnected: %s", client.ID)
            
        case msg := <-h.Broadcast:
            h.mu.RLock()
            if room, ok := h.Rooms[msg.RoomID]; ok {
                for clientID, client := range room {
                    if clientID == msg.Exclude {
                        continue
                    }
                    select {
                    case client.Send <- msg.Message:
                    default:
                        // 客户端发送队列满,断开连接
                        close(client.Send)
                        delete(room, clientID)
                    }
                }
            }
            h.mu.RUnlock()
            
        case <-ctx.Done():
            return
        }
    }
}

func (h *Hub) JoinRoom(clientID, roomID string) {
    h.mu.Lock()
    defer h.mu.Unlock()
    
    client, ok := h.Clients[clientID]
    if !ok {
        return
    }
    
    client.mu.Lock()
    client.Rooms[roomID] = true
    client.mu.Unlock()
    
    if _, ok := h.Rooms[roomID]; !ok {
        h.Rooms[roomID] = make(map[string]*Client)
    }
    h.Rooms[roomID][clientID] = client
}

func (h *Hub) LeaveRoom(clientID, roomID string) {
    h.mu.Lock()
    defer h.mu.Unlock()
    
    client, ok := h.Clients[clientID]
    if !ok {
        return
    }
    
    client.mu.Lock()
    delete(client.Rooms, roomID)
    client.mu.Unlock()
    
    if room, ok := h.Rooms[roomID]; ok {
        delete(room, clientID)
        if len(room) == 0 {
            delete(h.Rooms, roomID)
        }
    }
}

var hub = NewHub()
var redisClient *redis.Client

func init() {
    redisClient = redis.NewClient(&redis.Options{
        Addr: "localhost:6379",
    })
}

func handleWebSocket(w http.ResponseWriter, r *http.Request) {
    // 验证JWT token
    token := r.URL.Query().Get("token")
    userID, err := validateToken(token)
    if err != nil {
        http.Error(w, "Unauthorized", http.StatusUnauthorized)
        return
    }
    
    conn, err := upgrader.Upgrade(w, r, nil)
    if err != nil {
        log.Printf("Upgrade error: %v", err)
        return
    }
    
    client := &Client{
        ID:     generateID(),
        UserID: userID,
        Conn:   conn,
        Send:   make(chan []byte, 256),
        Rooms:  make(map[string]bool),
    }
    
    hub.Register <- client
    
    // 读取消息
    go client.readPump()
    
    // 写入消息
    go client.writePump()
}

func (c *Client) readPump() {
    defer func() {
        hub.Unregister <- c
        c.Conn.Close()
    }()
    
    c.Conn.SetReadLimit(4096)
    c.Conn.SetReadDeadline(time.Now().Add(60 * time.Second))
    c.Conn.SetPongHandler(func(string) error {
        c.Conn.SetReadDeadline(time.Now().Add(60 * time.Second))
        return nil
    })
    
    for {
        _, message, err := c.Conn.ReadMessage()
        if err != nil {
            if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) {
                log.Printf("Read error: %v", err)
            }
            break
        }
        
        var msg Message
        if err := json.Unmarshal(message, &msg); err != nil {
            log.Printf("Unmarshal error: %v", err)
            continue
        }
        
        msg.UserID = c.UserID
        msg.Time = time.Now().UnixMilli()
        
        switch msg.Type {
        case "join":
            hub.JoinRoom(c.ID, msg.RoomID)
            c.Send <- createResponse("joined", msg.RoomID)
            
        case "leave":
            hub.LeaveRoom(c.ID, msg.RoomID)
            c.Send <- createResponse("left", msg.RoomID)
            
        case "message":
            msgBytes, _ := json.Marshal(msg)
            hub.Broadcast <- &BroadcastMessage{
                RoomID:  msg.RoomID,
                Message: msgBytes,
                Exclude: c.ID,
            }
            
            // 保存到数据库
            go saveMessage(msg)
        }
    }
}

func (c *Client) writePump() {
    ticker := time.NewTicker(30 * time.Second)
    defer func() {
        ticker.Stop()
        c.Conn.Close()
    }()
    
    for {
        select {
        case message, ok := <-c.Send:
            c.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
            if !ok {
                c.Conn.WriteMessage(websocket.CloseMessage, []byte{})
                return
            }
            
            w, err := c.Conn.NextWriter(websocket.TextMessage)
            if err != nil {
                return
            }
            w.Write(message)
            
            if err := w.Close(); err != nil {
                return
            }
            
        case <-ticker.C:
            c.Conn.SetWriteDeadline(time.Now().Add(10 * time.Second))
            if err := c.Conn.WriteMessage(websocket.PingMessage, nil); err != nil {
                return
            }
        }
    }
}

func main() {
    ctx := context.Background()
    go hub.Run(ctx)
    
    http.HandleFunc("/ws", handleWebSocket)
    
    log.Println("Server starting on :8080")
    log.Fatal(http.ListenAndServe(":8080", nil))
}

水平扩展方案

Redis Pub/Sub实现多服务器通信

// scaling/redis-adapter.js
const { createAdapter } = require('@socket.io/redis-adapter');
const { createClient } = require('redis');

class RedisAdapter {
  constructor() {
    this.pubClient = createClient({ url: process.env.REDIS_URL });
    this.subClient = this.pubClient.duplicate();
    
    this.pubClient.on('error', err => console.error('Redis Pub Error:', err));
    this.subClient.on('error', err => console.error('Redis Sub Error:', err));
  }
  
  async init(io) {
    await Promise.all([
      this.pubClient.connect(),
      this.subClient.connect()
    ]);
    
    io.adapter(createAdapter(this.pubClient, this.subClient));
    
    console.log('Redis adapter initialized');
  }
  
  // 跨服务器发送消息
  async broadcastToRoom(room, event, data) {
    await this.pubClient.publish(
      `room:${room}`,
      JSON.stringify({ event, data })
    );
  }
  
  // 获取在线用户数(跨服务器)
  async getOnlineUsers() {
    const keys = await this.pubClient.keys('socket:user:*');
    return keys.length;
  }
  
  // 获取房间成员(跨服务器)
  async getRoomMembers(room) {
    const members = await this.pubClient.sMembers(`room:${room}:members`);
    return members;
  }
}

module.exports = RedisAdapter;

// docker-compose.yml
version: '3.8'
services:
  redis:
    image: redis:7-alpine
    ports:
      - "6379:6379"
    volumes:
      - redis-data:/data
    
  nginx:
    image: nginx:alpine
    ports:
      - "80:80"
    volumes:
      - ./nginx.conf:/etc/nginx/nginx.conf
    depends_on:
      - ws-server-1
      - ws-server-2
      - ws-server-3
  
  ws-server-1:
    build: .
    environment:
      - REDIS_URL=redis://redis:6379
      - JWT_SECRET=${JWT_SECRET}
    depends_on:
      - redis
  
  ws-server-2:
    build: .
    environment:
      - REDIS_URL=redis://redis:6379
      - JWT_SECRET=${JWT_SECRET}
    depends_on:
      - redis
  
  ws-server-3:
    build: .
    environment:
      - REDIS_URL=redis://redis:6379
      - JWT_SECRET=${JWT_SECRET}
    depends_on:
      - redis

volumes:
  redis-data:

Nginx负载均衡配置

# nginx.conf
upstream websocket_backend {
    # 使用IP hash保证连接粘性
    ip_hash;
    
    server ws-server-1:3000 max_fails=3 fail_timeout=30s;
    server ws-server-2:3000 max_fails=3 fail_timeout=30s;
    server ws-server-3:3000 max_fails=3 fail_timeout=30s;
    
    keepalive 32;
}

server {
    listen 80;
    server_name ws.example.com;
    
    # WebSocket连接升级
    location /ws {
        proxy_pass http://websocket_backend;
        proxy_http_version 1.1;
        proxy_set_header Upgrade $http_upgrade;
        proxy_set_header Connection "upgrade";
        proxy_set_header Host $host;
        proxy_set_header X-Real-IP $remote_addr;
        proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
        proxy_set_header X-Forwarded-Proto $scheme;
        
        # 超时配置
        proxy_connect_timeout 7d;
        proxy_send_timeout 7d;
        proxy_read_timeout 7d;
        
        # 缓冲配置
        proxy_buffering off;
    }
    
    # 健康检查
    location /health {
        access_log off;
        return 200 "healthy\n";
    }
}

消息持久化与离线消息

// services/messageService.js
const { Pool } = require('pg');

const pool = new Pool({
  connectionString: process.env.DATABASE_URL
});

class MessageService {
  // 保存消息
  async saveMessage(message) {
    const query = `
      INSERT INTO messages (room_id, user_id, content, type, timestamp)
      VALUES ($1, $2, $3, $4, $5)
      RETURNING id
    `;
    
    const result = await pool.query(query, [
      message.roomId,
      message.userId,
      message.content,
      message.type,
      message.timestamp
    ]);
    
    return result.rows[0].id;
  }
  
  // 获取历史消息
  async getMessages(roomId, limit = 50, before = null) {
    let query = `
      SELECT m.*, u.username, u.avatar
      FROM messages m
      JOIN users u ON m.user_id = u.id
      WHERE m.room_id = $1
    `;
    
    const params = [roomId];
    
    if (before) {
      query += ' AND m.timestamp < $2';
      params.push(before);
    }
    
    query += ' ORDER BY m.timestamp DESC LIMIT $' + (params.length + 1);
    params.push(limit);
    
    const result = await pool.query(query, params);
    return result.rows.reverse();
  }
  
  // 获取离线消息
  async getOfflineMessages(userId, since) {
    const query = `
      SELECT m.*, r.name as room_name
      FROM messages m
      JOIN rooms r ON m.room_id = r.id
      JOIN room_members rm ON r.id = rm.room_id
      WHERE rm.user_id = $1
        AND m.timestamp > $2
        AND m.user_id != $1
      ORDER BY m.timestamp ASC
    `;
    
    const result = await pool.query(query, [userId, since]);
    return result.rows;
  }
  
  // 标记消息已读
  async markAsRead(userId, roomId, messageId) {
    const query = `
      INSERT INTO message_reads (user_id, room_id, message_id, read_at)
      VALUES ($1, $2, $3, NOW())
      ON CONFLICT (user_id, message_id) DO NOTHING
    `;
    
    await pool.query(query, [userId, roomId, messageId]);
  }
}

module.exports = new MessageService();

// server.js - 集成离线消息
io.on('connection', async (socket) => {
  // 发送离线消息
  const offlineMessages = await messageService.getOfflineMessages(
    socket.userId,
    socket.lastSeen || Date.now() - 7 * 24 * 60 * 60 * 1000 // 7天内
  );
  
  if (offlineMessages.length > 0) {
    socket.emit('messages:offline', { messages: offlineMessages });
  }
  
  // 更新最后在线时间
  socket.on('disconnect', async () => {
    await updateUserLastSeen(socket.userId, Date.now());
  });
});

性能优化与监控

// monitoring/metrics.js
const prometheus = require('prom-client');

// 定义指标
const connectedClients = new prometheus.Gauge({
  name: 'websocket_connected_clients',
  help: 'Number of connected WebSocket clients'
});

const messagesTotal = new prometheus.Counter({
  name: 'websocket_messages_total',
  help: 'Total number of WebSocket messages',
  labelNames: ['type']
});

const messageDuration = new prometheus.Histogram({
  name: 'websocket_message_duration_seconds',
  help: 'Message processing duration',
  labelNames: ['type'],
  buckets: [0.001, 0.005, 0.01, 0.05, 0.1, 0.5, 1]
});

const roomMembers = new prometheus.Gauge({
  name: 'websocket_room_members',
  help: 'Number of members in each room',
  labelNames: ['room']
});

// 更新指标
io.on('connection', (socket) => {
  connectedClients.inc();
  
  socket.on('message:send', (data) => {
    const end = messageDuration.startTimer({ type: 'send' });
    messagesTotal.inc({ type: 'send' });
    
    // 处理消息...
    
    end();
  });
  
  socket.on('join:room', (roomId) => {
    roomMembers.inc({ room: roomId });
  });
  
  socket.on('leave:room', (roomId) => {
    roomMembers.dec({ room: roomId });
  });
  
  socket.on('disconnect', () => {
    connectedClients.dec();
  });
});

// 暴露metrics端点
app.get('/metrics', async (req, res) => {
  res.set('Content-Type', prometheus.register.contentType);
  res.end(await prometheus.register.metrics());
});

// monitoring/health.js
app.get('/health', (req, res) => {
  const health = {
    status: 'ok',
    timestamp: new Date().toISOString(),
    uptime: process.uptime(),
    connections: io.sockets.sockets.size,
    memory: process.memoryUsage(),
    cpu: process.cpuUsage()
  };
  
  res.json(health);
});

// 连接限制
const MAX_CONNECTIONS_PER_USER = 5;

io.use((socket, next) => {
  const userConnections = Array.from(io.sockets.sockets.values())
    .filter(s => s.userId === socket.userId)
    .length;
  
  if (userConnections >= MAX_CONNECTIONS_PER_USER) {
    return next(new Error('Too many connections'));
  }
  
  next();
});

延伸阅读

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「backend」更多文章

  1. 幂等性设计模式:构建可靠的分布式系统
  2. BFF架构模式:为不同前端定制专属后端服务
  3. 蓝绿部署与金丝雀发布:零停机部署策略实战