Skip to content

WebSocket 与 Gin 集成与生产级应用

上一篇我们用纯 net/http + gorilla/websocket 实现了聊天室。真实业务中,WebSocket 往往不是孤立存在的,它要和 REST API、认证、限流、日志等共享一套 Web 框架基础设施。本篇把 WebSocket 接入 Gin 框架,实现认证中间件、限流、连接管理器、私聊/群组/通知分发,并引入断线重连、消息缓冲、Redis Pub/Sub 跨节点广播、消息持久化等生产级特性,最后给出一个完整的实时通知系统示例。

一、在 Gin 中使用 WebSocket

Gin 基于 net/http,而 gorilla/websocket 的 Upgrader.Upgrade 接收的正是 http.ResponseWriter*http.Request。Gin 的 *gin.Context 直接暴露了这两个底层对象(c.Writerc.Request),所以集成非常自然。

1. Gin Handler 中升级 WebSocket

go
package main

import (
	"fmt"
	"net/http"
	"github.com/gin-gonic/gin"
	"github.com/gorilla/websocket"
)

var upgrader = websocket.Upgrader{
	CheckOrigin: func(r *http.Request) bool { return true },
}

func main() {
	r := gin.Default()

	r.GET("/ws", func(c *gin.Context) {
		// 把 Gin 的 ResponseWriter 和 Request 传给 Upgrader
		conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
		if err != nil {
			fmt.Println("升级失败:", err)
			return
		}
		defer conn.Close()

		for {
			msgType, msg, err := conn.ReadMessage()
			if err != nil {
				break
			}
			conn.WriteMessage(msgType, msg)
		}
	})

	r.Run(":8080")
}

这里的关键是 upgrader.Upgrade(c.Writer, c.Request, nil)。注意一个重要陷阱:调用 Upgrade 之后,不要再对 c.Writer 写任何 HTTP 响应,因为连接已经被「劫持」为 WebSocket 了。Gin 的中间件链在 Upgrade 之后对当前请求也失效了。

2. 路由注册

WebSocket 路由和普通路由一样注册,通常用 GET 方法。可以把 WebSocket 路由放在单独的路由组里,统一挂中间件:

go
package main

import (
	"net/http"
	"github.com/gin-gonic/gin"
	"github.com/gorilla/websocket"
)

var upgrader = websocket.Upgrader{
	CheckOrigin: func(r *http.Request) bool { return true },
}

func wsHandler(c *gin.Context) {
	conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
	if err != nil {
		return
	}
	defer conn.Close()
	for {
		if _, _, err := conn.ReadMessage(); err != nil {
			break
		}
	}
}

func main() {
	r := gin.Default()
	// WebSocket 路由组,挂载认证等中间件
	wsGroup := r.Group("/ws")
	wsGroup.GET("/chat", wsHandler)
	wsGroup.GET("/notify", wsHandler)

	r.Run(":8080")
}

二、中间件集成

中间件是 Gin 的核心能力,WebSocket 路由同样可以享受。注意中间件只在握手阶段(HTTP 请求)生效,握手完成后连接就是 WebSocket 了,中间件不再介入。

1. 认证中间件:从 query/header 获取 token

浏览器原生 WebSocket API 不支持自定义请求头,所以认证 token 通常放在 query 参数里(如 ws://host/ws?token=xxx),或放在 Cookie/子协议里。最常见的是 query 参数。

go
package main

import (
	"net/http"
	"strings"
	"github.com/gin-gonic/gin"
	"github.com/gorilla/websocket"
)

var upgrader = websocket.Upgrader{
	CheckOrigin: func(r *http.Request) bool { return true },
}

// 模拟 token 校验:返回用户 ID
func parseToken(token string) (string, bool) {
	// 实际项目应调用 JWT 解析或查 Redis
	if strings.HasPrefix(token, "user-") {
		return strings.TrimPrefix(token, "user-"), true
	}
	return "", false
}

// 认证中间件
func authMiddleware() gin.HandlerFunc {
	return func(c *gin.Context) {
		token := c.Query("token")
		if token == "" {
			// 也支持 header(非浏览器客户端可用)
			token = c.GetHeader("Sec-WebSocket-Protocol")
		}
		uid, ok := parseToken(token)
		if !ok {
			c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "无效 token"})
			return
		}
		// 把用户 ID 存到 context,后续 handler 可取
		c.Set("uid", uid)
		c.Next()
	}
}

