Files
Meshray-Manager/core/connect/ws.go
T
2026-06-30 15:14:37 +08:00

217 lines
4.3 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}