// client.go 提供 wsc 服务端的客户端封装(Client)。 // 对称设计:cli.On == router.On,cli.Send == ctx.WriteMessage。 // 客户端与服务端共用同一 Message 信封与协议约定,改协议两边一起改。 package wsc import ( "context" "encoding/json" "errors" "net/http" "reflect" "sync" "sync/atomic" "time" "github.com/gorilla/websocket" "github.com/sirupsen/logrus" ) // ============================================================ // 配置 // ============================================================ // ClientConfig 客户端配置。 type ClientConfig struct { URL string // WebSocket 地址(必填) Header http.Header // 自定义请求头(如 token) DialTimeout time.Duration // 连接超时,默认 10s AutoReconnect bool // 自动重连 MaxRetry int // 最大重试次数(0 = 无限),默认 0 MinReconDelay time.Duration // 最小重连延迟,默认 1s MaxReconDelay time.Duration // 最大重连延迟,默认 30s PingInterval time.Duration // 心跳间隔(0 = 不发送),默认 25s PongTimeout time.Duration // 心跳超时,默认 60s WriteTimeout time.Duration // 写超时,默认 10s SendBufferSize int // 发送缓冲区大小,默认 256 } // DefaultClientConfig 返回推荐默认配置。 func DefaultClientConfig(url string) ClientConfig { return ClientConfig{ URL: url, DialTimeout: 10 * time.Second, AutoReconnect: true, MaxRetry: 0, MinReconDelay: 1 * time.Second, MaxReconDelay: 30 * time.Second, PingInterval: 25 * time.Second, PongTimeout: 60 * time.Second, WriteTimeout: 10 * time.Second, SendBufferSize: 256, } } // ============================================================ // PayloadHandler — 客户端消息处理器 // ============================================================ // PayloadHandler 客户端消息处理函数。payload 为服务端发来的 JSON,已去掉 action 包装。 type PayloadHandler func(payload json.RawMessage) // ============================================================ // Client // ============================================================ // Client WebSocket 客户端,并发安全。 // 所有公开方法可在任意 goroutine 调用。 type Client struct { url string urlMu sync.RWMutex cfg ClientConfig // 消息路由 mu sync.RWMutex routes map[string]PayloadHandler // 连接 conn *websocket.Conn connMu sync.Mutex writeCh chan []byte closeCh chan struct{} doneCh chan struct{} // 重连 reconnecting atomic.Bool reconMu sync.Mutex // 保护 reconDelay / reconCount / reconStableTimer reconDelay time.Duration reconCount int reconStop chan struct{} reconDone chan struct{} reconStableTimer *time.Timer // 连接稳定后重置退避的定时器 // 回调 OnConnected func() OnDisconnected func(err error) OnReconnecting func(attempt int, delay time.Duration) // 生命周期控制 closed atomic.Bool closeOnce sync.Once wg sync.WaitGroup } // New 创建客户端(尚未连接)。 func NewClient(cfg ClientConfig) *Client { if cfg.DialTimeout == 0 { cfg.DialTimeout = 10 * time.Second } if cfg.MinReconDelay == 0 { cfg.MinReconDelay = time.Second } if cfg.MaxReconDelay == 0 { cfg.MaxReconDelay = 30 * time.Second } if cfg.PingInterval == 0 { cfg.PingInterval = 25 * time.Second } if cfg.PongTimeout == 0 { cfg.PongTimeout = 60 * time.Second } if cfg.WriteTimeout == 0 { cfg.WriteTimeout = 10 * time.Second } if cfg.SendBufferSize == 0 { cfg.SendBufferSize = 256 } return &Client{ url: cfg.URL, cfg: cfg, routes: make(map[string]PayloadHandler), writeCh: make(chan []byte, cfg.SendBufferSize), closeCh: make(chan struct{}), doneCh: make(chan struct{}), } } // ============================================================ // 连接控制 // ============================================================ // Connect 连接 WebSocket 服务端。AutoReconnect=true 时阻塞直到连接成功或达到最大重试。 func (c *Client) Connect() error { return c.connect() } // Close 关闭连接,停止重连。 func (c *Client) Close() { c.closeOnce.Do(func() { c.closed.Store(true) close(c.closeCh) // 停止重连 select { case c.reconStop <- struct{}{}: default: } c.cancelBackoffReset() c.connMu.Lock() if c.conn != nil { c.conn.Close() } c.connMu.Unlock() c.wg.Wait() close(c.doneCh) }) } // Connected 返回当前是否已连接。 func (c *Client) Connected() bool { c.connMu.Lock() defer c.connMu.Unlock() return c.conn != nil } // Done 返回一个通道,Client 完全关闭后关闭。 func (c *Client) Done() <-chan struct{} { return c.doneCh } // SetURL 动态更新连接地址(含 token)。并发安全,下一次 dial/重连时生效。 // 注意:已建立的连接不会立即断开,仍使用旧地址直到下一次重连。 func (c *Client) SetURL(url string) { c.urlMu.Lock() c.url = url c.urlMu.Unlock() } // getURL 并发安全地读取当前连接地址。 func (c *Client) getURL() string { c.urlMu.RLock() defer c.urlMu.RUnlock() return c.url } // ============================================================ // 消息路由(注册后再 Connect) // ============================================================ // On 注册 action 对应的消息处理器。 func (c *Client) On(action string, handler PayloadHandler) { c.mu.Lock() c.routes[action] = handler c.mu.Unlock() } // Send 发送结构化消息。并发安全。 // 与服务端同用 Message 信封:{"action":"...","payload":...},payload 为 nil 时字段省略。 func (c *Client) Send(action string, payload any) error { msg := Message{Action: action} if payload != nil { raw, err := json.Marshal(payload) if err != nil { return err } msg.Payload = raw } data, err := json.Marshal(msg) if err != nil { return err } return c.SendRaw(data) } // SendRaw 发送已序列化的字节。并发安全。 func (c *Client) SendRaw(data []byte) error { if c.closed.Load() { return errors.New("wsc: client closed") } select { case c.writeCh <- data: return nil case <-c.closeCh: return errors.New("wsc: client closed") default: logrus.Warnf("[wsc-client] send buffer full, dropping message") return errors.New("wsc: send buffer full") } } // ============================================================ // Bind — 自动 JSON 反序列化 // ============================================================ // BindPayload 将带类型的函数包装为 PayloadHandler,自动 JSON 反序列化 payload。 // // cli.On("ping.resp", wsc.BindPayload(func(resp *PingResp) { // log.Println(resp.Message) // })) func BindPayload(fn any) PayloadHandler { fnVal := reflect.ValueOf(fn) fnType := fnVal.Type() if fnType.Kind() != reflect.Func || fnType.NumIn() != 1 { panic("wsc.BindPayload: function must have 1 parameter") } reqType := fnType.In(0) if reqType.Kind() != reflect.Ptr { panic("wsc.BindPayload: parameter must be a pointer") } return func(raw json.RawMessage) { req := reflect.New(reqType.Elem()).Interface() if len(raw) > 0 { if err := json.Unmarshal(raw, req); err != nil { logrus.Errorf("[wsc-client] Bind unmarshal error: %v, raw=%s", err, string(raw)) return } } fnVal.Call([]reflect.Value{reflect.ValueOf(req)}) } } // ============================================================ // 内部:连接与重连 // ============================================================ func (c *Client) connect() error { for { err := c.dial() if err == nil { return nil } if !c.cfg.AutoReconnect { return err } if c.closed.Load() { return errors.New("wsc: closed") } c.reconCount++ if c.cfg.MaxRetry > 0 && c.reconCount > c.cfg.MaxRetry { return err } delay := c.nextDelay() logrus.Warnf("[wsc-client] connect failed (attempt %d): %v, retry in %v", c.reconCount, err, delay) if c.OnReconnecting != nil { c.OnReconnecting(c.reconCount, delay) } select { case <-time.After(delay): case <-c.closeCh: return errors.New("wsc: closed during reconnect") } } } func (c *Client) dial() error { dialer := websocket.Dialer{ HandshakeTimeout: c.cfg.DialTimeout, } if c.cfg.Header != nil { dialer.Proxy = http.ProxyFromEnvironment // 无操作,只是占位 } ctx, cancel := context.WithTimeout(context.Background(), c.cfg.DialTimeout) defer cancel() conn, _, err := dialer.DialContext(ctx, c.getURL(), c.cfg.Header) if err != nil { return err } c.connMu.Lock() c.conn = conn c.connMu.Unlock() c.scheduleBackoffReset() // 重启内部通道(每次重连重新创建) if c.reconStop != nil { select { case c.reconStop <- struct{}{}: default: } } c.reconStop = make(chan struct{}, 1) c.reconDone = make(chan struct{}) // 启动读写协程 c.wg.Add(2) go c.readPump() go c.writePump() if c.OnConnected != nil { c.OnConnected() } return nil } // ============================================================ // 内部:读写协程 // ============================================================ func (c *Client) readPump() { defer c.wg.Done() defer func() { if r := recover(); r != nil { logrus.Errorf("[wsc-client] readPump panic: %v", r) } c.onDisconnect() }() c.connMu.Lock() conn := c.conn c.connMu.Unlock() if conn == nil { return } conn.SetReadLimit(65536) conn.SetReadDeadline(time.Now().Add(c.cfg.PongTimeout)) conn.SetPongHandler(func(string) error { conn.SetReadDeadline(time.Now().Add(c.cfg.PongTimeout)) return nil }) for { _, data, err := conn.ReadMessage() if err != nil { if !c.closed.Load() { if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) && !errors.Is(err, context.DeadlineExceeded) { logrus.Errorf("[wsc-client] read error: %v", err) } } return } var msg Message if err := json.Unmarshal(data, &msg); err != nil { continue } c.mu.RLock() handler, ok := c.routes[msg.Action] c.mu.RUnlock() if ok { handler(msg.Payload) } else { logrus.Warnf("[wsc-client] unhandled message action=%q, payload=%s", msg.Action, string(msg.Payload)) } } } func (c *Client) writePump() { defer c.wg.Done() c.connMu.Lock() conn := c.conn c.connMu.Unlock() if conn == nil { return } defer conn.Close() pingTicker := time.NewTicker(c.cfg.PingInterval) defer pingTicker.Stop() for { select { case data, ok := <-c.writeCh: if !ok { return } conn.SetWriteDeadline(time.Now().Add(c.cfg.WriteTimeout)) if err := conn.WriteMessage(websocket.TextMessage, data); err != nil { logrus.Errorf("[wsc-client] write error: %v", err) return } case <-pingTicker.C: conn.SetWriteDeadline(time.Now().Add(c.cfg.WriteTimeout)) if err := conn.WriteMessage(websocket.PingMessage, nil); err != nil { return } case <-c.reconStop: return case <-c.closeCh: return } } } // ============================================================ // 内部:断连处理 // ============================================================ func (c *Client) onDisconnect() { c.connMu.Lock() c.conn = nil c.connMu.Unlock() if c.OnDisconnected != nil { c.OnDisconnected(errors.New("connection lost")) } if c.closed.Load() || !c.cfg.AutoReconnect { close(c.writeCh) return } // 连接已断开,取消“稳定后重置退避”的定时器,使退避继续累积 c.cancelBackoffReset() // 启动重连 goroutine if c.reconnecting.CompareAndSwap(false, true) { go c.reconnectLoop() } } func (c *Client) reconnectLoop() { defer c.reconnecting.Store(false) for { if c.closed.Load() { return } conn, _, err := (&websocket.Dialer{ HandshakeTimeout: c.cfg.DialTimeout, }).DialContext(context.Background(), c.getURL(), c.cfg.Header) if err == nil { c.scheduleBackoffReset() c.connMu.Lock() c.conn = conn c.connMu.Unlock() // 新 writeCh(旧的可能还有残留,丢弃) c.writeCh = make(chan []byte, c.cfg.SendBufferSize) c.reconStop = make(chan struct{}, 1) c.wg.Add(2) go c.readPump() go c.writePump() if c.OnConnected != nil { c.OnConnected() } return } c.reconCount++ if c.cfg.MaxRetry > 0 && c.reconCount > c.cfg.MaxRetry { logrus.Errorf("[wsc-client] reconnect max retry exceeded (%d)", c.cfg.MaxRetry) c.Close() return } delay := c.nextDelay() logrus.Warnf("[wsc-client] reconnect attempt %d failed: %v, retry in %v", c.reconCount, err, delay) if c.OnReconnecting != nil { c.OnReconnecting(c.reconCount, delay) } select { case <-time.After(delay): case <-c.reconStop: return case <-c.closeCh: return } } } // ============================================================ // 内部:辅助 // ============================================================ // defaultReconStableGrace 连接持续稳定达到该时长后,重连退避才重置为初始值。 // 避免在连接频繁抖动(握手成功但随即断开)时退避被反复清零,导致“永远 1 秒重连一次”。 const defaultReconStableGrace = 10 * time.Second // scheduleBackoffReset 在连接稳定持续 grace 后,将退避延迟与重试计数重置为初始值。 // 若在此期间连接再次断开(cancelBackoffReset),则退避继续累积,不会被清零。 func (c *Client) scheduleBackoffReset() { c.reconMu.Lock() defer c.reconMu.Unlock() if c.reconStableTimer != nil { c.reconStableTimer.Stop() } c.reconStableTimer = time.AfterFunc(defaultReconStableGrace, func() { c.reconMu.Lock() c.reconDelay = 0 c.reconCount = 0 c.reconStableTimer = nil c.reconMu.Unlock() }) } // cancelBackoffReset 取消待执行的退避重置(连接再次断开时调用)。 func (c *Client) cancelBackoffReset() { c.reconMu.Lock() defer c.reconMu.Unlock() if c.reconStableTimer != nil { c.reconStableTimer.Stop() c.reconStableTimer = nil } } func (c *Client) nextDelay() time.Duration { c.reconMu.Lock() defer c.reconMu.Unlock() if c.reconDelay == 0 { c.reconDelay = c.cfg.MinReconDelay } delay := c.reconDelay c.reconDelay *= 2 if c.reconDelay > c.cfg.MaxReconDelay { c.reconDelay = c.cfg.MaxReconDelay } return delay }