func wsHandler(c *gin.Context) {
	uid, _ := c.Get("uid")
	conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
	if err != nil {
		return
	}
	defer conn.Close()

	// 把用户 ID 发回去,证明认证成功
	conn.WriteMessage(websocket.TextMessage, []byte("欢迎用户: "+uid.(string)))
	for {
		if _, _, err := conn.ReadMessage(); err != nil {
			break
		}
	}
}

func main() {
	r := gin.Default()
	r.GET("/ws", authMiddleware(), wsHandler)
	r.Run(":8080")
}

安全提示:token 放在 URL 里会被记进访问日志和浏览器历史。生产环境建议用短期 token(如 5 分钟有效的握手专用 token),握手后改用 Cookie 或在协议内重新颁发会话凭证。

2. 限流中间件

WebSocket 连接建立后就长期占用,容易成为攻击目标(恶意建立大量连接耗尽资源)。可以在握手前做限流:

go
package main

import (
	"net/http"
	"sync"
	"time"
	"github.com/gin-gonic/gin"
	"github.com/gorilla/websocket"
)

var upgrader = websocket.Upgrader{
	CheckOrigin: func(r *http.Request) bool { return true },
}

// 基于 IP 的简单限流器(令牌桶思路)
type ipLimiter struct {
	mu    sync.Mutex
	count map[string]int
	limit int
	window time.Duration
}

func newIPLimiter(limit int, window time.Duration) *ipLimiter {
	l := &ipLimiter{count: make(map[string]int), limit: limit, window: window}
	go func() {
		for {
			time.Sleep(window)
			l.mu.Lock()
			l.count = make(map[string]int)
			l.mu.Unlock()
		}
	}()
	return l
}

func (l *ipLimiter) allow(ip string) bool {
	l.mu.Lock()
	defer l.mu.Unlock()
	if l.count[ip] >= l.limit {
		return false
	}
	l.count[ip]++
	return true
}

func rateLimitMiddleware(limiter *ipLimiter) gin.HandlerFunc {
	return func(c *gin.Context) {
		if !limiter.allow(c.ClientIP()) {
			c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{"error": "请求过于频繁"})
			return
		}
		c.Next()
	}
}

func wsHandler(c *gin.Context) {
	conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
	if err != nil {
		return
	}
	defer conn.Close()
	for {
		if _, _, err := conn.ReadMessage(); err != nil {
			break
		}
	}
}

func main() {
	r := gin.Default()
	limiter := newIPLimiter(10, time.Minute) // 每个 IP 每分钟最多 10 次握手
	r.GET("/ws", rateLimitMiddleware(limiter), wsHandler)
	r.Run(":8080")
}

3. 连接数限制

除了限流,还要限制全局并发连接数,防止资源耗尽:

go
package main

import (
	"net/http"
	"sync/atomic"
	"github.com/gin-gonic/gin"
	"github.com/gorilla/websocket"
)

var upgrader = websocket.Upgrader{
	CheckOrigin: func(r *http.Request) bool { return true },
}

const maxConnections int64 = 10000
var currentConns int64

func connLimitMiddleware() gin.HandlerFunc {
	return func(c *gin.Context) {
		if atomic.LoadInt64(&currentConns) >= maxConnections {
			c.AbortWithStatusJSON(http.StatusServiceUnavailable, gin.H{"error": "连接数已满"})
			return
		}
		c.Next()
	}
}

func wsHandler(c *gin.Context) {
	conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
	if err != nil {
		return
	}
	atomic.AddInt64(&currentConns, 1)
	defer func() {
		atomic.AddInt64(&currentConns, -1)
		conn.Close()
	}()
	for {
		if _, _, err := conn.ReadMessage(); err != nil {
			break
		}
	}
}

func main() {
	r := gin.Default()
	r.GET("/ws", connLimitMiddleware(), wsHandler)
	r.Run(":8080")
}

三、连接管理器(Connection Manager)

