100 lines
2.0 KiB
Go
100 lines
2.0 KiB
Go
package transport
|
|
|
|
import (
|
|
"net"
|
|
"sync"
|
|
|
|
"go.uber.org/zap"
|
|
)
|
|
|
|
// ConnManager 连接管理器
|
|
// 维护 peer_key → net.Conn 的映射关系
|
|
type ConnManager struct {
|
|
conns map[string]net.Conn
|
|
mu sync.RWMutex
|
|
logger *zap.Logger
|
|
}
|
|
|
|
// NewConnManager 创建连接管理器
|
|
func NewConnManager(logger *zap.Logger) *ConnManager {
|
|
return &ConnManager{
|
|
conns: make(map[string]net.Conn),
|
|
logger: logger,
|
|
}
|
|
}
|
|
|
|
// Add 添加连接
|
|
func (m *ConnManager) Add(peerKey string, conn net.Conn) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
|
|
// 如果已存在,先关闭旧连接
|
|
if oldConn, ok := m.conns[peerKey]; ok {
|
|
oldConn.Close()
|
|
m.logger.Debug("关闭旧连接", zap.String("peer_key", peerKey))
|
|
}
|
|
|
|
m.conns[peerKey] = conn
|
|
m.logger.Info("添加新连接",
|
|
zap.String("peer_key", peerKey),
|
|
zap.String("remote_addr", conn.RemoteAddr().String()))
|
|
}
|
|
|
|
// Get 获取连接
|
|
func (m *ConnManager) Get(peerKey string) (net.Conn, bool) {
|
|
m.mu.RLock()
|
|
defer m.mu.RUnlock()
|
|
|
|
conn, ok := m.conns[peerKey]
|
|
return conn, ok
|
|
}
|
|
|
|
// Remove 移除连接
|
|
func (m *ConnManager) Remove(peerKey string) {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
|
|
if conn, ok := m.conns[peerKey]; ok {
|
|
conn.Close()
|
|
delete(m.conns, peerKey)
|
|
m.logger.Info("移除连接", zap.String("peer_key", peerKey))
|
|
}
|
|
}
|
|
|
|
// Count 获取连接数量
|
|
func (m *ConnManager) Count() int {
|
|
m.mu.RLock()
|
|
defer m.mu.RUnlock()
|
|
return len(m.conns)
|
|
}
|
|
|
|
// CloseAll 关闭所有连接
|
|
func (m *ConnManager) CloseAll() {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
|
|
for peerKey, conn := range m.conns {
|
|
conn.Close()
|
|
m.logger.Debug("关闭连接", zap.String("peer_key", peerKey))
|
|
}
|
|
|
|
m.conns = make(map[string]net.Conn)
|
|
}
|
|
|
|
// List 列出所有连接(返回副本)
|
|
func (m *ConnManager) List() map[string]net.Conn {
|
|
m.mu.RLock()
|
|
defer m.mu.RUnlock()
|
|
|
|
result := make(map[string]net.Conn)
|
|
for k, v := range m.conns {
|
|
result[k] = v
|
|
}
|
|
return result
|
|
}
|
|
|
|
// GetAll 获取所有连接(同 List,为了兼容)
|
|
func (m *ConnManager) GetAll() map[string]net.Conn {
|
|
return m.List()
|
|
}
|