Files
common/websocket/server.go
T
19904408334 a66a38e074 feat: 新增通用工具框架与WebSocket服务
* 新增 tools 包:统一工具定义、注册表与 Server 接口,对齐 MCP 规范
* 新增 websocket 包:泛化连接管理、心跳、并发写锁与优雅关闭
* 新增参数读取工具函数,避免类型断言静默失败
* 新增 OSS 路径识别与 JSON 扁平映射还原工具
* 修复租户 SQL 条件插入位置,正确处理 GROUP BY 与 ORDER BY 同时出现的场景
2026-08-21 09:26:35 +08:00

327 lines
7.8 KiB
Go

package websocket
import (
"context"
"errors"
"fmt"
"sync"
"sync/atomic"
"time"
"github.com/gogf/gf/v2/container/gmap"
"github.com/gogf/gf/v2/encoding/gjson"
"github.com/gogf/gf/v2/frame/g"
"github.com/gogf/gf/v2/net/ghttp"
"github.com/gogf/gf/v2/os/glog"
"github.com/gogf/gf/v2/os/grpool"
"github.com/google/uuid"
"github.com/gorilla/websocket"
)
// WsServer 泛化 WebSocket 服务器
type WsServer struct {
connections *gmap.StrAnyMap
upgrader websocket.Upgrader
workerPool *grpool.Pool
handlers map[string]MessageHandler
handlerMu sync.RWMutex
opts ServerOptions
svcClosed int32
closeOnce sync.Once
}
// NewWsServer 创建泛化 WebSocket 服务器
func NewWsServer(opts ...ServerOption) *WsServer {
o := defaultOptions()
for _, opt := range opts {
opt(&o)
}
return &WsServer{
connections: gmap.NewStrAnyMap(true),
upgrader: websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
CheckOrigin: o.checkOrigin,
},
workerPool: grpool.New(o.workerPoolSize),
handlers: make(map[string]MessageHandler),
opts: o,
svcClosed: 0,
}
}
// OnMessage 注册业务消息处理器
func (s *WsServer) OnMessage(msgType string, handler MessageHandler) {
s.handlerMu.Lock()
defer s.handlerMu.Unlock()
s.handlers[msgType] = handler
}
// Upgrade 将 HTTP 连接升级为 WebSocket 并注册到连接池
func (s *WsServer) Upgrade(ctx context.Context, r *ghttp.Request, sessionId string) (*WsConnection, error) {
if g.IsEmpty(sessionId) {
sessionId = uuid.NewString()
}
if atomic.LoadInt32(&s.svcClosed) == 1 {
return nil, errors.New("websocket server is closed")
}
if s.connections.Size() >= s.opts.maxConnections {
return nil, errors.New("too many online websocket connections")
}
wsConn, err := s.upgrader.Upgrade(r.Response.Writer, r.Request, nil)
if err != nil {
return nil, fmt.Errorf("upgrade failed: %w", err)
}
headers := make(map[string]string)
for k, v := range r.Request.Header {
if len(v) > 0 {
headers[k] = v[0]
}
}
key := s.opts.connKeyPrefix + sessionId
// 踢下线旧连接
s.kickOld(key)
baseCtx := context.WithoutCancel(ctx)
closeCtx, closeCancel := context.WithCancel(baseCtx)
wc := &WsConnection{
SessionId: sessionId,
Conn: wsConn,
Headers: headers,
closeCancel: closeCancel,
closed: 0,
}
s.connections.Set(key, wc)
// 连接成功回执
_ = s.writeJSON(closeCtx, wc, &WsPushMsg{Type: "ack", Message: "WebSocket连接成功", Data: map[string]any{
"sessionId": sessionId,
}})
// Pong 心跳回调,重置读超时
wsConn.SetPongHandler(func(string) error {
_ = wsConn.SetReadDeadline(time.Now().Add(s.opts.readTimeout))
return nil
})
go s.handleConnection(closeCtx, key, wc)
return wc, nil
}
// PushToSession 向指定会话推送消息
func (s *WsServer) PushToSession(ctx context.Context, sessionId string, msg *WsPushMsg) {
key := s.opts.connKeyPrefix + sessionId
val := s.connections.Get(key)
if val == nil {
return
}
wc, ok := val.(*WsConnection)
if !ok || wc.IsClosed() {
return
}
_ = s.writeJSON(ctx, wc, msg)
}
// GetOnlineSessions 获取在线会话列表
func (s *WsServer) GetOnlineSessions() []string {
var sessions []string
prefixLen := len(s.opts.connKeyPrefix)
s.connections.Iterator(func(key string, _ interface{}) bool {
if len(key) > prefixLen {
sessions = append(sessions, key[prefixLen:])
}
return true
})
return sessions
}
// Close 全局优雅关闭
func (s *WsServer) Close() {
s.closeOnce.Do(func() {
atomic.StoreInt32(&s.svcClosed, 1)
s.workerPool.Close()
s.connections.LockFunc(func(m map[string]interface{}) {
for _, val := range m {
wc, ok := val.(*WsConnection)
if !ok {
continue
}
if atomic.CompareAndSwapInt32(&wc.closed, 0, 1) {
if wc.closeCancel != nil {
wc.closeCancel()
}
_ = wc.Conn.Close()
}
}
})
s.connections.Clear()
})
}
// ====================== 内部方法 ======================
// kickOld 踢掉同session旧连接,不再主动remove,由旧连接defer清理
func (s *WsServer) kickOld(key string) {
val := s.connections.Get(key)
if val == nil {
return
}
old, ok := val.(*WsConnection)
if !ok {
return
}
if atomic.CompareAndSwapInt32(&old.closed, 0, 1) {
if old.closeCancel != nil {
old.closeCancel()
}
_ = old.Conn.Close()
}
}
// heartbeatLoop 心跳发送协程,入参改为 *WsConnection,复用写锁
func (s *WsServer) heartbeatLoop(ctx context.Context, wc *WsConnection, done <-chan struct{}) {
ticker := time.NewTicker(s.opts.heartbeatInterval)
defer ticker.Stop()
conn := wc.Conn
for {
select {
case <-ticker.C:
wc.writeMu.Lock()
_ = conn.SetWriteDeadline(time.Now().Add(s.opts.writeTimeout))
err := conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(s.opts.writeTimeout))
wc.writeMu.Unlock()
if err != nil {
glog.Debugf(ctx, "heartbeat ping failed: %v", err)
return
}
case <-done:
return
case <-ctx.Done():
return
}
}
}
func (s *WsServer) handleConnection(ctx context.Context, key string, wc *WsConnection) {
conn := wc.Conn
defer func() {
if atomic.CompareAndSwapInt32(&wc.closed, 0, 1) {
if wc.closeCancel != nil {
wc.closeCancel()
}
_ = conn.Close()
}
// 关键修复:只删除自身实例,防止旧连接误删新连接
s.connections.LockFunc(func(m map[string]interface{}) {
if v, exist := m[key]; exist && v == wc {
delete(m, key)
}
})
}()
done := make(chan struct{})
defer close(done)
go s.heartbeatLoop(ctx, wc, done)
_ = conn.SetReadDeadline(time.Now().Add(s.opts.readTimeout))
for {
select {
case <-ctx.Done():
return
default:
}
msgType, data, err := conn.ReadMessage()
if err != nil {
// 正常关闭不打error日志
if !websocket.IsUnexpectedCloseError(err,
websocket.CloseNormalClosure,
websocket.CloseGoingAway,
websocket.CloseNoStatusReceived,
) {
glog.Debugf(ctx, "normal close: %s, err: %v", key, err)
} else {
glog.Infof(ctx, "unexpected close: %s, err: %v", key, err)
}
break
}
_ = conn.SetReadDeadline(time.Now().Add(s.opts.readTimeout))
switch msgType {
case websocket.PingMessage:
wc.writeMu.Lock()
_ = conn.SetWriteDeadline(time.Now().Add(s.opts.writeTimeout))
_ = conn.WriteMessage(websocket.PongMessage, nil)
wc.writeMu.Unlock()
continue
case websocket.CloseMessage:
return
case websocket.BinaryMessage, websocket.TextMessage:
default:
continue
}
if len(data) == 0 {
continue
}
var msg WsMessage
if err := gjson.Unmarshal(data, &msg); err != nil {
_ = s.writeJSON(ctx, wc, &WsPushMsg{Type: "error", Message: "消息格式错误", Error: err.Error()})
continue
}
s.handlerMu.RLock()
handler, exists := s.handlers[msg.Type]
s.handlerMu.RUnlock()
if !exists {
_ = s.writeJSON(ctx, wc, &WsPushMsg{Type: "error", Message: fmt.Sprintf("未知消息类型: %s", msg.Type)})
continue
}
// 【重要修复】投递到workerPool,避免业务阻塞读循环
taskCtx := ctx
payload := msg.Payload
if err := s.workerPool.Add(taskCtx, func(ctx context.Context) {
handler(ctx, wc, payload)
}); err != nil {
_ = s.writeJSON(ctx, wc, &WsPushMsg{
Type: "error",
Message: "服务繁忙,任务队列已满",
})
}
}
}
// writeJSON 统一写入消息,入参改为 *WsConnection,带并发写锁
func (s *WsServer) writeJSON(ctx context.Context, wc *WsConnection, data interface{}) error {
wc.writeMu.Lock()
defer wc.writeMu.Unlock()
jsonBytes, err := gjson.Encode(data)
if err != nil {
glog.Errorf(ctx, "json encode failed: %v", err)
return err
}
_ = wc.Conn.SetWriteDeadline(time.Now().Add(s.opts.writeTimeout))
if err = wc.Conn.WriteMessage(websocket.TextMessage, jsonBytes); err != nil {
glog.Debugf(ctx, "websocket write failed: %v", err)
return err
}
return nil
}