上一篇的 Hub 是「广播」导向的。生产场景往往需要更精细的控制:按用户 ID 推送、按群组推送、查询某个用户是否在线。这就需要一个 Connection Manager

1. 全局连接池

go
package main

import (
	"sync"
	"github.com/gorilla/websocket"
)

// Conn 封装单个连接
type Conn struct {
	conn *websocket.Conn
	uid  string
	send chan []byte
}

// Manager 全局连接管理器
type Manager struct {
	mu    sync.RWMutex
	conns map[string]*Conn // uid -> 连接(单设备登录)
}

func NewManager() *Manager {
	return &Manager{conns: make(map[string]*Conn)}
}

func (m *Manager) Register(uid string, c *Conn) {
	m.mu.Lock()
	defer m.mu.Unlock()
	if old, ok := m.conns[uid]; ok {
		// 踢掉旧连接(单点登录)
		close(old.send)
	}
	m.conns[uid] = c
}

func (m *Manager) Unregister(uid string, c *Conn) {
	m.mu.Lock()
	defer m.mu.Unlock()
	if cur, ok := m.conns[uid]; ok && cur == c {
		delete(m.conns, uid)
	}
}

func (m *Manager) Get(uid string) (*Conn, bool) {
	m.mu.RLock()
	defer m.mu.RUnlock()
	c, ok := m.conns[uid]
	return c, ok
}

func (m *Manager) Count() int {
	m.mu.RLock()
	defer m.mu.RUnlock()
	return len(m.conns)
}

2. 用户与连接的映射

上面用 map[uid]*Conn 实现了一对一映射(单设备登录)。如果要支持多设备同时在线,改成 map[uid]map[*Conn]bool 即可。

go
package main

import (
	"sync"
	"github.com/gorilla/websocket"
)

type Conn struct {
	conn *websocket.Conn
	uid  string
	send chan []byte
}

// MultiManager 支持一个用户多个连接(多端登录)
type MultiManager struct {
	mu    sync.RWMutex
	conns map[string]map[*Conn]bool
}

func NewMultiManager() *MultiManager {
	return &MultiManager{conns: make(map[string]map[*Conn]bool)}
}

func (m *MultiManager) Register(uid string, c *Conn) {
	m.mu.Lock()
	defer m.mu.Unlock()
	if m.conns[uid] == nil {
		m.conns[uid] = make(map[*Conn]bool)
	}
	m.conns[uid][c] = true
}

func (m *MultiManager) Unregister(uid string, c *Conn) {
	m.mu.Lock()
	defer m.mu.Unlock()
	if set, ok := m.conns[uid]; ok {
		delete(set, c)
		if len(set) == 0 {
			delete(m.conns, uid)
		}
	}
}

func (m *MultiManager) IsOnline(uid string) bool {
	m.mu.RLock()
	defer m.mu.RUnlock()
	set, ok := m.conns[uid]
	return ok && len(set) > 0
}

func main() {
	mgr := NewMultiManager()
	c1 := &Conn{uid: "alice", send: make(chan []byte, 8)}
	c2 := &Conn{uid: "alice", send: make(chan []byte, 8)} // 多端登录
	mgr.Register("alice", c1)
	mgr.Register("alice", c2)
	_ = mgr.IsOnline("alice")
	mgr.Unregister("alice", c1)
	mgr.Unregister("alice", c2)
}

3. 按 ID 推送消息

go
package main

import (
	"sync"
	"time"
	"github.com/gorilla/websocket"
)

type Conn struct {
	conn *websocket.Conn
	uid  string
	send chan []byte
}

type Manager struct {
	mu    sync.RWMutex
	conns map[string]*Conn
}

func NewManager() *Manager {
	return &Manager{conns: make(map[string]*Conn)}
}

func (m *Manager) Register(uid string, c *Conn) {
	m.mu.Lock()
	defer m.mu.Unlock()
	m.conns[uid] = c
}

func (m *Manager) Unregister(uid string) {
	m.mu.Lock()
	defer m.mu.Unlock()
	delete(m.conns, uid)
}

// PushToUser 给指定用户推送消息
func (m *Manager) PushToUser(uid string, msg []byte) bool {
	m.mu.RLock()
	c, ok := m.conns[uid]
	m.mu.RUnlock()
	if !ok {
		return false // 用户不在线
	}
	select {
	case c.send <- msg:
		return true
	default:
		// 缓冲满,推送失败
		return false
	}
}

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

