217 lines
4.3 KiB
Go
217 lines
4.3 KiB
Go
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
|
||
}
|