* 新增 tools 包:统一工具定义、注册表与 Server 接口,对齐 MCP 规范 * 新增 websocket 包:泛化连接管理、心跳、并发写锁与优雅关闭 * 新增参数读取工具函数,避免类型断言静默失败 * 新增 OSS 路径识别与 JSON 扁平映射还原工具 * 修复租户 SQL 条件插入位置,正确处理 GROUP BY 与 ORDER BY 同时出现的场景
327 lines
7.8 KiB
Go
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
|
|
}
|