func main() {
	mgr := NewManager()
	c := &Conn{uid: "alice", send: make(chan []byte, 8)}
	mgr.Register("alice", c)
	mgr.PushToUser("alice", []byte("hello"))
	mgr.Unregister("alice")
}

四、消息分发

有了连接管理器,就可以实现更丰富的分发逻辑。

1. 一对一私聊

go
package main

import (
	"encoding/json"
	"sync"
	"time"
	"github.com/gorilla/websocket"
)

type Conn struct {
	conn *websocket.Conn
	uid  string
	send chan []byte
}

type Manager struct {
	mu    sync.RWMutex
	conns map[string]*Conn
}

func NewManager() *Manager { return &Manager{conns: make(map[string]*Conn)} }

func (m *Manager) get(uid string) (*Conn, bool) {
	m.mu.RLock()
	defer m.mu.RUnlock()
	c, ok := m.conns[uid]
	return c, ok
}

func (m *Manager) register(uid string, c *Conn) {
	m.mu.Lock()
	m.conns[uid] = c
	m.mu.Unlock()
}

func (m *Manager) unregister(uid string) {
	m.mu.Lock()
	delete(m.conns, uid)
	m.mu.Unlock()
}

type PrivateMsg struct {
	From string `json:"from"`
	To   string `json:"to"`
	Text string `json:"text"`
	Time string `json:"time"`
}

// SendPrivate 发送私聊:给目标用户推送,给发送者也回一份
func (m *Manager) SendPrivate(from, to, text string) bool {
	msg := PrivateMsg{From: from, To: to, Text: text,
		Time: time.Now().Format("15:04:05")}
	data, _ := json.Marshal(msg)

	ok := false
	if c, found := m.get(to); found {
		select {
		case c.send <- data:
			ok = true
		default:
		}
	}
	// 给发送者也回一份(多端同步)
	if c, found := m.get(from); found {
		select {
		case c.send <- data:
		default:
		}
	}
	return ok
}

func main() {
	mgr := NewManager()
	alice := &Conn{uid: "alice", send: make(chan []byte, 8)}
	bob := &Conn{uid: "bob", send: make(chan []byte, 8)}
	mgr.register("alice", alice)
	mgr.register("bob", bob)
	mgr.SendPrivate("alice", "bob", "你好")
}

2. 群组广播

群组用 map[groupID]map[uid]bool 维护成员关系,广播时遍历成员逐个推送:

go
package main

import (
	"encoding/json"
	"sync"
	"time"
	"github.com/gorilla/websocket"
)

type Conn struct {
	conn *websocket.Conn
	uid  string
	send chan []byte
}

type Manager struct {
	mu     sync.RWMutex
	conns  map[string]*Conn
	groups map[string]map[string]bool // groupID -> uid set
}

func NewManager() *Manager {
	return &Manager{conns: make(map[string]*Conn), groups: make(map[string]map[string]bool)}
}

func (m *Manager) joinGroup(groupID, uid string) {
	m.mu.Lock()
	defer m.mu.Unlock()
	if m.groups[groupID] == nil {
		m.groups[groupID] = make(map[string]bool)
	}
	m.groups[groupID][uid] = true
}

func (m *Manager) leaveGroup(groupID, uid string) {
	m.mu.Lock()
	defer m.mu.Unlock()
	if set, ok := m.groups[groupID]; ok {
		delete(set, uid)
	}
}

type GroupMsg struct {
	Group string `json:"group"`
	From  string `json:"from"`
	Text  string `json:"text"`
	Time  string `json:"time"`
}

func (m *Manager) BroadcastGroup(groupID, from, text string) int {
	m.mu.RLock()
	set := m.groups[groupID]
	if set == nil {
		m.mu.RUnlock()
		return 0
	}
	uids := make([]string, 0, len(set))
	for uid := range set {
		uids = append(uids, uid)
	}
	m.mu.RUnlock()

	msg := GroupMsg{Group: groupID, From: from, Text: text,
		Time: time.Now().Format("15:04:05")}
	data, _ := json.Marshal(msg)

	sent := 0
	for _, uid := range uids {
		m.mu.RLock()
		c, ok := m.conns[uid]
		m.mu.RUnlock()
		if !ok {
			continue
		}
		select {
		case c.send <- data:
			sent++
		default:
		}
	}
	return sent
}

