Initial commit
This commit is contained in:
@@ -0,0 +1,216 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user