Files

124 lines
2.8 KiB
Go

package wsc
import (
"net/http"
"sync"
"sync/atomic"
"github.com/gin-gonic/gin"
"github.com/gorilla/websocket"
)
// ============================================================
// 默认配置
// ============================================================
var defaultUpgrader = websocket.Upgrader{
ReadBufferSize: 4096,
WriteBufferSize: 4096,
CheckOrigin: func(r *http.Request) bool {
return true // 生产环境请限制 Origin
},
}
// ============================================================
// Server — 连接管理器
// ============================================================
// Server 管理 WebSocket 连接的全局状态。
type Server struct {
upgrader websocket.Upgrader
// 连接生命周期回调
OnConnect func(ctx *Context)
OnDisconnect func(ctx *Context)
// OnUpgrade 升级前钩子:可用于鉴权(如校验 URL token),返回 error 即拒绝连接(401)。
// 钩子中通过 c.Set(k, v) 写入的值会自动透传到 *Context(见 newContext)。
OnUpgrade func(c *gin.Context) error
// 房间管理
mu sync.RWMutex
rooms map[string]*Room
// 连接统计
connCount atomic.Int64
}
// NewServer 创建 Server。
func NewServer() *Server {
return &Server{
upgrader: defaultUpgrader,
rooms: make(map[string]*Room),
}
}
// SetUpgrader 自定义升级器。
func (s *Server) SetUpgrader(u websocket.Upgrader) {
s.upgrader = u
}
// ============================================================
// 连接统计
// ============================================================
// ActiveCount 返回当前活跃连接数。
func (s *Server) ActiveCount() int64 {
return s.connCount.Load()
}
func (s *Server) onSessionOpen() {
s.connCount.Add(1)
}
func (s *Server) onSessionClosed() {
s.connCount.Add(-1)
}
// ============================================================
// 房间管理
// ============================================================
// GetOrCreateRoom 获取或创建房间。
func (s *Server) GetOrCreateRoom(name string) *Room {
s.mu.Lock()
defer s.mu.Unlock()
if room, ok := s.rooms[name]; ok {
return room
}
room := newRoom(name)
room.onEmpty = func(name string) {
s.mu.Lock()
defer s.mu.Unlock()
// 二次判空:防止 leave→onEmpty 之间新成员加入
if r, exists := s.rooms[name]; exists && r.Len() == 0 {
delete(s.rooms, name)
}
}
s.rooms[name] = room
return room
}
// GetRoom 获取房间(不存在返回 nil)。
func (s *Server) GetRoom(name string) *Room {
s.mu.RLock()
defer s.mu.RUnlock()
return s.rooms[name]
}
// RemoveRoom 移除房间。
func (s *Server) RemoveRoom(name string) {
s.mu.Lock()
defer s.mu.Unlock()
delete(s.rooms, name)
}
// RoomCount 返回当前房间数。
func (s *Server) RoomCount() int {
s.mu.RLock()
defer s.mu.RUnlock()
return len(s.rooms)
}