func main() {
	mgr := NewManager()
	mgr.joinGroup("room1", "alice")
	mgr.joinGroup("room1", "bob")
	mgr.BroadcastGroup("room1", "alice", "大家好")
	mgr.leaveGroup("room1", "bob")
}

3. 系统通知

系统通知是给所有在线用户推一条消息(如「服务器将于 10 分钟后维护」):

go
package main

import (
	"encoding/json"
	"sync"
	"time"
	"github.com/gorilla/websocket"
)

type Conn struct {
	conn *websocket.Conn
	uid  string
	send chan []byte
}

type Manager struct {
	mu    sync.RWMutex
	conns map[string]*Conn
}

func NewManager() *Manager { return &Manager{conns: make(map[string]*Conn)} }

type Notice struct {
	Type string `json:"type"`
	Text string `json:"text"`
	Time string `json:"time"`
}

func (m *Manager) BroadcastNotice(text string) int {
	msg := Notice{Type: "notice", Text: text,
		Time: time.Now().Format("15:04:05")}
	data, _ := json.Marshal(msg)

	m.mu.RLock()
	targets := make([]*Conn, 0, len(m.conns))
	for _, c := range m.conns {
		targets = append(targets, c)
	}
	m.mu.RUnlock()

	sent := 0
	for _, c := range targets {
		select {
		case c.send <- data:
			sent++
		default:
		}
	}
	return sent
}

func main() {
	mgr := NewManager()
	alice := &Conn{uid: "alice", send: make(chan []byte, 8)}
	mgr.conns["alice"] = alice
	mgr.BroadcastNotice("服务器将于 10 分钟后维护")
}

五、生产级特性

1. 断线重连(客户端实现)

网络抖动会导致连接断开,客户端需要自动重连。重连要点:

  • 指数退避:第一次 1 秒、第二次 2 秒、第三次 4 秒……避免雪崩。
  • 上限:退避时间封顶(如 30 秒)。
  • 重连后恢复订阅:把之前的群组、私聊关系重新建立。
go
package main

import (
	"fmt"
	"math"
	"net/url"
	"time"
	"github.com/gorilla/websocket"
)

func connectWithRetry(addr string, maxRetries int) (*websocket.Conn, error) {
	var lastErr error
	for i := 0; i < maxRetries; i++ {
		u := url.URL{Scheme: "ws", Host: addr, Path: "/ws"}
		c, _, err := websocket.DefaultDialer.Dial(u.String(), nil)
		if err == nil {
			return c, nil
		}
		lastErr = err
		// 指数退避:2^i 秒,上限 30 秒
		backoff := time.Duration(math.Pow(2, float64(i))) * time.Second
		if backoff > 30*time.Second {
			backoff = 30 * time.Second
		}
		fmt.Printf("第 %d 次重连失败,%v 后重试: %v\n", i+1, backoff, err)
		time.Sleep(backoff)
	}
	return nil, lastErr
}

func main() {
	c, err := connectWithRetry("localhost:8080", 6)
	if err != nil {
		fmt.Println("重连失败:", err)
		return
	}
	defer c.Close()
	fmt.Println("已连接")

	// 模拟断线后自动重连
	go func() {
		for {
			_, _, err := c.ReadMessage()
			if err != nil {
				fmt.Println("断线,尝试重连...")
				c, err = connectWithRetry("localhost:8080", 6)
				if err != nil {
					return
				}
			}
		}
	}()

	select {}
}

2. 消息队列缓冲

每个 Client 的 send channel 就是天然的缓冲队列。但 channel 容量有限,如果消费者持续跟不上,消息会被丢弃。生产场景可以引入外部消息队列做持久缓冲,或对慢消费者做背压(关闭连接、降级推送)。

go
package main

import (
	"sync"
	"github.com/gorilla/websocket"
)

type Conn struct {
	conn *websocket.Conn
	send chan []byte
}

