package connect import ( "context" "fmt" "net" "sync" "time" "github.com/gorilla/websocket" "go.uber.org/zap" ) // WSClient WebSocket 客户端 type WSClient struct { servers []string logger *zap.Logger } // NewWSClient 创建 WebSocket 客户端 func NewWSClient(servers []string, logger *zap.Logger) *WSClient { return &WSClient{ servers: servers, logger: logger, } } // Connect 连接到 WebSocket 服务器 func (c *WSClient) Connect(ctx context.Context) (net.Conn, error) { for _, server := range c.servers { conn, err := c.connectServer(ctx, server) if err == nil { return conn, nil } c.logger.Warn("WebSocket 服务器连接失败", zap.String("server", server), zap.Error(err)) } return nil, fmt.Errorf("所有 WebSocket 服务器均连接失败") } // connectServer 连接单个服务器 func (c *WSClient) connectServer(ctx context.Context, server string) (net.Conn, error) { dialer := websocket.Dialer{ HandshakeTimeout: 10 * time.Second, } ctx, cancel := context.WithTimeout(ctx, 10*time.Second) defer cancel() wsConn, _, err := dialer.DialContext(ctx, server, nil) if err != nil { return nil, fmt.Errorf("WebSocket 握手失败: %w", err) } c.logger.Debug("WebSocket 连接已建立", zap.String("local_addr", wsConn.LocalAddr().String()), zap.String("remote_addr", wsConn.RemoteAddr().String())) return NewWSConn(wsConn, c.logger), nil } // WSFactory WebSocket 传输工厂 type WSFactory struct { client *WSClient logger *zap.Logger } // NewWSFactory 创建 WebSocket 工厂 func NewWSFactory(servers []string, logger *zap.Logger) *WSFactory { return &WSFactory{ client: NewWSClient(servers, logger), logger: logger, } } // Layer 返回传输层类型 func (f *WSFactory) Layer() Layer { return LayerWS } // Name 返回名称 func (f *WSFactory) Name() string { return "WS/WSS" } // Dial 建立 WebSocket 连接 func (f *WSFactory) Dial(ctx context.Context, config *DialConfig) (net.Conn, error) { f.logger.Info("开始建立 WebSocket 连接", zap.String("peer_id", config.PeerID), zap.Strings("ws_servers", config.WSServers)) servers := config.WSServers if len(servers) == 0 { servers = f.client.servers } if len(servers) == 0 { return nil, fmt.Errorf("未配置 WebSocket 服务器") } f.client.servers = servers return f.client.Connect(ctx) } // WSConn WebSocket 连接包装器(实现 net.Conn) type WSConn struct { conn *websocket.Conn localAddr net.Addr remoteAddr net.Addr readBuf []byte mu sync.Mutex closed bool logger *zap.Logger } // NewWSConn 创建 WebSocket net.Conn 包装器 func NewWSConn(wsConn *websocket.Conn, logger *zap.Logger) *WSConn { return &WSConn{ conn: wsConn, localAddr: wsConn.LocalAddr(), remoteAddr: wsConn.RemoteAddr(), readBuf: make([]byte, 0), logger: logger, } } // Read 实现 net.Conn func (c *WSConn) Read(b []byte) (n int, err error) { c.mu.Lock() defer c.mu.Unlock() if c.closed { return 0, net.ErrClosed } if len(c.readBuf) > 0 { n = copy(b, c.readBuf) c.readBuf = c.readBuf[n:] return n, nil } _, message, err := c.conn.ReadMessage() if err != nil { return 0, err } n = copy(b, message) if n < len(message) { c.readBuf = append(c.readBuf, message[n:]...) } return n, nil } // Write 实现 net.Conn func (c *WSConn) Write(b []byte) (n int, err error) { c.mu.Lock() defer c.mu.Unlock() if c.closed { return 0, net.ErrClosed } err = c.conn.WriteMessage(websocket.BinaryMessage, b) if err != nil { return 0, err } return len(b), nil } // Close 实现 net.Conn func (c *WSConn) Close() error { c.mu.Lock() defer c.mu.Unlock() if c.closed { return nil } c.closed = true return c.conn.Close() } // LocalAddr 实现 net.Conn func (c *WSConn) LocalAddr() net.Addr { return c.localAddr } // RemoteAddr 实现 net.Conn func (c *WSConn) RemoteAddr() net.Addr { return c.remoteAddr } // SetDeadline 实现 net.Conn func (c *WSConn) SetDeadline(t time.Time) error { return c.conn.SetReadDeadline(t) } // SetReadDeadline 实现 net.Conn func (c *WSConn) SetReadDeadline(t time.Time) error { return c.conn.SetReadDeadline(t) } // SetWriteDeadline 实现 net.Conn func (c *WSConn) SetWriteDeadline(t time.Time) error { return c.conn.SetWriteDeadline(t) } // IsClosed 检查是否已关闭 func (c *WSConn) IsClosed() bool { c.mu.Lock() defer c.mu.Unlock() return c.closed }