feat: wscclient 并入 wsc(wsc.Client), 信封/心跳约定单一来源; 客户端 API 改名 NewClient/ClientConfig/BindPayload/PayloadHandler; 补各包 README

This commit is contained in:
w11
2026-09-19 20:11:27 +08:00
parent f0f5421262
commit 9ef71f51d4
6 changed files with 180 additions and 49 deletions
+590
View File
@@ -0,0 +1,590 @@
// 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
}