feat(pool): 添加SQLite连接池实现并集成到TCP和WebSocket服务
- 新增pool包,包含ConnType连接类型和ConnectionInfo连接信息结构体 - 实现SQLitePool连接池,支持添加、获取、删除、更新连接操作 - 为TCP服务器集成SQLite连接池,存储连接信息到数据库 - 为WebSocket管理器集成SQLite连接池,存储连接信息到数据库 - 在TCP和WebSocket的连接生命周期中同步更新SQLite连接状态 - 添加GetAllConnIDs方法获取所有在线连接ID列表 - 在示例代码中添加错误处理和测试功能
This commit is contained in:
+36
-1
@@ -19,7 +19,10 @@ func NewWs() *Manager {
|
||||
}
|
||||
|
||||
// 2. 创建管理器
|
||||
m := NewManager(customConfig)
|
||||
m, err := NewManager(customConfig)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to create manager: %v", err)
|
||||
}
|
||||
|
||||
// 3. 覆盖业务回调(核心:自定义消息处理逻辑)
|
||||
// 连接建立回调
|
||||
@@ -71,3 +74,35 @@ func main() {
|
||||
log.Println("WebSocket服务启动:http://localhost:8080/ws")
|
||||
log.Fatal(http.ListenAndServe(":8080", nil))
|
||||
}
|
||||
|
||||
// TestWebSocket 测试WebSocket连接
|
||||
func TestWebSocket() {
|
||||
log.Println("=== 测试WebSocket连接 ===")
|
||||
log.Println("1. 创建WebSocket管理器")
|
||||
m, err := NewManager(DefaultConfig())
|
||||
if err != nil {
|
||||
log.Fatalf("创建管理器失败:%v", err)
|
||||
}
|
||||
log.Println("2. 管理器创建成功")
|
||||
log.Println("3. 获取在线连接数")
|
||||
count, err := m.sqlitePool.Count()
|
||||
if err != nil {
|
||||
log.Printf("获取在线连接数失败:%v", err)
|
||||
} else {
|
||||
log.Printf("当前在线连接数:%d", count)
|
||||
}
|
||||
log.Println("4. 获取所有在线连接ID")
|
||||
connIDs, err := m.GetAllConnIDs()
|
||||
if err != nil {
|
||||
log.Printf("获取在线连接ID失败:%v", err)
|
||||
} else {
|
||||
log.Printf("在线连接ID:%v", connIDs)
|
||||
}
|
||||
log.Println("5. 关闭管理器")
|
||||
if err := m.Close(); err != nil {
|
||||
log.Printf("关闭管理器失败:%v", err)
|
||||
} else {
|
||||
log.Println("管理器关闭成功")
|
||||
}
|
||||
log.Println("=== WebSocket测试完成 ===")
|
||||
}
|
||||
|
||||
+87
-12
@@ -9,6 +9,7 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.magicany.cc/black1552/gin-base/pool"
|
||||
"github.com/gogf/gf/v2/encoding/gjson"
|
||||
"github.com/gogf/gf/v2/os/gctx"
|
||||
"github.com/gogf/gf/v2/os/gtime"
|
||||
@@ -20,20 +21,20 @@ import (
|
||||
|
||||
// 常量定义:默认配置
|
||||
const (
|
||||
// DefaultReadBufferSize 默认读写缓冲区大小(字节)
|
||||
// 默认读写缓冲区大小(字节)
|
||||
DefaultReadBufferSize = 1024
|
||||
DefaultWriteBufferSize = 1024
|
||||
// DefaultHeartbeatInterval 默认心跳间隔(秒):每30秒发送一次心跳
|
||||
// 默认心跳间隔(秒):每30秒发送一次心跳
|
||||
DefaultHeartbeatInterval = 30 * time.Second
|
||||
// DefaultHeartbeatTimeout 默认心跳超时(秒):60秒未收到客户端心跳响应则关闭连接
|
||||
// 默认心跳超时(秒):60秒未收到客户端心跳响应则关闭连接
|
||||
DefaultHeartbeatTimeout = 60 * time.Second
|
||||
// DefaultReadTimeout 默认读写超时(秒)
|
||||
// 默认读写超时(秒)
|
||||
DefaultReadTimeout = 60 * time.Second
|
||||
DefaultWriteTimeout = 10 * time.Second
|
||||
// MessageTypeText 消息类型
|
||||
// 消息类型
|
||||
MessageTypeText = websocket.TextMessage
|
||||
MessageTypeBinary = websocket.BinaryMessage
|
||||
// HeartbeatMaxRetry 心跳最大重试次数
|
||||
// 心跳最大重试次数
|
||||
HeartbeatMaxRetry = 3
|
||||
)
|
||||
|
||||
@@ -92,7 +93,8 @@ type Connection struct {
|
||||
type Manager struct {
|
||||
config *Config // 配置
|
||||
upgrader *websocket.Upgrader // HTTP升级器
|
||||
connections map[string]*Connection // 所有在线连接(connID -> Connection)
|
||||
connections map[string]*Connection // 内存中的连接(connID -> Connection)
|
||||
sqlitePool *pool.SQLitePool // SQLite连接池
|
||||
mutex sync.RWMutex // 读写锁(保护connections)
|
||||
// 业务回调:收到消息时触发(用户自定义处理逻辑)
|
||||
OnMessage func(connID string, msgType int, data any)
|
||||
@@ -148,16 +150,16 @@ func (c *Config) Merge(other *Config) *Config {
|
||||
}
|
||||
|
||||
// NewManager 创建连接管理器
|
||||
func NewManager(config *Config) *Manager {
|
||||
func NewManager(config *Config) (*Manager, error) {
|
||||
defaultConfig := DefaultConfig()
|
||||
finalConfig := defaultConfig.Merge(config)
|
||||
// 初始化升级器
|
||||
upgrader := &websocket.Upgrader{
|
||||
ReadBufferSize: config.ReadBufferSize,
|
||||
WriteBufferSize: config.WriteBufferSize,
|
||||
ReadBufferSize: finalConfig.ReadBufferSize,
|
||||
WriteBufferSize: finalConfig.WriteBufferSize,
|
||||
CheckOrigin: func(r *http.Request) bool {
|
||||
// 跨域检查
|
||||
if config.AllowAllOrigins {
|
||||
if finalConfig.AllowAllOrigins {
|
||||
return true
|
||||
}
|
||||
origin := r.Header.Get("Origin")
|
||||
@@ -170,10 +172,17 @@ func NewManager(config *Config) *Manager {
|
||||
},
|
||||
}
|
||||
|
||||
// 初始化SQLite连接池
|
||||
sqlitePool, err := pool.NewSQLitePool()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create sqlite pool: %w", err)
|
||||
}
|
||||
|
||||
return &Manager{
|
||||
config: finalConfig,
|
||||
upgrader: upgrader,
|
||||
connections: make(map[string]*Connection),
|
||||
sqlitePool: sqlitePool,
|
||||
mutex: sync.RWMutex{},
|
||||
// 默认回调(用户可覆盖)
|
||||
OnMessage: func(connID string, msgType int, data any) {
|
||||
@@ -185,7 +194,7 @@ func NewManager(config *Config) *Manager {
|
||||
OnDisconnect: func(connID string, err error) {
|
||||
log.Printf("[默认回调] 连接[%s]已关闭:%v", connID, err)
|
||||
},
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Upgrade HTTP升级为WebSocket连接
|
||||
@@ -237,6 +246,24 @@ func (m *Manager) Upgrade(w http.ResponseWriter, r *http.Request, connID string)
|
||||
m.connections[connID] = wsConn
|
||||
m.mutex.Unlock()
|
||||
|
||||
// 存储到SQLite
|
||||
connInfo := &pool.ConnectionInfo{
|
||||
ID: connID,
|
||||
Type: pool.ConnTypeWebSocket,
|
||||
Address: r.RemoteAddr,
|
||||
IsActive: true,
|
||||
LastUsed: time.Now(),
|
||||
CreatedAt: time.Now(),
|
||||
Data: map[string]interface{}{
|
||||
"origin": r.Header.Get("Origin"),
|
||||
"userAgent": r.Header.Get("User-Agent"),
|
||||
},
|
||||
}
|
||||
if err := m.sqlitePool.Add(connInfo); err != nil {
|
||||
log.Printf("[错误] 存储连接到SQLite失败:%v", err)
|
||||
// 不影响连接建立,仅记录错误
|
||||
}
|
||||
|
||||
// 触发连接建立回调
|
||||
m.OnConnect(connID)
|
||||
|
||||
@@ -282,6 +309,18 @@ func (c *Connection) ReadPump() {
|
||||
return
|
||||
}
|
||||
|
||||
// 更新最后使用时间
|
||||
now := time.Now()
|
||||
// 从SQLite获取连接信息并更新
|
||||
connInfo, err := c.manager.sqlitePool.Get(c.connID)
|
||||
if err == nil && connInfo != nil {
|
||||
connInfo.LastUsed = now
|
||||
if err := c.manager.sqlitePool.Update(connInfo); err != nil {
|
||||
log.Printf("[错误] 更新SQLite连接信息失败:%v", err)
|
||||
// 不影响消息处理,仅记录错误
|
||||
}
|
||||
}
|
||||
|
||||
// 尝试解析JSON格式的心跳消息(精准判断,替代包含判断)
|
||||
isHeartbeat := false
|
||||
// 先尝试解析为JSON对象
|
||||
@@ -369,6 +408,19 @@ func (c *Connection) Send(data []byte) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("发送消息失败:%w", err)
|
||||
}
|
||||
|
||||
// 更新最后使用时间
|
||||
now := time.Now()
|
||||
// 从SQLite获取连接信息并更新
|
||||
connInfo, err := c.manager.sqlitePool.Get(c.connID)
|
||||
if err == nil && connInfo != nil {
|
||||
connInfo.LastUsed = now
|
||||
if err := c.manager.sqlitePool.Update(connInfo); err != nil {
|
||||
log.Printf("[错误] 更新SQLite连接信息失败:%v", err)
|
||||
// 不影响消息发送,仅记录错误
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
}
|
||||
@@ -394,6 +446,12 @@ func (c *Connection) Close(err error) {
|
||||
delete(c.manager.connections, c.connID)
|
||||
c.manager.mutex.Unlock()
|
||||
|
||||
// 从SQLite移除
|
||||
if err := c.manager.sqlitePool.Remove(c.connID); err != nil {
|
||||
log.Printf("[错误] 从SQLite移除连接失败:%v", err)
|
||||
// 不影响连接关闭,仅记录错误
|
||||
}
|
||||
|
||||
// 触发断开回调
|
||||
c.manager.OnDisconnect(c.connID, err)
|
||||
|
||||
@@ -462,12 +520,18 @@ func (m *Manager) GetAllConn() map[string]*Connection {
|
||||
return connCopy
|
||||
}
|
||||
|
||||
// GetConn 获取指定连接
|
||||
func (m *Manager) GetConn(connID string) *Connection {
|
||||
m.mutex.RLock()
|
||||
defer m.mutex.RUnlock()
|
||||
return m.connections[connID]
|
||||
}
|
||||
|
||||
// GetAllConnIDs 获取所有在线连接的ID列表
|
||||
func (m *Manager) GetAllConnIDs() ([]string, error) {
|
||||
return m.sqlitePool.GetAllConnIDs()
|
||||
}
|
||||
|
||||
// CloseAll 关闭所有连接
|
||||
func (m *Manager) CloseAll() {
|
||||
m.mutex.RLock()
|
||||
@@ -486,3 +550,14 @@ func (m *Manager) CloseAll() {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Close 关闭管理器,清理资源
|
||||
func (m *Manager) Close() error {
|
||||
// 关闭所有连接
|
||||
m.CloseAll()
|
||||
// 关闭SQLite连接池
|
||||
if m.sqlitePool != nil {
|
||||
return m.sqlitePool.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user