type Manager struct {
	mu    sync.RWMutex
	conns map[string]*Conn
}

// SafeSend 带背压的发送:缓冲满时返回 false,调用方决定降级策略
func (m *Manager) SafeSend(c *Conn, msg []byte) bool {
	select {
	case c.send <- msg:
		return true
	default:
		return false // 触发背压:降级或断开
	}
}

func main() {
	mgr := &Manager{conns: make(map[string]*Conn)}
	c := &Conn{send: make(chan []byte, 2)}
	mgr.SafeSend(c, []byte("msg1"))
	mgr.SafeSend(c, []byte("msg2"))
	mgr.SafeSend(c, []byte("msg3")) // 缓冲满,触发背压返回 false
}

3. 分布式 WebSocket:Redis Pub/Sub 跨节点广播

单机 WebSocket 只能服务本机连接。多节点部署时,节点 A 上的用户发的消息要让节点 B 上的用户也收到,就需要跨节点广播。最常用的方案是 Redis Pub/Sub:每个节点订阅同一个频道,发消息时 publish 到 Redis,所有节点收到后推给本机连接。

go
package main

import (
	"context"
	"encoding/json"
	"fmt"
	"sync"
	"github.com/go-redis/redis/v8"
	"github.com/gorilla/websocket"
)

type Conn struct {
	conn *websocket.Conn
	send chan []byte
}

type Manager struct {
	mu    sync.RWMutex
	conns map[string]*Conn
}

func NewManager() *Manager {
	return &Manager{conns: make(map[string]*Conn)}
}

func (m *Manager) localPush(uid string, msg []byte) {
	m.mu.RLock()
	c, ok := m.conns[uid]
	m.mu.RUnlock()
	if ok {
		select {
		case c.send <- msg:
		default:
		}
	}
}

// Redis 跨节点广播
type RedisBridge struct {
	rdb    *redis.Client
	mgr    *Manager
	channel string
}

func NewRedisBridge(addr, channel string, mgr *Manager) *RedisBridge {
	rdb := redis.NewClient(&redis.Options{Addr: addr})
	return &RedisBridge{rdb: rdb, mgr: mgr, channel: channel}
}

// Subscribe 订阅 Redis 频道,把收到的消息推给本机对应用户
func (b *RedisBridge) Subscribe(ctx context.Context) {
	sub := b.rdb.Subscribe(ctx, b.channel)
	ch := sub.Channel()
	for msg := range ch {
		// payload 格式: {"uid":"user1","data":"..."}
		var payload struct {
			UID  string `json:"uid"`
			Data []byte `json:"data"`
		}
		if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil {
			continue
		}
		// 只推给本机存在的连接
		b.mgr.localPush(payload.UID, payload.Data)
	}
}

// Publish 把消息发布到 Redis,所有节点都会收到
func (b *RedisBridge) Publish(ctx context.Context, uid string, data []byte) error {
	payload, _ := json.Marshal(struct {
		UID  string `json:"uid"`
		Data []byte `json:"data"`
	}{UID: uid, Data: data})
	return b.rdb.Publish(ctx, b.channel, payload).Err()
}

func main() {
	mgr := NewManager()
	bridge := NewRedisBridge("localhost:6379", "ws_broadcast", mgr)

	ctx := context.Background()
	go bridge.Subscribe(ctx)

	// 模拟收到一条要广播的消息
	if err := bridge.Publish(ctx, "user1", []byte("跨节点推送")); err != nil {
		fmt.Println("发布失败:", err)
	}
	_ = websocket.TextMessage
	select {}
}

4. 消息持久化

聊天消息通常要落库,方便历史消息查询和审计。可以在分发消息的同时异步写入数据库:

go
package main

import (
	"sync"
	"time"
	"github.com/gorilla/websocket"
)

type Conn struct {
	conn *websocket.Conn
	send chan []byte
}

type ChatRecord struct {
	From    string
	To      string
	Content string
	Time    time.Time
}

type Manager struct {
	mu      sync.RWMutex
	conns   map[string]*Conn
	saveCh  chan ChatRecord // 异步保存队列
}

func NewManager() *Manager {
	m := &Manager{
		conns:  make(map[string]*Conn),
		saveCh: make(chan ChatRecord, 1024),
	}
	go m.persistLoop()
	return m
}

