Files
zogo/wsc/client.go
T

591 lines
14 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}