591 lines
14 KiB
Go
591 lines
14 KiB
Go
// 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
|
||
}
|