// 后台消费保存队列,批量写入数据库
func (m *Manager) persistLoop() {
	batch := make([]ChatRecord, 0, 100)
	ticker := time.NewTicker(time.Second)
	defer ticker.Stop()
	for {
		select {
		case r := <-m.saveCh:
			batch = append(batch, r)
			if len(batch) >= 100 {
				m.flush(batch)
				batch = batch[:0]
			}
		case <-ticker.C:
			if len(batch) > 0 {
				m.flush(batch)
				batch = batch[:0]
			}
		}
	}
}

func (m *Manager) flush(batch []ChatRecord) {
	// 实际项目:批量 INSERT 到数据库
	// db.Exec("INSERT INTO messages ...", ...)
}

func main() {
	mgr := NewManager()
	mgr.saveCh <- ChatRecord{From: "alice", To: "bob", Content: "hi"}
}

异步批量写入比每条消息同步写库高效得多,是高吞吐场景的标准做法。

六、完整示例:Gin + WebSocket 实时通知系统

把前面的能力组合起来,实现一个带认证、连接管理、私聊、群组、Redis 跨节点广播的通知系统。为控制篇幅,这里给出核心结构(省略部分重复实现)。

go
package main

import (
	"context"
	"encoding/json"
	"fmt"
	"log"
	"net/http"
	"time"
	"github.com/gin-gonic/gin"
	"github.com/go-redis/redis/v8"
	"github.com/gorilla/websocket"
)

const (
	writeWait  = 10 * time.Second
	pongWait   = 60 * time.Second
	pingPeriod = 50 * time.Second
)

var upgrader = websocket.Upgrader{
	CheckOrigin: func(r *http.Request) bool { return true },
}

// ---- 连接 ----
type Conn struct {
	ws   *websocket.Conn
	uid  string
	send chan []byte
}

// ---- 管理器 ----
type Manager struct {
	conns  map[string]*Conn
	add    chan *Conn
	remove chan *Conn
	bcast  chan []byte
}

func NewManager() *Manager {
	return &Manager{
		conns:  make(map[string]*Conn),
		add:    make(chan *Conn),
		remove: make(chan *Conn),
		bcast:  make(chan []byte, 256),
	}
}

func (m *Manager) Run(ctx context.Context, bridge *RedisBridge) {
	for {
		select {
		case c := <-m.add:
			m.conns[c.uid] = c
		case c := <-m.remove:
			if cur, ok := m.conns[c.uid]; ok && cur == c {
				delete(m.conns, c.uid)
			}
		case msg := <-m.bcast:
			for _, c := range m.conns {
				select {
				case c.send <- msg:
				default:
				}
			}
		case <-ctx.Done():
			return
		}
	}
}

// ---- Redis 桥 ----
type RedisBridge struct {
	rdb     *redis.Client
	channel string
	mgr     *Manager
}

func NewRedisBridge(addr, channel string, mgr *Manager) *RedisBridge {
	return &RedisBridge{
		rdb:     redis.NewClient(&redis.Options{Addr: addr}),
		channel: channel,
		mgr:     mgr,
	}
}

func (b *RedisBridge) Subscribe(ctx context.Context) {
	sub := b.rdb.Subscribe(ctx, b.channel)
	for msg := range sub.Channel() {
		// 收到 Redis 消息,投递到本机 bcast
		b.mgr.bcast <- []byte(msg.Payload)
	}
}

func (b *RedisBridge) Publish(ctx context.Context, payload []byte) error {
	return b.rdb.Publish(ctx, b.channel, payload).Err()
}

// ---- 读写 goroutine ----
func (c *Conn) readPump(mgr *Manager, bridge *RedisBridge) {
	defer func() {
		mgr.remove <- c
		c.ws.Close()
	}()
	c.ws.SetReadLimit(4096)
	c.ws.SetReadDeadline(time.Now().Add(pongWait))
	c.ws.SetPongHandler(func(string) error {
		c.ws.SetReadDeadline(time.Now().Add(pongWait))
		return nil
	})
	for {
		_, data, err := c.ws.ReadMessage()
		if err != nil {
			break
		}
		// 把消息发布到 Redis,所有节点广播
		_ = bridge.Publish(context.Background(), data)
	}
}

func (c *Conn) writePump() {
	ticker := time.NewTicker(pingPeriod)
	defer func() {
		ticker.Stop()
		c.ws.Close()
	}()
	for {
		select {
		case msg, ok := <-c.send:
			c.ws.SetWriteDeadline(time.Now().Add(writeWait))
			if !ok {
				c.ws.WriteMessage(websocket.CloseMessage, []byte{})
				return
			}
			if err := c.ws.WriteMessage(websocket.TextMessage, msg); err != nil {
				return
			}
		case <-ticker.C:
			c.ws.SetWriteDeadline(time.Now().Add(writeWait))
			if err := c.ws.WriteMessage(websocket.PingMessage, nil); err != nil {
				return
			}
		}
	}
}

// ---- 认证 ----
func authMiddleware() gin.HandlerFunc {
	return func(c *gin.Context) {
		token := c.Query("token")
		if token == "" {
			c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "缺少 token"})
			return
		}
		uid := token // 简化:实际应解析 JWT
		c.Set("uid", uid)
		c.Next()
	}
}

func main() {
	ctx := context.Background()
	mgr := NewManager()
	bridge := NewRedisBridge("localhost:6379", "notify", mgr)

	go mgr.Run(ctx, bridge)
	go bridge.Subscribe(ctx)

	r := gin.Default()
	r.GET("/ws", authMiddleware(), func(c *gin.Context) {
		uid, _ := c.Get("uid")
		conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
		if err != nil {
			return
		}
		client := &Conn{
			ws:   conn,
			uid:  uid.(string),
			send: make(chan []byte, 256),
		}
		mgr.add <- client
		go client.writePump()
		go client.readPump(mgr, bridge)

		// 推送欢迎消息
		hello, _ := json.Marshal(map[string]string{
			"type": "welcome", "text": "已连接通知服务"})
		client.send <- hello
	})

	// HTTP 接口:向所有人广播一条通知
	r.POST("/notify", func(c *gin.Context) {
		var body struct {
			Text string `json:"text"`
		}
		if err := c.ShouldBindJSON(&body); err != nil {
			c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
			return
		}
		payload, _ := json.Marshal(map[string]string{
			"type": "notice", "text": body.Text,
			"time": time.Now().Format("15:04:05")})
		// 发布到 Redis,所有节点都会推给本机连接
		if err := bridge.Publish(ctx, payload); err != nil {
			c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
			return
		}
		c.JSON(http.StatusOK, gin.H{"ok": true})
	})

	fmt.Println("通知服务监听 :8080")
	log.Fatal(r.Run(":8080"))
}

这个示例的核心数据流:HTTP /notify → Redis Publish → 所有节点 Subscribe → 各节点 Manager.bcast → 本机所有连接。这就实现了多节点部署下的全局广播。

七、小结

本篇我们把 WebSocket 接入 Gin,构建了一套生产级架构。重点包括:用 c.Writer/c.Request 在 Gin Handler 中完成升级;用认证、限流、连接数限制中间件保护握手入口;用 Connection Manager 维护用户与连接的映射,支持按 ID 推送、私聊、群组广播、系统通知;用 Redis Pub/Sub 实现跨节点广播,支持多节点水平扩展;用异步批量写入实现消息持久化;客户端用指数退避实现断线重连。

关键要点回顾:

  • Gin 集成upgrader.Upgrade(c.Writer, c.Request, nil) 是核心,握手后中间件失效。
  • 认证:token 放 query 最通用,注意短期 token 的安全实践。
  • Connection Manager:用 map 维护 uid→连接,支持单设备/多设备两种模式。
  • 分发:私聊查 map、群组遍历成员、通知遍历全部,均用非阻塞发送防卡死。
  • 分布式:Redis Pub/Sub 是跨节点广播的最简方案,多节点共享一个频道。
  • 生产特性:断线重连(指数退避)、消息缓冲(channel + 背压)、持久化(异步批量写库)。

下一篇我们将聚焦性能与高并发,分析 WebSocket 的瓶颈、优化策略、分布式架构演进,以及如何压测 10 万连接。